1 #ifndef __STAN__MCMC__NUTS_MASSGIVEN_H__
2 #define __STAN__MCMC__NUTS_MASSGIVEN_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>
19 #include <Eigen/Dense>
33 template <
class BaseRNG = boost::mt19937>
38 const double _maxchange;
48 Eigen::MatrixXd _cov_L;
66 inline bool compute_criterion(std::vector<double>& xplus,
67 std::vector<double>& xminus,
68 std::vector<double>& mplus,
69 std::vector<double>& mminus) {
70 std::vector<double> total_direction;
72 Eigen::Map<Eigen::VectorXd> total_direction_mat(&total_direction[0],total_direction.size());
73 Eigen::Map<Eigen::VectorXd> mplus_mat(&mplus[0],mplus.size());
74 Eigen::Map<Eigen::VectorXd> mminus_mat(&mminus[0],mminus.size());
75 _cov_L.triangularView<Eigen::Lower>().solveInPlace(total_direction_mat);
76 return total_direction_mat.dot(mplus_mat) > 0
77 && total_direction_mat.dot(mminus_mat) > 0;
109 const std::vector<double>& params_r,
110 const std::vector<int>& params_i,
111 std::string cov_file,
114 double epsilon_pm = 0.0,
115 bool epsilon_adapt =
true,
118 BaseRNG base_rng = BaseRNG(std::time(0)))
161 for (
size_t i = 0; i < mminus.size(); ++i)
163 std::vector<double> mplus(mminus);
168 std::vector<double> gradminus(this->
_g);
169 std::vector<double> gradplus(this->
_g);
170 std::vector<double> xminus(this->
_x);
171 std::vector<double> xplus(this->
_x);
177 bool criterion =
true;
180 std::vector<double> newx, newgrad, dummy1, dummy2, dummy3;
182 double prob_sum = -1;
184 int n_considered = 0;
193 double range = high - low;
198 while (criterion && (_maxdepth < 0 || depth <= _maxdepth)) {
201 build_tree(xminus, mminus, gradminus, u, direction, depth,
202 H0, xminus, mminus, gradminus, dummy1, dummy2, dummy3,
203 newx, newgrad, newlogp, newnvalid, criterion, prob_sum,
206 build_tree(xplus, mplus, gradplus, u, direction, depth,
207 H0, dummy1, dummy2, dummy3, xplus, mplus, gradplus,
208 newx, newgrad, newlogp, newnvalid, criterion, prob_sum,
213 criterion = compute_criterion(xplus, xminus, mplus, mminus);
219 this->
_logp = newlogp;
228 double adapt_stat = prob_sum / float(n_considered);
231 double adapt_g = adapt_stat - this->
_delta;
232 std::vector<double> gvec(1, -adapt_g);
233 std::vector<double> result;
239 std::vector<double> result;
241 double avg_eta = 1.0 / this->
n_steps();
254 o << _lastdepth <<
',';
260 o <<
"# (mcmc::nuts_massgiven) adaptation finished" <<
'\n';
261 o <<
"# step size=" << this->
_epsilon <<
'\n';
262 o <<
"# Preset covariance matrix:\n";
263 Eigen::MatrixXd _cov_mat = _cov_L * _cov_L.transpose();
264 for(
int i=0; i<_cov_mat.rows(); i++){
266 for(
int j=0; j<_cov_mat.cols(); j++)
267 o << _cov_mat(i,j) <<
",";
274 names.push_back(
"treedepth__");
276 names.push_back(
"stepsize__");
281 values.push_back(_lastdepth);
321 const std::vector<double>& m,
322 const std::vector<double>&
grad,
327 std::vector<double>& xminus,
328 std::vector<double>& mminus,
329 std::vector<double>& gradminus,
330 std::vector<double>& xplus,
331 std::vector<double>& mplus,
332 std::vector<double>& gradplus,
333 std::vector<double>& newx,
334 std::vector<double>& newgrad,
345 xminus, mminus, gradminus,
352 gradplus = gradminus;
355 newH = -std::numeric_limits<double>::infinity();
357 criterion = newH - u > _maxchange;
364 build_tree(x, m, grad, u, direction, depth-1, H0, xminus, mminus,
365 gradminus, xplus, mplus, gradplus, newx, newgrad, newlogp,
366 nvalid, criterion, prob_sum, n_considered);
368 std::vector<double> dummy1, dummy2, dummy3;
369 std::vector<double> newx2;
370 std::vector<double> newgrad2;
377 build_tree(xminus, mminus, gradminus, u, direction, depth-1, H0,
378 xminus, mminus, gradminus, dummy1, dummy2, dummy3,
379 newx2, newgrad2, newlogp2, nvalid2, criterion2,
380 prob_sum2, n_considered2);
382 build_tree(xplus, mplus, gradplus, u, direction, depth-1, H0,
383 dummy1, dummy2, dummy3, xplus, mplus, gradplus,
384 newx2, newgrad2, newlogp2, nvalid2, criterion2,
385 prob_sum2, n_considered2);
388 <
float(nvalid2) /
float(nvalid+nvalid2))){
393 n_considered += n_considered2;
394 prob_sum += prob_sum2;
395 criterion &= criterion2;
398 criterion &= compute_criterion(xplus, xminus, mplus, mminus);