2#ifndef STOCHTREE_TREE_SAMPLER_H_
3#define STOCHTREE_TREE_SAMPLER_H_
5#include <stochtree/container.h>
6#include <stochtree/cutpoint_candidates.h>
7#include <stochtree/data.h>
8#include <stochtree/discrete_sampler.h>
9#include <stochtree/distributions.h>
10#include <stochtree/ensemble.h>
11#include <stochtree/leaf_model.h>
12#include <stochtree/openmp_utils.h>
13#include <stochtree/partition_tracker.h>
14#include <stochtree/prior.h>
49 var_min = std::numeric_limits<double>::max();
50 var_max = std::numeric_limits<double>::min();
81 int p =
dataset.GetCovariates().cols();
90 for (
int j = 0;
j < p;
j++) {
124 int p =
dataset.GetCovariates().cols();
130 for (
int j = 0;
j < p;
j++) {
133 var_max = std::numeric_limits<double>::min();
134 var_min = std::numeric_limits<double>::max();
155 if (
tree->OutputDimension() > 1) {
172 if (
tree->OutputDimension() > 1) {
184static inline double ComputeMeanOutcome(ColumnVector& residual) {
185 int n = residual.NumRows();
189 y = residual.GetElement(
i);
192 return sum_y /
static_cast<double>(n);
195static inline double ComputeVarianceOutcome(ColumnVector& residual) {
196 int n = residual.NumRows();
201 y = residual.GetElement(
i);
205 return sum_y_sq /
static_cast<double>(n) - (
sum_y *
sum_y) / (
static_cast<double>(n) *
static_cast<double>(n));
208static inline void UpdateModelVarianceForest(ForestTracker&
tracker, ForestDataset&
dataset, ColumnVector& residual,
209 TreeEnsemble* forest,
bool requires_basis, std::function<
double(
double,
double)>
op) {
216 for (
int j = 0;
j < forest->NumTrees();
j++) {
217 Tree*
tree = forest->GetTree(
j);
235static inline void UpdateResidualNoTrackerUpdate(ForestTracker&
tracker, ForestDataset&
dataset, ColumnVector& residual, TreeEnsemble* forest,
243 for (
int j = 0;
j < forest->NumTrees();
j++) {
244 Tree*
tree = forest->GetTree(
j);
260static inline void UpdateResidualEntireForest(ForestTracker&
tracker, ForestDataset&
dataset, ColumnVector& residual, TreeEnsemble* forest,
268 for (
int j = 0;
j < forest->NumTrees();
j++) {
269 Tree*
tree = forest->GetTree(
j);
287static inline void UpdateResidualNewOutcome(ForestTracker&
tracker, ColumnVector& residual) {
301static inline void UpdateMeanModelTree(ForestTracker&
tracker, ForestDataset&
dataset, ColumnVector& residual, Tree*
tree,
int tree_num,
332static inline void UpdateResidualNewBasis(ForestTracker&
tracker, ForestDataset&
dataset, ColumnVector& residual, TreeEnsemble* forest) {
364static inline void UpdateVarModelTree(ForestTracker&
tracker, ForestDataset&
dataset, ColumnVector& residual, Tree*
tree,
400static inline void UpdateCLogLogModelTree(ForestTracker&
tracker, ForestDataset&
dataset, ColumnVector& residual, Tree*
tree,
int tree_num,
435static inline std::tuple<double, double, data_size_t, data_size_t> EvaluateProposedSplit(
460static inline std::tuple<double, double, data_size_t, data_size_t> EvaluateExistingSplit(
482template <
typename LeafModel>
485 if constexpr (std::is_same_v<LeafModel, CloglogOrdinalLeafModel>) {
495template <
typename LeafModel>
498 if constexpr (std::is_same_v<LeafModel, CloglogOrdinalLeafModel>) {
522 int p =
dataset.NumCovariates();
542 Eigen::VectorXd&
outcome = residual.GetData();
552 if (num_threads == -1) {
560 StochTree::ParallelFor(0,
covariates.cols(), num_threads, [&](
int j) {
561 if ((std::abs(variable_weights.at(j)) > kEpsilon) && (feature_subset[j])) {
563 cutpoint_grid_container.CalculateStrides(covariates, outcome, tracker.GetSortedNodeSampleTracker(), node_id, node_begin, node_end, j, feature_types);
566 LeafSuffStat left_suff_stat = LeafSuffStat(leaf_suff_stat_args...);
567 LeafSuffStat right_suff_stat = LeafSuffStat(leaf_suff_stat_args...);
570 int32_t num_feature_cutpoints = cutpoint_grid_container.NumCutpoints(j);
571 FeatureType feature_type = feature_types[j];
573 for (data_size_t cutpoint_idx = 0; cutpoint_idx < (num_feature_cutpoints - 1); cutpoint_idx++) {
574 data_size_t current_bin_begin = cutpoint_grid_container.BinStartIndex(cutpoint_idx, j);
575 data_size_t current_bin_size = cutpoint_grid_container.BinLength(cutpoint_idx, j);
576 data_size_t next_bin_begin = cutpoint_grid_container.BinStartIndex(cutpoint_idx + 1, j);
579 AccumulateCutpointBinSuffStat<LeafSuffStat>(left_suff_stat, tracker, cutpoint_grid_container, dataset, residual,
580 global_variance, tree_num, node_id, j, cutpoint_idx);
583 right_suff_stat.SubtractSuffStat(node_suff_stat, left_suff_stat);
587 double cutoff_value = cutpoint_idx;
590 bool valid_split = (left_suff_stat.SampleGreaterThanEqual(min_samples_in_leaf) &&
591 right_suff_stat.SampleGreaterThanEqual(min_samples_in_leaf));
593 feature_cutpoint_counts[j]++;
595 feature_cutpoint_values[j].push_back(cutoff_value);
597 double split_log_ml = leaf_model.SplitLogMarginalLikelihood(left_suff_stat, right_suff_stat, global_variance);
598 feature_log_cutpoint_evaluations[j].push_back(split_log_ml);
623 double largest_ml = -std::numeric_limits<double>::infinity();
624 for (
int j = 0;
j < p + 1;
j++) {
632 for (
int j = 0;
j < p + 1;
j++) {
643 for (
int j = 0;
j < p + 1;
j++) {
681 if (
feature_type == FeatureType::kUnorderedCategorical) {
686 }
else if (
feature_type == FeatureType::kOrderedCategorical) {
695 Log::Fatal(
"Invalid split type");
699 AddSplitToModel(
tracker,
dataset,
tree_prior,
tree_split,
gen,
tree,
tree_num,
node_id,
feature_split,
true, num_threads);
721template <
typename LeafModel,
typename LeafSuffStat,
typename... LeafSuffStatConstructorArgs>
722static inline void GFRSampleTreeOneIter(Tree* tree, ForestTracker& tracker, ForestContainer& forests, LeafModel& leaf_model, ForestDataset& dataset,
723 ColumnVector& residual, TreePrior& tree_prior, std::mt19937& gen, std::vector<double>& variable_weights,
724 int tree_num,
double global_variance, std::vector<FeatureType>& feature_types,
int cutpoint_grid_size,
725 int num_features_subsample,
int num_threads, LeafSuffStatConstructorArgs&... leaf_suff_stat_args) {
726 int root_id = Tree::kRoot;
728 data_size_t curr_node_begin;
729 data_size_t curr_node_end;
730 data_size_t n = dataset.GetCovariates().rows();
731 int p = dataset.GetCovariates().cols();
734 std::vector<bool> feature_subset(p,
true);
735 if (num_features_subsample < p) {
737 int number_nonzero_weights = 0;
738 for (
int j = 0; j < p; j++) {
739 if (std::abs(variable_weights.at(j)) > kEpsilon) {
740 number_nonzero_weights++;
743 if (number_nonzero_weights > num_features_subsample) {
745 std::vector<int> feature_indices(p);
746 std::iota(feature_indices.begin(), feature_indices.end(), 0);
747 std::vector<int> features_selected(num_features_subsample);
748 sample_without_replacement<int, double>(
749 features_selected.data(), variable_weights.data(), feature_indices.data(),
750 p, num_features_subsample, gen);
751 for (
int i = 0; i < p; i++) {
752 feature_subset.at(i) =
false;
754 for (
const auto& feat : features_selected) {
755 feature_subset.at(feat) =
true;
761 std::unordered_map<int, std::pair<data_size_t, data_size_t>> node_index_map;
762 node_index_map.insert({root_id, std::make_pair(0, n)});
763 std::pair<data_size_t, data_size_t> begin_end;
765 std::deque<node_t> split_queue;
766 split_queue.push_back(Tree::kRoot);
768 while (!split_queue.empty()) {
770 curr_node_id = split_queue.front();
771 split_queue.pop_front();
773 begin_end = node_index_map[curr_node_id];
774 curr_node_begin = begin_end.first;
775 curr_node_end = begin_end.second;
777 SampleSplitRule<LeafModel, LeafSuffStat, LeafSuffStatConstructorArgs...>(
778 tree, tracker, leaf_model, dataset, residual, tree_prior, gen, tree_num, global_variance, cutpoint_grid_size,
779 node_index_map, split_queue, curr_node_id, curr_node_begin, curr_node_end, variable_weights, feature_types,
780 feature_subset, num_threads, leaf_suff_stat_args...);
815template <
typename LeafModel,
typename LeafSuffStat,
typename... LeafSuffStatConstructorArgs>
817 ColumnVector& residual,
TreePrior& tree_prior, std::mt19937& gen, std::vector<double>& variable_weights,
818 std::vector<int>& sweep_update_indices,
double global_variance, std::vector<FeatureType>& feature_types,
int cutpoint_grid_size,
819 bool keep_forest,
bool pre_initialized,
bool backfitting,
int num_features_subsample,
int num_threads, LeafSuffStatConstructorArgs&... leaf_suff_stat_args) {
821 int num_trees = forests.NumTrees();
822 for (
const int& i : sweep_update_indices) {
828 AdjustStateBeforeTreeSampling<LeafModel>(tracker, leaf_model, dataset, residual, tree_prior, backfitting, tree, i);
832 tracker.ResetRoot(dataset.
GetCovariates(), feature_types, i);
833 tree = active_forest.
GetTree(i);
836 GFRSampleTreeOneIter<LeafModel, LeafSuffStat, LeafSuffStatConstructorArgs...>(
837 tree, tracker, forests, leaf_model, dataset, residual, tree_prior, gen,
838 variable_weights, i, global_variance, feature_types, cutpoint_grid_size,
839 num_features_subsample, num_threads, leaf_suff_stat_args...);
842 tree = active_forest.
GetTree(i);
843 leaf_model.SampleLeafParameters(dataset, tracker, residual, tree, i, global_variance, gen);
849 AdjustStateAfterTreeSampling<LeafModel>(tracker, leaf_model, dataset, residual, tree_prior, backfitting, tree, i);
857template <
typename LeafModel,
typename LeafSuffStat,
typename... LeafSuffStatConstructorArgs>
858static inline void MCMCGrowTreeOneIter(Tree* tree, ForestTracker& tracker, LeafModel& leaf_model, ForestDataset& dataset, ColumnVector& residual,
859 TreePrior& tree_prior, std::mt19937& gen,
int tree_num, std::vector<double>& variable_weights,
860 double global_variance,
double prob_grow_old,
int num_threads, LeafSuffStatConstructorArgs&... leaf_suff_stat_args) {
862 data_size_t n = dataset.GetCovariates().rows();
865 int num_leaves = tree->NumLeaves();
866 std::vector<int> leaves = tree->GetLeaves();
867 std::vector<double> leaf_weights(num_leaves);
868 std::fill(leaf_weights.begin(), leaf_weights.end(), 1.0 / num_leaves);
869 walker_vose leaf_dist(leaf_weights.begin(), leaf_weights.end());
870 int leaf_chosen = leaves[leaf_dist(gen)];
871 int leaf_depth = tree->GetDepth(leaf_chosen);
874 int32_t max_depth = tree_prior.GetMaxDepth();
878 if ((leaf_depth >= max_depth) && (max_depth != -1)) {
882 int p = dataset.GetCovariates().cols();
883 CHECK_EQ(variable_weights.size(), p);
884 walker_vose var_dist(variable_weights.begin(), variable_weights.end());
885 int var_chosen = var_dist(gen);
889 double var_min, var_max;
890 VarSplitRange(tracker, dataset, tree_num, leaf_chosen, var_chosen, var_min, var_max);
891 if (var_max <= var_min) {
899 TreeSplit split = TreeSplit(split_point_chosen);
902 std::tuple<double, double, int32_t, int32_t> split_eval = EvaluateProposedSplit<LeafModel, LeafSuffStat, LeafSuffStatConstructorArgs...>(
903 dataset, tracker, residual, leaf_model, split, tree_num, leaf_chosen, var_chosen, global_variance, num_threads, leaf_suff_stat_args...);
904 double split_log_marginal_likelihood = std::get<0>(split_eval);
905 double no_split_log_marginal_likelihood = std::get<1>(split_eval);
906 int32_t left_n = std::get<2>(split_eval);
907 int32_t right_n = std::get<3>(split_eval);
910 bool left_node_sample_cutoff = left_n >= tree_prior.GetMinSamplesLeaf();
911 bool right_node_sample_cutoff = right_n >= tree_prior.GetMinSamplesLeaf();
912 if ((left_node_sample_cutoff) && (right_node_sample_cutoff)) {
914 double pg = tree_prior.GetAlpha() * std::pow(1 + leaf_depth, -tree_prior.GetBeta());
915 double pgl = tree_prior.GetAlpha() * std::pow(1 + leaf_depth + 1, -tree_prior.GetBeta());
916 double pgr = tree_prior.GetAlpha() * std::pow(1 + leaf_depth + 1, -tree_prior.GetBeta());
922 bool min_samples_left_check = left_n >= 2 * tree_prior.GetMinSamplesLeaf();
923 bool min_samples_right_check = right_n >= 2 * tree_prior.GetMinSamplesLeaf();
924 double prob_prune_new;
925 if (non_constant && (min_samples_left_check || min_samples_right_check)) {
926 prob_prune_new = 0.5;
928 prob_prune_new = 1.0;
932 int num_leaf_parents = tree->NumLeafParents();
933 double p_leaf = 1 /
static_cast<double>(num_leaves);
934 double p_leaf_parent = 1 /
static_cast<double>(num_leaf_parents + 1);
937 double log_mh_ratio = (std::log(pg) + std::log(1 - pgl) + std::log(1 - pgr) - std::log(1 - pg) + std::log(prob_prune_new) +
938 std::log(p_leaf_parent) - std::log(prob_grow_old) - std::log(p_leaf) - no_split_log_marginal_likelihood + split_log_marginal_likelihood);
940 if (log_mh_ratio > 0) {
946 if (log_acceptance_prob <= log_mh_ratio) {
948 AddSplitToModel(tracker, dataset, tree_prior, split, gen, tree, tree_num, leaf_chosen, var_chosen,
false, num_threads);
959template <
typename LeafModel,
typename LeafSuffStat,
typename... LeafSuffStatConstructorArgs>
960static inline void MCMCPruneTreeOneIter(Tree* tree, ForestTracker& tracker, LeafModel& leaf_model, ForestDataset& dataset, ColumnVector& residual,
961 TreePrior& tree_prior, std::mt19937& gen,
int tree_num,
double global_variance,
int num_threads,
962 LeafSuffStatConstructorArgs&... leaf_suff_stat_args) {
964 int num_leaves = tree->NumLeaves();
965 int num_leaf_parents = tree->NumLeafParents();
966 std::vector<int> leaf_parents = tree->GetLeafParents();
967 std::vector<double> leaf_parent_weights(num_leaf_parents);
968 std::fill(leaf_parent_weights.begin(), leaf_parent_weights.end(), 1.0 / num_leaf_parents);
969 walker_vose leaf_parent_dist(leaf_parent_weights.begin(), leaf_parent_weights.end());
970 int leaf_parent_chosen = leaf_parents[leaf_parent_dist(gen)];
971 int leaf_parent_depth = tree->GetDepth(leaf_parent_chosen);
972 int left_node = tree->LeftChild(leaf_parent_chosen);
973 int right_node = tree->RightChild(leaf_parent_chosen);
974 int feature_split = tree->SplitIndex(leaf_parent_chosen);
977 std::tuple<double, double, int32_t, int32_t> split_eval = EvaluateExistingSplit<LeafModel, LeafSuffStat, LeafSuffStatConstructorArgs...>(
978 dataset, tracker, residual, leaf_model, global_variance, tree_num, leaf_parent_chosen, left_node, right_node, leaf_suff_stat_args...);
979 double split_log_marginal_likelihood = std::get<0>(split_eval);
980 double no_split_log_marginal_likelihood = std::get<1>(split_eval);
981 int32_t left_n = std::get<2>(split_eval);
982 int32_t right_n = std::get<3>(split_eval);
985 double pg = tree_prior.GetAlpha() * std::pow(1 + leaf_parent_depth, -tree_prior.GetBeta());
986 double pgl = tree_prior.GetAlpha() * std::pow(1 + leaf_parent_depth + 1, -tree_prior.GetBeta());
987 double pgr = tree_prior.GetAlpha() * std::pow(1 + leaf_parent_depth + 1, -tree_prior.GetBeta());
992 bool non_root_tree = tree->NumNodes() > 1;
993 double prob_grow_new;
1002 bool non_constant_left = NodeNonConstant(dataset, tracker, tree_num, left_node);
1003 bool non_constant_right = NodeNonConstant(dataset, tracker, tree_num, right_node);
1004 double prob_prune_old;
1005 if (non_constant_left && non_constant_right) {
1006 prob_prune_old = 0.5;
1008 prob_prune_old = 1.0;
1012 double p_leaf = 1 /
static_cast<double>(num_leaves - 1);
1013 double p_leaf_parent = 1 /
static_cast<double>(num_leaf_parents);
1016 double log_mh_ratio = (std::log(1 - pg) - std::log(pg) - std::log(1 - pgl) - std::log(1 - pgr) + std::log(prob_prune_old) +
1017 std::log(p_leaf) - std::log(prob_grow_new) - std::log(p_leaf_parent) + no_split_log_marginal_likelihood - split_log_marginal_likelihood);
1019 if (log_mh_ratio > 0) {
1026 if (log_acceptance_prob <= log_mh_ratio) {
1028 RemoveSplitFromModel(tracker, dataset, tree_prior, gen, tree, tree_num, leaf_parent_chosen, left_node, right_node,
false);
1034template <
typename LeafModel,
typename LeafSuffStat,
typename... LeafSuffStatConstructorArgs>
1035static inline void MCMCSampleTreeOneIter(Tree* tree, ForestTracker& tracker, ForestContainer& forests, LeafModel& leaf_model, ForestDataset& dataset,
1036 ColumnVector& residual, TreePrior& tree_prior, std::mt19937& gen, std::vector<double>& variable_weights,
1037 int tree_num,
double global_variance,
int num_threads, LeafSuffStatConstructorArgs&... leaf_suff_stat_args) {
1039 bool grow_possible =
false;
1040 std::vector<int> leaves = tree->GetLeaves();
1041 for (
auto& leaf : leaves) {
1042 if (tracker.UnsortedNodeSize(tree_num, leaf) > 2 * tree_prior.GetMinSamplesLeaf()) {
1043 grow_possible =
true;
1049 bool prune_possible =
false;
1050 if (tree->NumValidNodes() > 1) {
1051 prune_possible =
true;
1056 std::vector<double> step_probs(2);
1057 if (grow_possible && prune_possible) {
1058 step_probs = {0.5, 0.5};
1060 }
else if (!grow_possible && prune_possible) {
1061 step_probs = {0.0, 1.0};
1063 }
else if (grow_possible && !prune_possible) {
1064 step_probs = {1.0, 0.0};
1067 Log::Fatal(
"In this tree, neither grow nor prune is possible");
1069 walker_vose step_dist(step_probs.begin(), step_probs.end());
1072 data_size_t step_chosen = step_dist(gen);
1075 if (step_chosen == 0) {
1076 MCMCGrowTreeOneIter<LeafModel, LeafSuffStat, LeafSuffStatConstructorArgs...>(
1077 tree, tracker, leaf_model, dataset, residual, tree_prior, gen, tree_num, variable_weights, global_variance, prob_grow, num_threads, leaf_suff_stat_args...);
1079 MCMCPruneTreeOneIter<LeafModel, LeafSuffStat, LeafSuffStatConstructorArgs...>(
1080 tree, tracker, leaf_model, dataset, residual, tree_prior, gen, tree_num, global_variance, num_threads, leaf_suff_stat_args...);
1112template <
typename LeafModel,
typename LeafSuffStat,
typename... LeafSuffStatConstructorArgs>
1114 ColumnVector& residual,
TreePrior& tree_prior, std::mt19937& gen, std::vector<double>& variable_weights,
1115 std::vector<int>& sweep_update_indices,
double global_variance,
bool keep_forest,
bool pre_initialized,
bool backfitting,
int num_threads,
1116 LeafSuffStatConstructorArgs&... leaf_suff_stat_args) {
1118 int num_trees = forests.NumTrees();
1119 for (
const int& i : sweep_update_indices) {
1125 AdjustStateBeforeTreeSampling<LeafModel>(tracker, leaf_model, dataset, residual, tree_prior, backfitting, tree, i);
1128 tree = active_forest.
GetTree(i);
1129 MCMCSampleTreeOneIter<LeafModel, LeafSuffStat, LeafSuffStatConstructorArgs...>(
1130 tree, tracker, forests, leaf_model, dataset, residual, tree_prior, gen, variable_weights, i,
1131 global_variance, num_threads, leaf_suff_stat_args...);
1134 tree = active_forest.
GetTree(i);
1135 leaf_model.SampleLeafParameters(dataset, tracker, residual, tree, i, global_variance, gen);
1141 AdjustStateAfterTreeSampling<LeafModel>(tracker, leaf_model, dataset, residual, tree_prior, backfitting, tree, i);
Internal wrapper around Eigen::VectorXd interface for univariate floating point data....
Definition data.h:193
Container of TreeEnsemble forest objects. This is the primary (in-memory) storage interface for multi...
Definition container.h:24
void AddSample(TreeEnsemble &forest)
Add a new forest to the container by copying forest.
API for loading and accessing data used to sample tree ensembles The covariates / bases / weights use...
Definition data.h:272
Eigen::MatrixXd & GetCovariates()
Return a reference to the raw Eigen::MatrixXd storing the covariate data.
Definition data.h:384
"Superclass" wrapper around tracking data structures for forest sampling algorithms
Definition partition_tracker.h:46
Class storing a "forest," or an ensemble of decision trees.
Definition ensemble.h:31
Tree * GetTree(int i)
Return a pointer to a tree in the forest.
Definition ensemble.h:134
void ResetInitTree(int i)
Reset a single tree in an ensemble.
Definition ensemble.h:163
Representation of arbitrary tree split rules, including numeric split rules (X[,i] <= c) and categori...
Definition tree.h:956
Decision tree data structure.
Definition tree.h:66
static void MCMCSampleOneIter(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, bool keep_forest, bool pre_initialized, bool backfitting, int num_threads, LeafSuffStatConstructorArgs &... leaf_suff_stat_args)
Runs one iteration of the MCMC sampler for a tree ensemble model, which consists of two steps for eve...
Definition tree_sampler.h:1113
static void VarSplitRange(ForestTracker &tracker, ForestDataset &dataset, int tree_num, int leaf_split, int feature_split, double &var_min, double &var_max)
Computer the range of available split values for a continuous variable, given the current structure o...
Definition tree_sampler.h:48
static bool NodesNonConstantAfterSplit(ForestDataset &dataset, ForestTracker &tracker, TreeSplit &split, int tree_num, int leaf_split, int feature_split)
Determines whether a proposed split creates two leaf nodes with constant values for every feature (th...
Definition tree_sampler.h:80
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
double standard_uniform_draw_53bit(std::mt19937 &gen)
Definition distributions.h:37