10#ifndef STOCHTREE_ENSEMBLE_H_
11#define STOCHTREE_ENSEMBLE_H_
13#include <stochtree/data.h>
14#include <stochtree/tree.h>
15#include <nlohmann/json.hpp>
17using json = nlohmann::json;
43 trees_ = std::vector<std::unique_ptr<Tree>>(
num_trees);
45 trees_[
i].reset(
new Tree());
63 output_dimension_ =
ensemble.output_dimension_;
64 is_leaf_constant_ =
ensemble.is_leaf_constant_;
65 is_exponentiated_ =
ensemble.is_exponentiated_;
67 trees_ = std::vector<std::unique_ptr<Tree>>(num_trees_);
68 for (
int i = 0;
i < num_trees_;
i++) {
69 trees_[
i].reset(
new Tree());
72 for (
int j = 0;
j < num_trees_;
j++) {
93 trees_.resize(num_trees_);
95 trees_[
i].reset(
new Tree());
110 for (
int j = 0;
j < num_trees_;
j++) {
122 for (
int j = 0;
j < num_trees_;
j++) {
135 return trees_[
i].get();
142 for (
int i = 0;
i < num_trees_;
i++) {
154 trees_[
i].reset(
new Tree());
164 trees_[
i].reset(
new Tree());
165 trees_[
i]->Init(output_dimension_, is_exponentiated_);
175 return trees_[
i]->CloneFromTree(
tree);
188 output_dimension_ =
ensemble.output_dimension_;
189 is_leaf_constant_ =
ensemble.is_leaf_constant_;
190 is_exponentiated_ =
ensemble.is_exponentiated_;
192 trees_ = std::vector<std::unique_ptr<Tree>>(num_trees_);
193 for (
int i = 0;
i < num_trees_;
i++) {
194 trees_[
i].reset(
new Tree());
197 for (
int j = 0;
j < num_trees_;
j++) {
205 std::vector<double>
output(n);
210 std::vector<double> PredictRaw(ForestDataset&
dataset,
bool row_major =
true) {
222 inline void PredictInplace(ForestDataset&
dataset, std::vector<double>&
output,
224 if (is_leaf_constant_) {
236 inline void PredictInplace(Eigen::MatrixXd&
covariates, Eigen::MatrixXd&
basis, std::vector<double>&
output,
240 CHECK_EQ(output_dimension_, trees_[0]->OutputDimension());
245 Log::Fatal(
"Mismatched size of prediction vector and training data");
250 auto&
tree = *trees_[
j];
252 for (
int32_t k = 0;
k < output_dimension_;
k++) {
256 if (is_exponentiated_)
272 Log::Fatal(
"Mismatched size of prediction vector and training data");
277 auto&
tree = *trees_[
j];
281 if (is_exponentiated_)
292 inline void PredictRawInplace(ForestDataset&
dataset, std::vector<double>&
output,
296 CHECK_EQ(output_dimension_, trees_[0]->OutputDimension());
300 Log::Fatal(
"Mismatched size of raw prediction vector and training data");
303 for (
int32_t k = 0;
k < output_dimension_;
k++) {
306 auto&
tree = *trees_[
j];
325 for (
int i = 0;
i < num_trees_;
i++) {
326 result += trees_[
i]->NumLeaves();
331 inline double SumLeafSquared() {
333 for (
int i = 0;
i < num_trees_;
i++) {
334 result += trees_[
i]->SumSquaredLeafValues();
339 inline int32_t OutputDimension() {
340 return output_dimension_;
343 inline bool IsLeafConstant() {
344 return is_leaf_constant_;
347 inline bool IsExponentiated() {
348 return is_exponentiated_;
352 return trees_[
tree_num]->MaxLeafDepth();
355 inline double AverageMaxDepth() {
358 for (
int i = 0;
i < num_trees_;
i++) {
359 numerator +=
static_cast<double>(TreeMaxDepth(
i));
365 inline bool AllRoots() {
366 for (
int i = 0;
i < num_trees_;
i++) {
367 if (!trees_[
i]->IsRoot()) {
376 for (
int i = 0;
i < num_trees_;
i++) {
377 CHECK(trees_[
i]->IsRoot());
382 inline void SetLeafVector(std::vector<double>&
leaf_vector) {
384 for (
int i = 0;
i < num_trees_;
i++) {
385 CHECK(trees_[
i]->IsRoot());
397 for (
int j = 0;
j < num_trees_;
j++) {
398 auto&
tree = *trees_[
j];
447 auto&
tree = *trees_[
j];
474 Eigen::Map<Eigen::Matrix<int, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor>>&
output,
480 auto&
tree = *trees_[
j];
510 auto&
tree = *trees_[
j];
533 result_obj.emplace(
"num_trees", this->num_trees_);
534 result_obj.emplace(
"output_dimension", this->output_dimension_);
535 result_obj.emplace(
"is_leaf_constant", this->is_leaf_constant_);
536 result_obj.emplace(
"is_exponentiated", this->is_exponentiated_);
539 for (
int i = 0;
i < trees_.size();
i++) {
550 this->output_dimension_ =
ensemble_json.at(
"output_dimension");
551 this->is_leaf_constant_ =
ensemble_json.at(
"is_leaf_constant");
552 this->is_exponentiated_ =
ensemble_json.at(
"is_exponentiated");
556 trees_.resize(this->num_trees_);
557 for (
int i = 0;
i < this->num_trees_;
i++) {
559 trees_[
i] = std::make_unique<Tree>();
565 std::vector<std::unique_ptr<Tree>> trees_;
567 int output_dimension_;
568 bool is_leaf_constant_;
569 bool is_exponentiated_;
API for loading and accessing data used to sample tree ensembles The covariates / bases / weights use...
Definition data.h:272
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
int GetMaxLeafIndex()
Obtain a 0-based "maximum" leaf index for an ensemble, which is equivalent to the sum of the number o...
Definition ensemble.h:395
TreeEnsemble(TreeEnsemble &ensemble)
Initialize an ensemble based on the state of an existing ensemble.
Definition ensemble.h:60
void ResetTree(int i)
Reset a single tree in an ensemble.
Definition ensemble.h:153
void ReconstituteFromForest(TreeEnsemble &ensemble)
Reset an ensemble to clone another ensemble.
Definition ensemble.h:183
void AddValueToLeaves(double constant_value)
Add a constant value to every leaf of every tree in an ensemble. If leaves are multi-dimensional,...
Definition ensemble.h:109
json to_json()
Save to JSON.
Definition ensemble.h:531
void MergeForest(TreeEnsemble &ensemble)
Combine two forests into a single forest by merging their trees.
Definition ensemble.h:85
void PredictLeafIndicesInplace(ForestDataset *dataset, std::vector< int32_t > &output, int num_trees, data_size_t n)
Obtain a 0-based leaf index for every tree in an ensemble and for each observation in a ForestDataset...
Definition ensemble.h:421
void PredictLeafIndicesInplace(Eigen::MatrixXd &covariates, std::vector< int32_t > &output, int num_trees, data_size_t n)
Obtain a 0-based leaf index for every tree in an ensemble and for each observation in a ForestDataset...
Definition ensemble.h:505
void ResetRoot()
Reset a TreeEnsemble to all single-node "root" trees.
Definition ensemble.h:141
void from_json(const json &ensemble_json)
Load from JSON.
Definition ensemble.h:548
std::vector< int32_t > PredictLeafIndices(ForestDataset *dataset)
Same as PredictLeafIndicesInplace but assumes responsibility for allocating and returning output vect...
Definition ensemble.h:522
void PredictLeafIndicesInplace(Eigen::Map< Eigen::Matrix< double, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor > > &covariates, std::vector< int32_t > &output, int num_trees, data_size_t n)
Obtain a 0-based leaf index for every tree in an ensemble and for each observation in a ForestDataset...
Definition ensemble.h:442
void MultiplyLeavesByValue(double constant_multiple)
Multiply every leaf of every tree by a constant value. If leaves are multi-dimensional,...
Definition ensemble.h:121
TreeEnsemble(int num_trees, int output_dimension=1, bool is_leaf_constant=true, bool is_exponentiated=false)
Initialize a new TreeEnsemble.
Definition ensemble.h:41
void CloneFromExistingTree(int i, Tree *tree)
Clone a single tree in an ensemble from an existing tree, overwriting current tree.
Definition ensemble.h:174
void PredictLeafIndicesInplace(Eigen::Map< Eigen::Matrix< double, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor > > &covariates, Eigen::Map< Eigen::Matrix< int, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor > > &output, int column_ind, int num_trees, data_size_t n)
Obtain a 0-based leaf index for every tree in an ensemble and for each observation in a ForestDataset...
Definition ensemble.h:473
Decision tree data structure.
Definition tree.h:66
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
int EvaluateTree(Tree const &tree, Eigen::MatrixXd &data, int row)
Definition tree.h:888
A collection of random number generation utilities.
Definition bart.h:15