1 #ifndef __STAN__MCMC__NUTS_H__
2 #define __STAN__MCMC__NUTS_H__
9 #include <boost/random/normal_distribution.hpp>
10 #include <boost/random/mersenne_twister.hpp>
11 #include <boost/random/variate_generator.hpp>
12 #include <boost/random/uniform_01.hpp>
34 template <
class BaseRNG = boost::mt19937>
39 const double _maxchange;
54 inline static bool compute_criterion(std::vector<double>& xplus,
55 std::vector<double>& xminus,
56 std::vector<double>& mplus,
57 std::vector<double>& mminus) {
58 std::vector<double> total_direction;
93 const std::vector<double>& params_r,
94 const std::vector<int>& params_i,
97 double epsilon_pm = 0.0,
98 bool epsilon_adapt =
true,
101 BaseRNG base_rng = BaseRNG(std::time(0)))
134 for (
size_t i = 0; i < mminus.size(); ++i)
136 std::vector<double> mplus(mminus);
141 std::vector<double> gradminus(this->
_g);
142 std::vector<double> gradplus(this->
_g);
143 std::vector<double> xminus(this->
_x);
144 std::vector<double> xplus(this->
_x);
150 bool criterion =
true;
153 std::vector<double> newx, newgrad, dummy1, dummy2, dummy3;
155 double prob_sum = -1;
157 int n_considered = 0;
166 double range = high - low;
171 while (criterion && (_maxdepth < 0 || depth < _maxdepth)) {
174 build_tree(xminus, mminus, gradminus, u, direction, depth,
175 H0, xminus, mminus, gradminus, dummy1, dummy2, dummy3,
176 newx, newgrad, newlogp, newnvalid, criterion, prob_sum,
179 build_tree(xplus, mplus, gradplus, u, direction, depth,
180 H0, dummy1, dummy2, dummy3, xplus, mplus, gradplus,
181 newx, newgrad, newlogp, newnvalid, criterion, prob_sum,
186 criterion = compute_criterion(xplus, xminus, mplus, mminus);
192 this->
_logp = newlogp;
201 double adapt_stat = prob_sum / float(n_considered);
203 double adapt_g = adapt_stat - this->
_delta;
204 std::vector<double> gvec(1, -adapt_g);
205 std::vector<double> result;
209 std::vector<double> result;
212 double avg_eta = 1.0 / this->
n_steps();
226 o << _lastdepth <<
',';
233 names.push_back(
"treedepth__");
235 names.push_back(
"stepsize__");
239 values.push_back(_lastdepth);
279 const std::vector<double>& m,
280 const std::vector<double>&
grad,
285 std::vector<double>& xminus,
286 std::vector<double>& mminus,
287 std::vector<double>& gradminus,
288 std::vector<double>& xplus,
289 std::vector<double>& mplus,
290 std::vector<double>& gradplus,
291 std::vector<double>& newx,
292 std::vector<double>& newgrad,
310 gradplus = gradminus;
313 newH = -std::numeric_limits<double>::infinity();
315 criterion = newH - u > _maxchange;
320 build_tree(x, m, grad, u, direction, depth-1, H0, xminus, mminus,
321 gradminus, xplus, mplus, gradplus, newx, newgrad, newlogp,
322 nvalid, criterion, prob_sum, n_considered);
324 std::vector<double> dummy1, dummy2, dummy3;
325 std::vector<double> newx2;
326 std::vector<double> newgrad2;
333 build_tree(xminus, mminus, gradminus, u, direction, depth-1, H0,
334 xminus, mminus, gradminus, dummy1, dummy2, dummy3,
335 newx2, newgrad2, newlogp2, nvalid2, criterion2,
336 prob_sum2, n_considered2);
338 build_tree(xplus, mplus, gradplus, u, direction, depth-1, H0,
339 dummy1, dummy2, dummy3, xplus, mplus, gradplus,
340 newx2, newgrad2, newlogp2, nvalid2, criterion2,
341 prob_sum2, n_considered2);
344 <
float(nvalid2) /
float(nvalid+nvalid2))) {
349 n_considered += n_considered2;
350 prob_sum += prob_sum2;
351 criterion &= criterion2;
354 criterion &= compute_criterion(xplus, xminus, mplus, mminus);