45 BARTSamples* warmstart_source =
nullptr,
int warmstart_sample_num = 0);
49 void run_gfr(
BARTSamples& samples,
int num_gfr,
bool keep_gfr,
int num_chains = 0);
52 void run_mcmc(
BARTSamples& samples,
int num_burnin,
int keep_every,
int num_mcmc);
55 void run_mcmc_chains(
BARTSamples& samples,
int num_chains,
int num_burnin,
int keep_every,
int num_mcmc);
58 void postprocess_samples(
BARTSamples& samples,
int start_sample = 0);
63 std::string GetRngState()
const {
64 std::ostringstream oss;
72 void SetRngState(
const std::string& state) {
73 std::istringstream iss(state);
88 void InitializeState(
BARTSamples& samples,
bool continuation =
false);
89 bool initialized_ =
false;
92 void RestoreStateFromGFRSnapshot(
BARTSamples& samples,
int snapshot_index);
102 void RestoreStateDefault();
105 void RunOneIteration(
BARTSamples& samples,
bool gfr,
bool keep_sample,
bool write_snapshot =
false);
108 struct MeanForestInitVisitor {
112 sampler.mean_forest_ = std::make_unique<TreeEnsemble>(sampler.config_.num_trees_mean, sampler.config_.leaf_dim_mean, sampler.config_.leaf_constant_mean, sampler.config_.exponentiated_leaf_mean);
113 samples.mean_forests = std::make_unique<ForestContainer>(sampler.config_.num_trees_mean, sampler.config_.leaf_dim_mean, sampler.config_.leaf_constant_mean, sampler.config_.exponentiated_leaf_mean);
114 sampler.mean_forest_tracker_ = std::make_unique<ForestTracker>(sampler.forest_dataset_->GetCovariates(), sampler.config_.feature_types, sampler.config_.num_trees_mean, sampler.data_.n_train);
115 sampler.tree_prior_mean_ = std::make_unique<TreePrior>(sampler.config_.alpha_mean, sampler.config_.beta_mean, sampler.config_.min_samples_leaf_mean, sampler.config_.max_depth_mean);
116 sampler.mean_forest_->SetLeafValue(sampler.init_val_mean_ / sampler.config_.num_trees_mean);
117 UpdateResidualEntireForest(*sampler.mean_forest_tracker_, *sampler.forest_dataset_, *sampler.residual_, sampler.mean_forest_.get(), !sampler.config_.leaf_constant_mean, std::minus<double>());
118 sampler.mean_forest_tracker_->UpdatePredictions(sampler.mean_forest_.get(), *sampler.forest_dataset_.get());
119 sampler.has_mean_forest_ =
true;
122 sampler.mean_forest_ = std::make_unique<TreeEnsemble>(sampler.config_.num_trees_mean, sampler.config_.leaf_dim_mean, sampler.config_.leaf_constant_mean, sampler.config_.exponentiated_leaf_mean);
123 samples.mean_forests = std::make_unique<ForestContainer>(sampler.config_.num_trees_mean, sampler.config_.leaf_dim_mean, sampler.config_.leaf_constant_mean, sampler.config_.exponentiated_leaf_mean);
124 sampler.mean_forest_tracker_ = std::make_unique<ForestTracker>(sampler.forest_dataset_->GetCovariates(), sampler.config_.feature_types, sampler.config_.num_trees_mean, sampler.data_.n_train);
125 sampler.tree_prior_mean_ = std::make_unique<TreePrior>(sampler.config_.alpha_mean, sampler.config_.beta_mean, sampler.config_.min_samples_leaf_mean, sampler.config_.max_depth_mean);
126 sampler.mean_forest_->SetLeafValue(sampler.init_val_mean_ / sampler.config_.num_trees_mean);
127 UpdateResidualEntireForest(*sampler.mean_forest_tracker_, *sampler.forest_dataset_, *sampler.residual_, sampler.mean_forest_.get(), !sampler.config_.leaf_constant_mean, std::minus<double>());
128 sampler.mean_forest_tracker_->UpdatePredictions(sampler.mean_forest_.get(), *sampler.forest_dataset_.get());
129 sampler.has_mean_forest_ =
true;
132 sampler.mean_forest_ = std::make_unique<TreeEnsemble>(sampler.config_.num_trees_mean, sampler.config_.leaf_dim_mean, sampler.config_.leaf_constant_mean, sampler.config_.exponentiated_leaf_mean);
133 samples.mean_forests = std::make_unique<ForestContainer>(sampler.config_.num_trees_mean, sampler.config_.leaf_dim_mean, sampler.config_.leaf_constant_mean, sampler.config_.exponentiated_leaf_mean);
134 sampler.mean_forest_tracker_ = std::make_unique<ForestTracker>(sampler.forest_dataset_->GetCovariates(), sampler.config_.feature_types, sampler.config_.num_trees_mean, sampler.data_.n_train);
135 sampler.tree_prior_mean_ = std::make_unique<TreePrior>(sampler.config_.alpha_mean, sampler.config_.beta_mean, sampler.config_.min_samples_leaf_mean, sampler.config_.max_depth_mean);
136 sampler.mean_forest_->SetLeafVector(sampler.init_val_mean_vec_);
137 UpdateResidualEntireForest(*sampler.mean_forest_tracker_, *sampler.forest_dataset_, *sampler.residual_, sampler.mean_forest_.get(),
true, std::minus<double>());
138 sampler.mean_forest_tracker_->UpdatePredictions(sampler.mean_forest_.get(), *sampler.forest_dataset_.get());
139 sampler.has_mean_forest_ =
true;
142 sampler.mean_forest_ = std::make_unique<TreeEnsemble>(sampler.config_.num_trees_mean, sampler.config_.leaf_dim_mean, sampler.config_.leaf_constant_mean, sampler.config_.exponentiated_leaf_mean);
143 samples.mean_forests = std::make_unique<ForestContainer>(sampler.config_.num_trees_mean, sampler.config_.leaf_dim_mean, sampler.config_.leaf_constant_mean, sampler.config_.exponentiated_leaf_mean);
144 sampler.mean_forest_tracker_ = std::make_unique<ForestTracker>(sampler.forest_dataset_->GetCovariates(), sampler.config_.feature_types, sampler.config_.num_trees_mean, sampler.data_.n_train);
145 sampler.tree_prior_mean_ = std::make_unique<TreePrior>(sampler.config_.alpha_mean, sampler.config_.beta_mean, sampler.config_.min_samples_leaf_mean, sampler.config_.max_depth_mean);
146 sampler.mean_forest_->SetLeafValue(sampler.init_val_mean_ / sampler.config_.num_trees_mean);
147 UpdateResidualEntireForest(*sampler.mean_forest_tracker_, *sampler.forest_dataset_, *sampler.residual_, sampler.mean_forest_.get(),
false, std::minus<double>());
148 sampler.mean_forest_tracker_->UpdatePredictions(sampler.mean_forest_.get(), *sampler.forest_dataset_.get());
149 sampler.has_mean_forest_ =
true;
154 struct MeanForestResetVisitor {
159 sampler.mean_forest_->ReconstituteFromForest(forest);
160 sampler.mean_forest_tracker_->ReconstituteFromForest(forest, *sampler.forest_dataset_, *sampler.residual_,
true);
161 sampler.mean_forest_tracker_->UpdatePredictions(sampler.mean_forest_.get(), *sampler.forest_dataset_.get());
164 sampler.mean_forest_->ReconstituteFromForest(forest);
165 sampler.mean_forest_tracker_->ReconstituteFromForest(forest, *sampler.forest_dataset_, *sampler.residual_,
true);
166 sampler.mean_forest_tracker_->UpdatePredictions(sampler.mean_forest_.get(), *sampler.forest_dataset_.get());
169 sampler.mean_forest_->ReconstituteFromForest(forest);
170 sampler.mean_forest_tracker_->ReconstituteFromForest(forest, *sampler.forest_dataset_, *sampler.residual_,
true);
171 sampler.mean_forest_tracker_->UpdatePredictions(sampler.mean_forest_.get(), *sampler.forest_dataset_.get());
174 sampler.mean_forest_->ReconstituteFromForest(forest);
175 sampler.mean_forest_tracker_->ReconstituteFromForest(forest, *sampler.forest_dataset_, *sampler.residual_,
true);
176 sampler.mean_forest_tracker_->UpdatePredictions(sampler.mean_forest_.get(), *sampler.forest_dataset_.get());
185 struct MeanForestContinuationInitVisitor {
190 bool fresh_container;
193 sampler.mean_forest_ = std::make_unique<TreeEnsemble>(sampler.config_.num_trees_mean, sampler.config_.leaf_dim_mean, sampler.config_.leaf_constant_mean, sampler.config_.exponentiated_leaf_mean);
194 sampler.mean_forest_tracker_ = std::make_unique<ForestTracker>(sampler.forest_dataset_->GetCovariates(), sampler.config_.feature_types, sampler.config_.num_trees_mean, sampler.data_.n_train);
195 sampler.tree_prior_mean_ = std::make_unique<TreePrior>(sampler.config_.alpha_mean, sampler.config_.beta_mean, sampler.config_.min_samples_leaf_mean, sampler.config_.max_depth_mean);
196 if (fresh_container) {
197 samples.mean_forests = std::make_unique<ForestContainer>(sampler.config_.num_trees_mean, sampler.config_.leaf_dim_mean, sampler.config_.leaf_constant_mean, sampler.config_.exponentiated_leaf_mean);
199 sampler.has_mean_forest_ =
true;
201 TreeEnsemble& seed_forest = *source.mean_forests->GetEnsemble(idx);
202 sampler.mean_forest_->ReconstituteFromForest(seed_forest);
203 sampler.mean_forest_tracker_->ReconstituteFromForest(seed_forest, *sampler.forest_dataset_, *sampler.residual_,
true);
204 sampler.mean_forest_tracker_->UpdatePredictions(sampler.mean_forest_.get(), *sampler.forest_dataset_.get());
225 struct GFROneIterationVisitor {
231 *sampler.mean_forest_, *sampler.mean_forest_tracker_, *samples.mean_forests,
model,
232 *sampler.forest_dataset_, *sampler.residual_, *sampler.tree_prior_mean_, sampler.rng_,
233 sampler.config_.var_weights_mean, sampler.config_.sweep_update_indices_mean, sampler.global_variance_, sampler.config_.feature_types,
234 sampler.config_.cutpoint_grid_size, keep_sample,
236 sampler.config_.num_features_subsample_mean, sampler.config_.num_threads);
240 *sampler.mean_forest_, *sampler.mean_forest_tracker_, *samples.mean_forests,
model,
241 *sampler.forest_dataset_, *sampler.residual_, *sampler.tree_prior_mean_, sampler.rng_,
242 sampler.config_.var_weights_mean, sampler.config_.sweep_update_indices_mean, sampler.global_variance_, sampler.config_.feature_types,
243 sampler.config_.cutpoint_grid_size, keep_sample,
245 sampler.config_.num_features_subsample_mean, sampler.config_.num_threads);
249 *sampler.mean_forest_, *sampler.mean_forest_tracker_, *samples.mean_forests,
model,
250 *sampler.forest_dataset_, *sampler.residual_, *sampler.tree_prior_mean_, sampler.rng_,
251 sampler.config_.var_weights_mean, sampler.config_.sweep_update_indices_mean, sampler.global_variance_, sampler.config_.feature_types,
252 sampler.config_.cutpoint_grid_size, keep_sample,
254 sampler.config_.num_features_subsample_mean, sampler.config_.num_threads,
255 sampler.config_.leaf_dim_mean);
259 *sampler.mean_forest_, *sampler.mean_forest_tracker_, *samples.mean_forests,
model,
260 *sampler.forest_dataset_, *sampler.residual_, *sampler.tree_prior_mean_, sampler.rng_,
261 sampler.config_.var_weights_mean, sampler.config_.sweep_update_indices_mean, sampler.global_variance_, sampler.config_.feature_types,
262 sampler.config_.cutpoint_grid_size, keep_sample,
264 sampler.config_.num_features_subsample_mean, sampler.config_.num_threads);
269 struct MCMCOneIterationVisitor {
275 *sampler.mean_forest_, *sampler.mean_forest_tracker_, *samples.mean_forests,
model,
276 *sampler.forest_dataset_, *sampler.residual_, *sampler.tree_prior_mean_, sampler.rng_,
277 sampler.config_.var_weights_mean, sampler.config_.sweep_update_indices_mean, sampler.global_variance_, keep_sample,
279 sampler.config_.num_threads);
283 *sampler.mean_forest_, *sampler.mean_forest_tracker_, *samples.mean_forests,
model,
284 *sampler.forest_dataset_, *sampler.residual_, *sampler.tree_prior_mean_, sampler.rng_,
285 sampler.config_.var_weights_mean, sampler.config_.sweep_update_indices_mean, sampler.global_variance_, keep_sample,
287 sampler.config_.num_threads);
291 *sampler.mean_forest_, *sampler.mean_forest_tracker_, *samples.mean_forests,
model,
292 *sampler.forest_dataset_, *sampler.residual_, *sampler.tree_prior_mean_, sampler.rng_,
293 sampler.config_.var_weights_mean, sampler.config_.sweep_update_indices_mean, sampler.global_variance_, keep_sample,
295 sampler.config_.num_threads,
296 sampler.config_.leaf_dim_mean);
300 *sampler.mean_forest_, *sampler.mean_forest_tracker_, *samples.mean_forests,
model,
301 *sampler.forest_dataset_, *sampler.residual_, *sampler.tree_prior_mean_, sampler.rng_,
302 sampler.config_.var_weights_mean, sampler.config_.sweep_update_indices_mean, sampler.global_variance_, keep_sample,
304 sampler.config_.num_threads);
309 struct ScaleUpdateVisitor {
313 model.SetScale(leaf_scale);
316 model.SetScale(leaf_scale);
333 int warmstart_sample_num_ = 0;
336 std::variant<GaussianConstantLeafModel, GaussianUnivariateRegressionLeafModel, GaussianMultivariateRegressionLeafModel, CloglogOrdinalLeafModel> mean_leaf_model_;
340 std::unique_ptr<TreeEnsemble> mean_forest_;
341 std::unique_ptr<ForestTracker> mean_forest_tracker_;
342 std::unique_ptr<TreePrior> tree_prior_mean_;
343 bool has_mean_forest_ =
false;
344 double init_val_mean_;
345 std::vector<double> init_val_mean_vec_;
346 std::unique_ptr<OrdinalSampler> ordinal_sampler_;
349 std::unique_ptr<TreeEnsemble> variance_forest_;
350 std::unique_ptr<ForestTracker> variance_forest_tracker_;
351 std::unique_ptr<TreePrior> tree_prior_variance_;
352 bool has_variance_forest_ =
false;
353 double init_val_variance_;
356 std::unique_ptr<MultivariateRegressionRandomEffectsModel> random_effects_model_;
357 std::unique_ptr<RandomEffectsTracker> random_effects_tracker_;
358 std::unique_ptr<RandomEffectsDataset> random_effects_dataset_;
359 bool has_random_effects_ =
false;
362 std::unique_ptr<ColumnVector> residual_;
363 std::unique_ptr<ColumnVector> outcome_raw_;
364 std::unique_ptr<ForestDataset> forest_dataset_;
365 std::unique_ptr<ForestDataset> forest_dataset_test_;
366 bool has_test_ =
false;
372 double global_variance_;
374 std::vector<double> leaf_scale_multivariate_;
377 std::unique_ptr<GlobalHomoskedasticVarianceModel> var_model_;
378 bool sample_sigma2_global_ =
false;
381 std::unique_ptr<LeafNodeHomoskedasticVarianceModel> leaf_scale_model_;
382 bool sample_sigma2_leaf_ =
false;
387 std::unique_ptr<TreeEnsemble> mean_forest;
388 std::unique_ptr<TreeEnsemble> variance_forest;
393 std::vector<double> leaf_scale_multivariate;
396 std::vector<double> residual;
399 std::vector<double> variance_weights;
402 std::vector<double> cloglog_forest_preds;
403 std::vector<double> cloglog_latent_outcome;
404 std::vector<double> cloglog_logscale_cutpoints;
407 Eigen::VectorXd rfx_working_parameter;
408 Eigen::MatrixXd rfx_group_parameters;
409 Eigen::MatrixXd rfx_group_parameter_covariance;
410 Eigen::MatrixXd rfx_working_parameter_covariance;
411 double rfx_variance_prior_shape;
412 double rfx_variance_prior_scale;
416 std::vector<GFRSnapshot> gfr_snapshots_;