Stan  1.3
probability, sampling & optimization
 All Classes Namespaces Files Functions Variables Typedefs Enumerator Friends Macros Pages
nuts.hpp
Go to the documentation of this file.
1 #ifndef __STAN__MCMC__NUTS_H__
2 #define __STAN__MCMC__NUTS_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 
20 #include <stan/mcmc/hmc_base.hpp>
21 #include <stan/mcmc/util.hpp>
22 
23 namespace stan {
24 
25  namespace mcmc {
26 
34  template <class BaseRNG = boost::mt19937>
35  class nuts : public hmc_base<BaseRNG> {
36  private:
37 
38  // Stop immediately if H < u - _maxchange
39  const double _maxchange;
40 
41  // Limit tree depth
42  const int _maxdepth;
43 
44  // depth of last sample taken (-1 before any samples)
45  int _lastdepth;
46 
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;
59  stan::math::sub(xplus, xminus, total_direction);
60  return stan::math::dot(total_direction, mminus) > 0
61  && stan::math::dot(total_direction, mplus) > 0;
62  }
63 
64  public:
65 
93  const std::vector<double>& params_r,
94  const std::vector<int>& params_i,
95  int maxdepth = 10,
96  double epsilon = -1,
97  double epsilon_pm = 0.0,
98  bool epsilon_adapt = true,
99  double delta = 0.6,
100  double gamma = 0.05,
101  BaseRNG base_rng = BaseRNG(std::time(0)))
102  : hmc_base<BaseRNG>(model,
103  params_r,
104  params_i,
105  epsilon,
106  epsilon_pm,
107  epsilon_adapt,
108  delta,
109  gamma,
110  base_rng),
111  _maxchange(-1000),
112  _maxdepth(maxdepth),
113  _lastdepth(-1)
114  {
115  // start at 10 * epsilon because NUTS cheaper for larger epsilon
116  this->adaptation_init(10.0);
117  }
118 
124  ~nuts() { }
125 
131  virtual sample next_impl() {
132  // Initialize the algorithm
133  std::vector<double> mminus(this->_model.num_params_r());
134  for (size_t i = 0; i < mminus.size(); ++i)
135  mminus[i] = this->_rand_unit_norm();
136  std::vector<double> mplus(mminus);
137  // The log-joint probability of the momentum and position terms, i.e.
138  // -(kinetic energy + potential energy)
139  double H0 = -0.5 * stan::math::dot_self(mminus) + this->_logp;
140 
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);
145 
146  // Sample the slice variable
147  double u = log(this->_rand_uniform_01()) + H0;
148  int nvalid = 1;
149  int direction = 2 * (this->_rand_uniform_01() > 0.5) - 1;
150  bool criterion = true;
151 
152  // Repeatedly double the set of points we've visited
153  std::vector<double> newx, newgrad, dummy1, dummy2, dummy3;
154  double newlogp = -1;
155  double prob_sum = -1;
156  int newnvalid = -1;
157  int n_considered = 0;
158  // for-loop with depth outside to set lastdepth
159  int depth = 0;
160 
161  double epsilon = this->_epsilon;
162  // only vary epsilon after done adapting
163  if (!this->adapting() && this->varying_epsilon()) {
164  double low = epsilon * (1.0 - this->_epsilon_pm);
165  double high = epsilon * (1.0 + this->_epsilon_pm);
166  double range = high - low;
167  epsilon = low + (range * this->_rand_uniform_01());
168  }
169  this->_epsilon_last = epsilon; // use epsilon_last in tree build
170 
171  while (criterion && (_maxdepth < 0 || depth < _maxdepth)) {
172  direction = 2 * (this->_rand_uniform_01() > 0.5) - 1;
173  if (direction == -1)
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,
177  n_considered);
178  else
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,
182  n_considered);
183  // We can't look at the results of this last doubling if criterion==false
184  if (!criterion)
185  break;
186  criterion = compute_criterion(xplus, xminus, mplus, mminus);
187  // Metropolis-Hastings to determine if we can jump to a point in
188  // the new half-tree
189  if (this->_rand_uniform_01() < float(newnvalid) / (1e-100+float(nvalid))) {
190  this->_x = newx;
191  this->_g = newgrad;
192  this->_logp = newlogp;
193  }
194  nvalid += newnvalid;
195 // fprintf(stderr, "depth = %d, this->_logp = %g\n", depth, this->_logp);
196  ++depth;
197  }
198  _lastdepth = depth;
199 
200  // Now we just have to update epsilon, if adaptation is on.
201  double adapt_stat = prob_sum / float(n_considered);
202  if (this->adapting()) {
203  double adapt_g = adapt_stat - this->_delta;
204  std::vector<double> gvec(1, -adapt_g);
205  std::vector<double> result;
206  this->_da.update(gvec, result);
207  this->_epsilon = exp(result[0]);
208  }
209  std::vector<double> result;
210  this->_da.xbar(result);
211 // fprintf(stderr, "xbar = %f\n", exp(result[0]));
212  double avg_eta = 1.0 / this->n_steps();
213  this->update_mean_stat(avg_eta,adapt_stat);
214 
215  mcmc::sample s(this->_x, this->_z, this->_logp);
216  return s;
217  }
218 
219  virtual void write_sampler_param_names(std::ostream& o) {
220  o << "treedepth__,";
221  if (this->_epsilon_adapt || this->varying_epsilon())
222  o << "stepsize__,";
223  }
224 
225  virtual void write_sampler_params(std::ostream& o) {
226  o << _lastdepth << ',';
227  if (this->_epsilon_adapt || this->varying_epsilon())
228  o << this->_epsilon_last << ',';
229  }
230 
231  virtual void get_sampler_param_names(std::vector<std::string>& names) {
232  names.clear();
233  names.push_back("treedepth__");
234  if (this->_epsilon_adapt || this->varying_epsilon())
235  names.push_back("stepsize__");
236  }
237  virtual void get_sampler_params(std::vector<double>& values) {
238  values.clear();
239  values.push_back(_lastdepth);
240  if (this->_epsilon_adapt || this->varying_epsilon())
241  values.push_back(this->_epsilon_last);
242  }
243 
278  void build_tree(const std::vector<double>& x,
279  const std::vector<double>& m,
280  const std::vector<double>& grad,
281  double u,
282  int direction,
283  int depth,
284  double H0,
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,
293  double& newlogp,
294  int& nvalid,
295  bool& criterion,
296  double& prob_sum,
297  int& n_considered) {
298  if (depth == 0) { // base case
299  xminus = x;
300  gradminus = grad;
301  mminus = m;
302  // FIXME: lepfrog needs +/- this->_epsilon_pm
303  newlogp = leapfrog(this->_model, this->_z, xminus, mminus, gradminus,
304  direction * this->_epsilon_last,
305  this->_error_msgs, this->_output_msgs);
306  newx = xminus;
307  newgrad = gradminus;
308  xplus = xminus;
309  mplus = mminus;
310  gradplus = gradminus;
311  double newH = newlogp - 0.5 * stan::math::dot_self(mminus);
312  if (newH != newH) // treat nan as -inf
313  newH = -std::numeric_limits<double>::infinity();
314  nvalid = newH > u;
315  criterion = newH - u > _maxchange;
316  prob_sum = stan::math::min(1, exp(newH - H0));
317  n_considered = 1;
318  this->nfevals_plus_eq(1);
319  } else { // depth >= 1
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);
323  if (criterion) {
324  std::vector<double> dummy1, dummy2, dummy3;
325  std::vector<double> newx2;
326  std::vector<double> newgrad2;
327  double newlogp2;
328  int nvalid2;
329  bool criterion2;
330  double prob_sum2;
331  int n_considered2;
332  if (direction == -1)
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);
337  else
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);
342  if (criterion &&
343  (this->_rand_uniform_01()
344  < float(nvalid2) / float(nvalid+nvalid2))) {
345  newx = newx2;
346  newgrad = newgrad2;
347  newlogp = newlogp2;
348  }
349  n_considered += n_considered2;
350  prob_sum += prob_sum2;
351  criterion &= criterion2;
352  nvalid += nvalid2;
353  }
354  criterion &= compute_criterion(xplus, xminus, mplus, mminus);
355  }
356  }
357 
358 
359 
360  };
361 
362 
363  }
364 
365 }
366 
367 #endif

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