Stan  1.3
probability, sampling & optimization
 All Classes Namespaces Files Functions Variables Typedefs Enumerator Friends Macros Pages
nuts_massgiven.hpp
Go to the documentation of this file.
1 #ifndef __STAN__MCMC__NUTS_MASSGIVEN_H__
2 #define __STAN__MCMC__NUTS_MASSGIVEN_H__
3 
4 #include <ctime>
5 #include <cstddef>
6 #include <iostream>
7 #include <vector>
8 
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>
13 
16 #include <stan/mcmc/hmc_base.hpp>
17 #include <stan/mcmc/util.hpp>
18 
19 #include <Eigen/Dense>
20 
21 namespace stan {
22 
23  namespace mcmc {
24 
33  template <class BaseRNG = boost::mt19937>
34  class nuts_massgiven : public hmc_base<BaseRNG> {
35  private:
36 
37  // Stop immediately if H < u - _maxchange
38  const double _maxchange;
39 
40  // Limit tree depth
41  const int _maxdepth;
42 
43  // Depth of last sample taken (-1 before any samples)
44  int _lastdepth;
45 
46  // Yuanjun : now we use a matrix to get the covariance matrix
47  // Yuanjun : The cholesky decomposition of _cov_mat
48  Eigen::MatrixXd _cov_L;
49  // Running statistics to estimate per-coordinate std. deviations.
50 
51 
52  //Eigen::Map<Eigen::VectorXd> _x_mat;
53  //Eigen::Map<Eigen::VectorXd> _m_mat;
54  //Eigen::Map<Eigen::VectorXd> _g_mat;
55 
56  // Next time we should adapt the per-parameter step sizes.
57 
58 
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;
71  stan::math::sub(xplus, xminus, 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;
78  }
79 
80  public:
81 
109  const std::vector<double>& params_r,
110  const std::vector<int>& params_i,
111  std::string cov_file,
112  int maxdepth = 10,
113  double epsilon = -1,
114  double epsilon_pm = 0.0,
115  bool epsilon_adapt = true,
116  double delta = 0.6,
117  double gamma = 0.05,
118  BaseRNG base_rng = BaseRNG(std::time(0)))
119  : hmc_base<BaseRNG>(model,
120  params_r,
121  params_i,
122  epsilon,
123  epsilon_pm,
124  epsilon_adapt,
125  delta,
126  gamma,
127  base_rng),
128 
129  _maxchange(-1000),
130  _maxdepth(maxdepth),
131  _lastdepth(-1)
132  //_x_mat(&(this->_x[0]), (this->_x).size())
133  {
134  // start at 10 * epsilon because NUTS cheaper for larger epsilon
135  this->adaptation_init(10.0);
136 
137  _cov_L = Eigen::MatrixXd::Identity(model.num_params_r(),model.num_params_r());
138  read_cov(cov_file, _cov_L);
139  // std::cout << _cov_L << "baby" << std::endl;
140 
141  //_x_mat = Eigen::Map<Eigen::VectorXd>(this->_x[0], this->_x.size());
142  //_m_mat = Eigen::Map<Eigen::VectorXd>(this->_m[0], this->_m.size());
143  //_g_mat = Eigen::Map<Eigen::VectorXd>(this->_g[0], this->_g.size());
144  }
145 
152 
158  virtual sample next_impl() {
159  // Initialize the algorithm
160  std::vector<double> mminus(this->_model.num_params_r());
161  for (size_t i = 0; i < mminus.size(); ++i)
162  mminus[i] = this->_rand_unit_norm();
163  std::vector<double> mplus(mminus);
164  // The log-joint probability of the momentum and position terms, i.e.
165  // -(kinetic energy + potential energy)
166  double H0 = -0.5 * stan::math::dot_self(mminus) + this->_logp;
167 
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);
172 
173  // Sample the slice variable
174  double u = log(this->_rand_uniform_01()) + H0;
175  int nvalid = 1;
176  int direction = 2 * (this->_rand_uniform_01() > 0.5) - 1;
177  bool criterion = true;
178 
179  // Repeatedly double the set of points we've visited
180  std::vector<double> newx, newgrad, dummy1, dummy2, dummy3;
181  double newlogp = -1;
182  double prob_sum = -1;
183  int newnvalid = -1;
184  int n_considered = 0;
185  // for-loop with depth outside to set lastdepth
186  int depth = 0;
187 
188  double epsilon = this->_epsilon;
189  // only vary epsilon after done adapting
190  if (!this->adapting() && this->varying_epsilon()) {
191  double low = epsilon * (1.0 - this->_epsilon_pm);
192  double high = epsilon * (1.0 + this->_epsilon_pm);
193  double range = high - low;
194  epsilon = low + (range * this->_rand_uniform_01());
195  }
196  this->_epsilon_last = epsilon; // use epsilon_last in tree build
197 
198  while (criterion && (_maxdepth < 0 || depth <= _maxdepth)) {
199  direction = 2 * (this->_rand_uniform_01() > 0.5) - 1;
200  if (direction == -1)
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,
204  n_considered);
205  else
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,
209  n_considered);
210  // We can't look at the results of this last doubling if criterion==false
211  if (!criterion)
212  break;
213  criterion = compute_criterion(xplus, xminus, mplus, mminus);
214  // Metropolis-Hastings to determine if we can jump to a point in
215  // the new half-tree
216  if (this->_rand_uniform_01() < float(newnvalid) / (1e-100+float(nvalid))) {
217  this->_x = newx;
218  this->_g = newgrad;
219  this->_logp = newlogp;
220  }
221  nvalid += newnvalid;
222  ++depth;
223  }
224  _lastdepth = depth;
225 
226  // Now we just have to update global (epsilon) and local
227  // (step_sizes) step sizes, if adaptation is on.
228  double adapt_stat = prob_sum / float(n_considered);
229  if (this->adapting()) {
230  // epsilon.
231  double adapt_g = adapt_stat - this->_delta;
232  std::vector<double> gvec(1, -adapt_g);
233  std::vector<double> result;
234  this->_da.update(gvec, result);
235  this->_epsilon = exp(result[0]);
236  // step_sizes. Doesn't happen every step.
237 
238  }
239  std::vector<double> result;
240  this->_da.xbar(result);
241  double avg_eta = 1.0 / this->n_steps();
242  this->update_mean_stat(avg_eta,adapt_stat);
243 
244  return mcmc::sample(this->_x, this->_z, this->_logp);
245  }
246 
247  virtual void write_sampler_param_names(std::ostream& o) {
248  o << "treedepth__,";
249  if (this->_epsilon_adapt || this->varying_epsilon())
250  o << "stepsize__,";
251  }
252 
253  virtual void write_sampler_params(std::ostream& o) {
254  o << _lastdepth << ',';
255  if (this->_epsilon_adapt || this->varying_epsilon())
256  o << this->_epsilon_last << ',';
257  }
258 
259  virtual void write_adaptation_params(std::ostream& o) {
260  o << "# (mcmc::nuts_massgiven) adaptation finished" << '\n';
261  o << "# step size=" << this->_epsilon << '\n';
262  o << "# Preset covariance matrix:\n"; // FIXME: names/delineation requires access to model
263  Eigen::MatrixXd _cov_mat = _cov_L * _cov_L.transpose();
264  for(int i=0; i<_cov_mat.rows(); i++){
265  o << "#";
266  for(int j=0; j<_cov_mat.cols(); j++)
267  o << _cov_mat(i,j) << ",";
268  o << std::endl;
269  }
270  }
271 
272  virtual void get_sampler_param_names(std::vector<std::string>& names) {
273  names.clear();
274  names.push_back("treedepth__");
275  if (this->_epsilon_adapt || this->varying_epsilon())
276  names.push_back("stepsize__");
277  }
278 
279  virtual void get_sampler_params(std::vector<double>& values) {
280  values.clear();
281  values.push_back(_lastdepth);
282  if (this->_epsilon_adapt || this->varying_epsilon())
283  values.push_back(this->_epsilon_last);
284  }
285 
320  void build_tree(const std::vector<double>& x,
321  const std::vector<double>& m,
322  const std::vector<double>& grad,
323  double u,
324  int direction,
325  int depth,
326  double H0,
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,
335  double& newlogp,
336  int& nvalid,
337  bool& criterion,
338  double& prob_sum,
339  int& n_considered) {
340  if (depth == 0) { // base case
341  xminus = x;
342  gradminus = grad;
343  mminus = m;
344  newlogp = nondiag_leapfrog(this->_model, this->_z, _cov_L, //Yuanjun : implement a new leapfrog function
345  xminus, mminus, gradminus,
346  direction * this->_epsilon_last,
347  this->_error_msgs, this->_output_msgs);
348  newx = xminus;
349  newgrad = gradminus;
350  xplus = xminus;
351  mplus = mminus;
352  gradplus = gradminus;
353  double newH = newlogp - 0.5 * stan::math::dot_self(mminus); //Yuanjun : No longer need to calculate H with a matrix
354  if (newH != newH) // treat nan as -inf
355  newH = -std::numeric_limits<double>::infinity();
356  nvalid = newH > u;
357  criterion = newH - u > _maxchange;
358  prob_sum = stan::math::min(1, exp(newH - H0));
359  n_considered = 1;
360  this->nfevals_plus_eq(1);
361  // Update running statistics if point is in slice
362 
363  } else { // depth >= 1
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);
367  if (criterion) {
368  std::vector<double> dummy1, dummy2, dummy3;
369  std::vector<double> newx2;
370  std::vector<double> newgrad2;
371  double newlogp2;
372  int nvalid2;
373  bool criterion2;
374  double prob_sum2;
375  int n_considered2;
376  if (direction == -1)
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);
381  else
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);
386  if (criterion &&
387  (this->_rand_uniform_01()
388  < float(nvalid2) / float(nvalid+nvalid2))){
389  newx = newx2;
390  newgrad = newgrad2;
391  newlogp = newlogp2;
392  }
393  n_considered += n_considered2;
394  prob_sum += prob_sum2;
395  criterion &= criterion2;
396  nvalid += nvalid2;
397  }
398  criterion &= compute_criterion(xplus, xminus, mplus, mminus);
399  }
400  }
401 
402 
403  };
404 
405 
406  }
407 
408 }
409 
410 #endif

     [ Stan Home Page ] © 2011–2013, Stan Development Team.