Stan  1.3
probability, sampling & optimization
 All Classes Namespaces Files Functions Variables Typedefs Enumerator Friends Macros Pages
nuts_diag.hpp
Go to the documentation of this file.
1 #ifndef __STAN__MCMC__NUTS_DIAG_H__
2 #define __STAN__MCMC__NUTS_DIAG_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 
35  template <class BaseRNG = boost::mt19937>
36  class nuts_diag : public hmc_base<BaseRNG> {
37  private:
38 
39  // Stop immediately if H < u - _maxchange
40  const double _maxchange;
41 
42  // Limit tree depth
43  const int _maxdepth;
44 
45  // Depth of last sample taken (-1 before any samples)
46  int _lastdepth;
47 
48  // Vector of per-parameter step sizes.
49  std::vector<double> _step_sizes;
50  // Running statistics to estimate per-coordinate std. deviations.
51  std::vector<double> _x_sum;
52  std::vector<double> _xsq_sum;
53  int _x_sum_n;
54  // Next time we should adapt the per-parameter step sizes.
55  int _next_diag_adapt;
56 
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;
71  stan::math::sub(xplus, xminus, total_direction);
72  // adjustment for U-turn due to step sizes
73  for (size_t i = 0; i < total_direction.size(); ++i)
74  total_direction[i] /= step_sizes[i];
75  return stan::math::dot(total_direction, mminus) > 0
76  && stan::math::dot(total_direction, mplus) > 0;
77  }
78 
79  public:
80 
108  const std::vector<double>& params_r,
109  const std::vector<int>& params_i,
110  int maxdepth = 10,
111  double epsilon = -1,
112  double epsilon_pm = 0.0,
113  bool epsilon_adapt = true,
114  double delta = 0.6,
115  double gamma = 0.05,
116  BaseRNG base_rng = BaseRNG(std::time(0)))
117  : hmc_base<BaseRNG>(model,
118  params_r,
119  params_i,
120  epsilon,
121  epsilon_pm,
122  epsilon_adapt,
123  delta,
124  gamma,
125  base_rng),
126  _maxchange(-1000),
127  _maxdepth(maxdepth),
128  _lastdepth(-1),
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),
132  _x_sum_n(0),
133  _next_diag_adapt(10)
134  {
135  // start at 10 * epsilon because NUTS cheaper for larger epsilon
136  this->adaptation_init(10.0);
137  }
138 
145 
151  virtual sample next_impl() {
152  // Initialize the algorithm
153  std::vector<double> mminus(this->_model.num_params_r());
154  for (size_t i = 0; i < mminus.size(); ++i)
155  mminus[i] = this->_rand_unit_norm();
156  std::vector<double> mplus(mminus);
157  // The log-joint probability of the momentum and position terms, i.e.
158  // -(kinetic energy + potential energy)
159  double H0 = -0.5 * stan::math::dot_self(mminus) + this->_logp;
160 
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);
165 
166  // Sample the slice variable
167  double u = log(this->_rand_uniform_01()) + H0;
168  int nvalid = 1;
169  int direction = 2 * (this->_rand_uniform_01() > 0.5) - 1;
170  bool criterion = true;
171 
172  // Repeatedly double the set of points we've visited
173  std::vector<double> newx, newgrad, dummy1, dummy2, dummy3;
174  double newlogp = -1;
175  double prob_sum = -1;
176  int newnvalid = -1;
177  int n_considered = 0;
178  // for-loop with depth outside to set lastdepth
179  int depth = 0;
180 
181  double epsilon = this->_epsilon;
182  // only vary epsilon after done adapting
183  if (!this->adapting() && this->varying_epsilon()) {
184  double low = epsilon * (1.0 - this->_epsilon_pm);
185  double high = epsilon * (1.0 + this->_epsilon_pm);
186  double range = high - low;
187  epsilon = low + (range * this->_rand_uniform_01());
188  }
189  this->_epsilon_last = epsilon; // use epsilon_last in tree build
190 
191  while (criterion && (_maxdepth < 0 || depth < _maxdepth)) {
192  direction = 2 * (this->_rand_uniform_01() > 0.5) - 1;
193  if (direction == -1)
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,
197  n_considered);
198  else
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,
202  n_considered);
203  // We can't look at the results of this last doubling if criterion==false
204  if (!criterion)
205  break;
206  criterion = compute_criterion(xplus, xminus, mplus, mminus,_step_sizes);
207  // Metropolis-Hastings to determine if we can jump to a point in
208  // the new half-tree
209  if (this->_rand_uniform_01() < float(newnvalid) / (1e-100+float(nvalid))) {
210  this->_x = newx;
211  this->_g = newgrad;
212  this->_logp = newlogp;
213  }
214  nvalid += newnvalid;
215  ++depth;
216  }
217  _lastdepth = depth;
218 
219  // Now we just have to update global (epsilon) and local
220  // (step_sizes) step sizes, if adaptation is on.
221  double adapt_stat = prob_sum / float(n_considered);
222  if (this->adapting()) {
223  // epsilon.
224  double adapt_g = adapt_stat - this->_delta;
225  std::vector<double> gvec(1, -adapt_g);
226  std::vector<double> result;
227  this->_da.update(gvec, result);
228  this->_epsilon = exp(result[0]);
229  // step_sizes. Doesn't happen every step.
230  if (this->_n_adapt_steps == _next_diag_adapt) {
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;
236  _x_sum[i] = 0;
237  _xsq_sum[i] = 0;
238  _step_sizes[i] = sqrt(Exsq - Ex*Ex);
239  step_size_sq_sum += _step_sizes[i] * _step_sizes[i];
240  }
241  if (step_size_sq_sum > 0.0) {
242  _x_sum_n = 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;
247  } else {
248  for (size_t i = 0; i < _step_sizes.size(); i++)
249  _step_sizes[i] = 1.0;
250  }
251  }
252  }
253  std::vector<double> result;
254  this->_da.xbar(result);
255  double avg_eta = 1.0 / this->n_steps();
256  this->update_mean_stat(avg_eta,adapt_stat);
257 
258  return mcmc::sample(this->_x, this->_z, this->_logp);
259  }
260 
261  virtual void write_sampler_param_names(std::ostream& o) {
262  o << "treedepth__,";
263  if (this->_epsilon_adapt || this->varying_epsilon())
264  o << "stepsize__,";
265  }
266 
267  virtual void write_sampler_params(std::ostream& o) {
268  o << _lastdepth << ',';
269  if (this->_epsilon_adapt || this->varying_epsilon())
270  o << this->_epsilon_last << ',';
271  }
272 
273  virtual void write_adaptation_params(std::ostream& o) {
274  o << "# (mcmc::nuts_diag) adaptation finished" << '\n';
275  o << "# step size=" << this->_epsilon << '\n';
276  o << "# parameter step size multipliers:\n"; // FIXME: names/delineation requires access to model
277  o << "# ";
278  for (size_t k = 0; k < _step_sizes.size(); ++k) {
279  if (k > 0) o << ',';
280  o << _step_sizes[k];
281  }
282  o << '\n';
283  }
284 
285  std::vector<double> get_step_sizes() {
286  return _step_sizes;
287  }
288 
289  virtual void get_sampler_param_names(std::vector<std::string>& names) {
290  names.clear();
291  names.push_back("treedepth__");
292  if (this->_epsilon_adapt || this->varying_epsilon())
293  names.push_back("stepsize__");
294  }
295 
296  virtual void get_sampler_params(std::vector<double>& values) {
297  values.clear();
298  values.push_back(_lastdepth);
299  if (this->_epsilon_adapt || this->varying_epsilon())
300  values.push_back(this->_epsilon_last);
301  }
302 
337  void build_tree(const std::vector<double>& x,
338  const std::vector<double>& m,
339  const std::vector<double>& grad,
340  double u,
341  int direction,
342  int depth,
343  double H0,
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,
352  double& newlogp,
353  int& nvalid,
354  bool& criterion,
355  double& prob_sum,
356  int& n_considered) {
357  if (depth == 0) { // base case
358  xminus = x;
359  gradminus = grad;
360  mminus = m;
361  newlogp = rescaled_leapfrog(this->_model, this->_z, _step_sizes,
362  xminus, mminus, gradminus,
363  direction * this->_epsilon_last,
364  this->_error_msgs, this->_output_msgs);
365  newx = xminus;
366  newgrad = gradminus;
367  xplus = xminus;
368  mplus = mminus;
369  gradplus = gradminus;
370  double newH = newlogp - 0.5 * stan::math::dot_self(mminus);
371  if (newH != newH) // treat nan as -inf
372  newH = -std::numeric_limits<double>::infinity();
373  nvalid = newH > u;
374  criterion = newH - u > _maxchange;
375  prob_sum = stan::math::min(1, exp(newH - H0));
376  n_considered = 1;
377  this->nfevals_plus_eq(1);
378  // Update running statistics if point is in slice
379  if (nvalid) {
380  _x_sum_n++;
381  for (size_t i = 0; i < newx.size(); i++) {
382  _x_sum[i] += newx[i];
383  _xsq_sum[i] += newx[i] * newx[i];
384  }
385  }
386  } else { // depth >= 1
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);
390  if (criterion) {
391  std::vector<double> dummy1, dummy2, dummy3;
392  std::vector<double> newx2;
393  std::vector<double> newgrad2;
394  double newlogp2;
395  int nvalid2;
396  bool criterion2;
397  double prob_sum2;
398  int n_considered2;
399  if (direction == -1)
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);
404  else
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);
409  if (criterion &&
410  (this->_rand_uniform_01()
411  < float(nvalid2) / float(nvalid+nvalid2))){
412  newx = newx2;
413  newgrad = newgrad2;
414  newlogp = newlogp2;
415  }
416  n_considered += n_considered2;
417  prob_sum += prob_sum2;
418  criterion &= criterion2;
419  nvalid += nvalid2;
420  }
421  criterion &= compute_criterion(xplus, xminus, mplus, mminus, _step_sizes);
422  }
423  }
424 
425 
426  };
427 
428 
429  }
430 
431 }
432 
433 #endif

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