1 #ifndef __STAN__MCMC__NUTS_DIAG_H__
2 #define __STAN__MCMC__NUTS_DIAG_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>
35 template <
class BaseRNG = boost::mt19937>
40 const double _maxchange;
49 std::vector<double> _step_sizes;
51 std::vector<double> _x_sum;
52 std::vector<double> _xsq_sum;
65 inline static bool compute_criterion(std::vector<double>& xplus,
66 std::vector<double>& xminus,
67 std::vector<double>& mplus,
68 std::vector<double>& mminus,
69 std::vector<double>& step_sizes) {
70 std::vector<double> total_direction;
73 for (
size_t i = 0; i < total_direction.size(); ++i)
74 total_direction[i] /= step_sizes[i];
108 const std::vector<double>& params_r,
109 const std::vector<int>& params_i,
112 double epsilon_pm = 0.0,
113 bool epsilon_adapt =
true,
116 BaseRNG base_rng = BaseRNG(std::time(0)))
129 _step_sizes(model.num_params_r(), 1.0),
130 _x_sum(model.num_params_r(), 0),
131 _xsq_sum(model.num_params_r(), 0),
154 for (
size_t i = 0; i < mminus.size(); ++i)
156 std::vector<double> mplus(mminus);
161 std::vector<double> gradminus(this->
_g);
162 std::vector<double> gradplus(this->
_g);
163 std::vector<double> xminus(this->
_x);
164 std::vector<double> xplus(this->
_x);
170 bool criterion =
true;
173 std::vector<double> newx, newgrad, dummy1, dummy2, dummy3;
175 double prob_sum = -1;
177 int n_considered = 0;
186 double range = high - low;
191 while (criterion && (_maxdepth < 0 || depth < _maxdepth)) {
194 build_tree(xminus, mminus, gradminus, u, direction, depth,
195 H0, xminus, mminus, gradminus, dummy1, dummy2, dummy3,
196 newx, newgrad, newlogp, newnvalid, criterion, prob_sum,
199 build_tree(xplus, mplus, gradplus, u, direction, depth,
200 H0, dummy1, dummy2, dummy3, xplus, mplus, gradplus,
201 newx, newgrad, newlogp, newnvalid, criterion, prob_sum,
206 criterion = compute_criterion(xplus, xminus, mplus, mminus,_step_sizes);
212 this->
_logp = newlogp;
221 double adapt_stat = prob_sum / float(n_considered);
224 double adapt_g = adapt_stat - this->
_delta;
225 std::vector<double> gvec(1, -adapt_g);
226 std::vector<double> result;
231 _next_diag_adapt *= 2;
232 double step_size_sq_sum = 0;
233 for (
size_t i = 0; i < _step_sizes.size(); i++) {
234 double Ex = _x_sum[i] / _x_sum_n;
235 double Exsq = _xsq_sum[i] / _x_sum_n;
238 _step_sizes[i] =
sqrt(Exsq - Ex*Ex);
239 step_size_sq_sum += _step_sizes[i] * _step_sizes[i];
241 if (step_size_sq_sum > 0.0) {
243 double normalizer =
sqrt((
double)_step_sizes.size())
244 /
sqrt(step_size_sq_sum);
245 for (
size_t i = 0; i < _step_sizes.size(); i++)
246 _step_sizes[i] *= normalizer;
248 for (
size_t i = 0; i < _step_sizes.size(); i++)
249 _step_sizes[i] = 1.0;
253 std::vector<double> result;
255 double avg_eta = 1.0 / this->
n_steps();
268 o << _lastdepth <<
',';
274 o <<
"# (mcmc::nuts_diag) adaptation finished" <<
'\n';
275 o <<
"# step size=" << this->
_epsilon <<
'\n';
276 o <<
"# parameter step size multipliers:\n";
278 for (
size_t k = 0; k < _step_sizes.size(); ++k) {
291 names.push_back(
"treedepth__");
293 names.push_back(
"stepsize__");
298 values.push_back(_lastdepth);
338 const std::vector<double>& m,
339 const std::vector<double>&
grad,
344 std::vector<double>& xminus,
345 std::vector<double>& mminus,
346 std::vector<double>& gradminus,
347 std::vector<double>& xplus,
348 std::vector<double>& mplus,
349 std::vector<double>& gradplus,
350 std::vector<double>& newx,
351 std::vector<double>& newgrad,
362 xminus, mminus, gradminus,
369 gradplus = gradminus;
372 newH = -std::numeric_limits<double>::infinity();
374 criterion = newH - u > _maxchange;
381 for (
size_t i = 0; i < newx.size(); i++) {
382 _x_sum[i] += newx[i];
383 _xsq_sum[i] += newx[i] * newx[i];
387 build_tree(x, m, grad, u, direction, depth-1, H0, xminus, mminus,
388 gradminus, xplus, mplus, gradplus, newx, newgrad, newlogp,
389 nvalid, criterion, prob_sum, n_considered);
391 std::vector<double> dummy1, dummy2, dummy3;
392 std::vector<double> newx2;
393 std::vector<double> newgrad2;
400 build_tree(xminus, mminus, gradminus, u, direction, depth-1, H0,
401 xminus, mminus, gradminus, dummy1, dummy2, dummy3,
402 newx2, newgrad2, newlogp2, nvalid2, criterion2,
403 prob_sum2, n_considered2);
405 build_tree(xplus, mplus, gradplus, u, direction, depth-1, H0,
406 dummy1, dummy2, dummy3, xplus, mplus, gradplus,
407 newx2, newgrad2, newlogp2, nvalid2, criterion2,
408 prob_sum2, n_considered2);
411 <
float(nvalid2) /
float(nvalid+nvalid2))){
416 n_considered += n_considered2;
417 prob_sum += prob_sum2;
418 criterion &= criterion2;
421 criterion &= compute_criterion(xplus, xminus, mplus, mminus, _step_sizes);