StochTree 0.5.0.9000
Loading...
Searching...
No Matches
variance_model.h
1
5#ifndef STOCHTREE_VARIANCE_MODEL_H_
6#define STOCHTREE_VARIANCE_MODEL_H_
7
8#include <Eigen/Dense>
9#include <stochtree/data.h>
10#include <stochtree/ensemble.h>
11#include <stochtree/gamma_sampler.h>
12#include <stochtree/ig_sampler.h>
13#include <stochtree/meta.h>
14
15#include <random>
16
17namespace StochTree {
18
21 public:
24 double PosteriorShape(Eigen::VectorXd& residuals, double a, double b) {
25 data_size_t n = residuals.rows();
26 return a + (0.5 * n);
27 }
28 double PosteriorScale(Eigen::VectorXd& residuals, double a, double b) {
29 data_size_t n = residuals.rows();
30 double sum_sq_resid = 0.;
31 for (data_size_t i = 0; i < n; i++) {
33 }
34 return b + (0.5 * sum_sq_resid);
35 }
36 double PosteriorShape(Eigen::VectorXd& residuals, Eigen::VectorXd& weights, double a, double b) {
37 data_size_t n = residuals.rows();
38 return a + (0.5 * n);
39 }
40 double PosteriorScale(Eigen::VectorXd& residuals, Eigen::VectorXd& weights, double a, double b) {
41 data_size_t n = residuals.rows();
42 double sum_sq_resid = 0.;
43 for (data_size_t i = 0; i < n; i++) {
45 }
46 return b + (0.5 * sum_sq_resid);
47 }
48 double SampleVarianceParameter(Eigen::VectorXd& residuals, double a, double b, std::mt19937& gen) {
49 double ig_shape = PosteriorShape(residuals, a, b);
50 double ig_scale = PosteriorScale(residuals, a, b);
51 return ig_sampler_.Sample(ig_shape, ig_scale, gen);
52 }
53 double SampleVarianceParameter(Eigen::VectorXd& residuals, Eigen::VectorXd& weights, double a, double b, std::mt19937& gen) {
54 double ig_shape = PosteriorShape(residuals, weights, a, b);
55 double ig_scale = PosteriorScale(residuals, weights, a, b);
56 return ig_sampler_.Sample(ig_shape, ig_scale, gen);
57 }
58
59 private:
60 InverseGammaSampler ig_sampler_;
61};
62
65 public:
68 double PosteriorShape(TreeEnsemble* ensemble, double a, double b) {
69 data_size_t num_leaves = ensemble->NumLeaves();
70 return (a / 2.0) + (num_leaves / 2.0);
71 }
72 double PosteriorScale(TreeEnsemble* ensemble, double a, double b) {
73 double mu_sq = ensemble->SumLeafSquared();
74 return (b / 2.0) + (mu_sq / 2.0);
75 }
76 double SampleVarianceParameter(TreeEnsemble* ensemble, double a, double b, std::mt19937& gen) {
77 double ig_shape = PosteriorShape(ensemble, a, b);
78 double ig_scale = PosteriorScale(ensemble, a, b);
79 return ig_sampler_.Sample(ig_shape, ig_scale, gen);
80 }
81
82 private:
83 InverseGammaSampler ig_sampler_;
84};
85
86} // namespace StochTree
87
88#endif // STOCHTREE_VARIANCE_MODEL_H_
Marginal likelihood and posterior computation for gaussian homoskedastic constant leaf outcome model.
Definition variance_model.h:20
Definition ig_sampler.h:10
Marginal likelihood and posterior computation for gaussian homoskedastic constant leaf outcome model.
Definition variance_model.h:64
Class storing a "forest," or an ensemble of decision trees.
Definition ensemble.h:31
static void GFRSampleOneIter(TreeEnsemble &active_forest, ForestTracker &tracker, ForestContainer &forests, LeafModel &leaf_model, ForestDataset &dataset, ColumnVector &residual, TreePrior &tree_prior, std::mt19937 &gen, std::vector< double > &variable_weights, std::vector< int > &sweep_update_indices, double global_variance, std::vector< FeatureType > &feature_types, int cutpoint_grid_size, bool keep_forest, bool pre_initialized, bool backfitting, int num_features_subsample, int num_threads, LeafSuffStatConstructorArgs &... leaf_suff_stat_args)
Definition tree_sampler.h:816
A collection of random number generation utilities.
Definition bart.h:15