63 std::string GetRngState()
const {
64 std::ostringstream
oss;
72 void SetRngState(
const std::string&
state) {
82 void RegenerateProbitLatent(
BCFSamples& samples);
87 bool initialized_ =
false;
100 void RestoreStateDefault();
106 void SampleParametricTreatmentEffect();
109 void SampleAdaptiveCodingParameters();
118 int warmstart_sample_num_ = 0;
122 std::variant<GaussianUnivariateRegressionLeafModel, GaussianMultivariateRegressionLeafModel> tau_leaf_model_;
126 std::unique_ptr<TreeEnsemble> mu_forest_;
127 std::unique_ptr<ForestTracker> mu_forest_tracker_;
128 std::unique_ptr<TreePrior> tree_prior_mu_;
129 std::unique_ptr<TreeEnsemble> tau_forest_;
130 std::unique_ptr<ForestTracker> tau_forest_tracker_;
131 std::unique_ptr<TreePrior> tree_prior_tau_;
133 double init_val_tau_;
134 std::vector<double> init_val_tau_vec_;
137 std::unique_ptr<TreeEnsemble> variance_forest_;
138 std::unique_ptr<ForestTracker> variance_forest_tracker_;
139 std::unique_ptr<TreePrior> tree_prior_variance_;
140 bool has_variance_forest_ =
false;
141 double init_val_variance_;
144 std::unique_ptr<MultivariateRegressionRandomEffectsModel> random_effects_model_;
145 std::unique_ptr<RandomEffectsTracker> random_effects_tracker_;
146 std::unique_ptr<RandomEffectsDataset> random_effects_dataset_;
147 bool has_random_effects_ =
false;
150 std::unique_ptr<ColumnVector> residual_;
151 std::unique_ptr<ColumnVector> outcome_raw_;
152 std::unique_ptr<ForestDataset> forest_dataset_;
153 std::unique_ptr<ForestDataset> forest_dataset_test_;
154 bool has_test_ =
false;
160 double global_variance_;
161 double leaf_scale_mu_;
162 double leaf_scale_tau_;
163 std::vector<double> leaf_scale_tau_multivariate_;
166 std::vector<double> model_preds_;
169 std::vector<double> tau_raw_sum_preds_;
172 std::unique_ptr<GlobalHomoskedasticVarianceModel> var_model_;
173 bool sample_sigma2_global_ =
false;
176 std::unique_ptr<LeafNodeHomoskedasticVarianceModel> leaf_scale_model_mu_;
177 bool sample_sigma2_leaf_mu_ =
false;
178 std::unique_ptr<LeafNodeHomoskedasticVarianceModel> leaf_scale_model_tau_;
179 bool sample_sigma2_leaf_tau_ =
false;
182 double tau_0_scalar_;
183 std::vector<double> tau_0_vector_;
184 bool sample_tau_0_ =
false;
189 bool adaptive_coding_ =
false;
190 std::vector<double> tau_basis_vector_train_;
191 std::vector<double> tau_basis_vector_test_;
194 struct GFROneIterationVisitorTau {
200 *sampler.tau_forest_, *sampler.tau_forest_tracker_, *samples.tau_forests,
model,
201 *sampler.forest_dataset_, *sampler.residual_, *sampler.tree_prior_tau_, sampler.rng_,
202 sampler.config_.var_weights_tau, sampler.config_.sweep_update_indices_tau, sampler.global_variance_, sampler.config_.feature_types,
203 sampler.config_.cutpoint_grid_size, keep_sample,
205 sampler.config_.num_features_subsample_tau, sampler.config_.num_threads);
209 *sampler.tau_forest_, *sampler.tau_forest_tracker_, *samples.tau_forests,
model,
210 *sampler.forest_dataset_, *sampler.residual_, *sampler.tree_prior_tau_, sampler.rng_,
211 sampler.config_.var_weights_tau, sampler.config_.sweep_update_indices_tau, sampler.global_variance_, sampler.config_.feature_types,
212 sampler.config_.cutpoint_grid_size, keep_sample,
214 sampler.config_.num_features_subsample_tau, sampler.config_.num_threads,
215 sampler.config_.leaf_dim_tau);
220 struct MCMCOneIterationVisitorTau {
226 *sampler.tau_forest_, *sampler.tau_forest_tracker_, *samples.tau_forests,
model,
227 *sampler.forest_dataset_, *sampler.residual_, *sampler.tree_prior_tau_, sampler.rng_,
228 sampler.config_.var_weights_tau, sampler.config_.sweep_update_indices_tau, sampler.global_variance_, keep_sample,
230 sampler.config_.num_threads);
234 *sampler.tau_forest_, *sampler.tau_forest_tracker_, *samples.tau_forests,
model,
235 *sampler.forest_dataset_, *sampler.residual_, *sampler.tree_prior_tau_, sampler.rng_,
236 sampler.config_.var_weights_tau, sampler.config_.sweep_update_indices_tau, sampler.global_variance_, keep_sample,
238 sampler.config_.num_threads, sampler.config_.leaf_dim_tau);
243 struct ScaleUpdateVisitor {
247 model.SetScale(leaf_scale);
250 model.SetScale(leaf_scale);
260 struct TauForestResetVisitor {
265 sampler.tau_forest_->ReconstituteFromForest(forest);
266 sampler.tau_forest_tracker_->ReconstituteFromForest(forest, *sampler.forest_dataset_, *sampler.residual_,
true);
267 sampler.tau_forest_tracker_->UpdatePredictions(sampler.tau_forest_.get(), *sampler.forest_dataset_.get());
270 sampler.tau_forest_->ReconstituteFromForest(forest);
271 sampler.tau_forest_tracker_->ReconstituteFromForest(forest, *sampler.forest_dataset_, *sampler.residual_,
true);
272 sampler.tau_forest_tracker_->UpdatePredictions(sampler.tau_forest_.get(), *sampler.forest_dataset_.get());
279 std::unique_ptr<TreeEnsemble> mu_forest;
280 std::unique_ptr<TreeEnsemble> tau_forest;
281 std::unique_ptr<TreeEnsemble> variance_forest;
285 double leaf_scale_mu;
286 double leaf_scale_tau;
287 std::vector<double> leaf_scale_tau_multivariate;
291 std::vector<double> tau_0_vector;
298 std::vector<double> residual;
301 std::vector<double> variance_weights;
304 Eigen::VectorXd rfx_working_parameter;
305 Eigen::MatrixXd rfx_group_parameters;
306 Eigen::MatrixXd rfx_group_parameter_covariance;
307 Eigen::MatrixXd rfx_working_parameter_covariance;
308 double rfx_variance_prior_shape;
309 double rfx_variance_prior_scale;
313 std::vector<GFRSnapshot> gfr_snapshots_;
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