StochTree 0.5.0.9000
Loading...
Searching...
No Matches
ensemble.h
1
10#ifndef STOCHTREE_ENSEMBLE_H_
11#define STOCHTREE_ENSEMBLE_H_
12
13#include <stochtree/data.h>
14#include <stochtree/tree.h>
15#include <nlohmann/json.hpp>
16
17using json = nlohmann::json;
18
19namespace StochTree {
20
32 public:
41 TreeEnsemble(int num_trees, int output_dimension = 1, bool is_leaf_constant = true, bool is_exponentiated = false) {
42 // Initialize trees in the ensemble
43 trees_ = std::vector<std::unique_ptr<Tree>>(num_trees);
44 for (int i = 0; i < num_trees; i++) {
45 trees_[i].reset(new Tree());
46 trees_[i]->Init(output_dimension, is_exponentiated);
47 }
48 // Store ensemble configurations
49 num_trees_ = num_trees;
50 output_dimension_ = output_dimension;
51 is_leaf_constant_ = is_leaf_constant;
52 is_exponentiated_ = is_exponentiated;
53 }
54
61 // Unpack ensemble configurations
62 num_trees_ = ensemble.num_trees_;
63 output_dimension_ = ensemble.output_dimension_;
64 is_leaf_constant_ = ensemble.is_leaf_constant_;
65 is_exponentiated_ = ensemble.is_exponentiated_;
66 // Initialize trees in the ensemble
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());
70 }
71 // Clone trees in the ensemble
72 for (int j = 0; j < num_trees_; j++) {
73 Tree* tree = ensemble.GetTree(j);
74 this->CloneFromExistingTree(j, tree);
75 }
76 }
77
78 ~TreeEnsemble() {}
79
86 // Unpack ensemble configurations
87 int old_num_trees = num_trees_;
88 num_trees_ += ensemble.num_trees_;
89 CHECK_EQ(output_dimension_, ensemble.output_dimension_);
90 CHECK_EQ(is_leaf_constant_, ensemble.is_leaf_constant_);
91 CHECK_EQ(is_exponentiated_, ensemble.is_exponentiated_);
92 // Resize tree vector and reset new trees
93 trees_.resize(num_trees_);
94 for (int i = old_num_trees; i < num_trees_; i++) {
95 trees_[i].reset(new Tree());
96 }
97 // Clone trees in the input ensemble
98 for (int j = 0; j < ensemble.num_trees_; j++) {
99 Tree* tree = ensemble.GetTree(j);
100 this->CloneFromExistingTree(old_num_trees + j, tree);
101 }
102 }
103
110 for (int j = 0; j < num_trees_; j++) {
111 Tree* tree = GetTree(j);
112 tree->AddValueToLeaves(constant_value);
113 }
114 }
115
122 for (int j = 0; j < num_trees_; j++) {
123 Tree* tree = GetTree(j);
124 tree->MultiplyLeavesByValue(constant_multiple);
125 }
126 }
127
134 inline Tree* GetTree(int i) {
135 return trees_[i].get();
136 }
137
141 inline void ResetRoot() {
142 for (int i = 0; i < num_trees_; i++) {
144 }
145 }
146
153 inline void ResetTree(int i) {
154 trees_[i].reset(new Tree());
155 }
156
163 inline void ResetInitTree(int i) {
164 trees_[i].reset(new Tree());
165 trees_[i]->Init(output_dimension_, is_exponentiated_);
166 }
167
174 inline void CloneFromExistingTree(int i, Tree* tree) {
175 return trees_[i]->CloneFromTree(tree);
176 }
177
184 // Delete old tree pointers
185 trees_.clear();
186 // Unpack ensemble configurations
187 num_trees_ = ensemble.num_trees_;
188 output_dimension_ = ensemble.output_dimension_;
189 is_leaf_constant_ = ensemble.is_leaf_constant_;
190 is_exponentiated_ = ensemble.is_exponentiated_;
191 // Initialize trees in the ensemble
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());
195 }
196 // Clone trees in the ensemble
197 for (int j = 0; j < num_trees_; j++) {
198 Tree* tree = ensemble.GetTree(j);
199 this->CloneFromExistingTree(j, tree);
200 }
201 }
202
203 std::vector<double> Predict(ForestDataset& dataset) {
204 data_size_t n = dataset.NumObservations();
205 std::vector<double> output(n);
206 PredictInplace(dataset, output, 0);
207 return output;
208 }
209
210 std::vector<double> PredictRaw(ForestDataset& dataset, bool row_major = true) {
211 data_size_t n = dataset.NumObservations();
212 data_size_t total_output_size = n * output_dimension_;
213 std::vector<double> output(total_output_size);
214 PredictRawInplace(dataset, output, 0, trees_.size(), 0, row_major);
215 return output;
216 }
217
218 inline void PredictInplace(ForestDataset& dataset, std::vector<double>& output, data_size_t offset = 0) {
219 PredictInplace(dataset, output, 0, trees_.size(), offset);
220 }
221
222 inline void PredictInplace(ForestDataset& dataset, std::vector<double>& output,
223 int tree_begin, int tree_end, data_size_t offset = 0) {
224 if (is_leaf_constant_) {
225 PredictInplace(dataset.GetCovariates(), output, tree_begin, tree_end, offset);
226 } else {
227 CHECK(dataset.HasBasis());
228 PredictInplace(dataset.GetCovariates(), dataset.GetBasis(), output, tree_begin, tree_end, offset);
229 }
230 }
231
232 inline void PredictInplace(Eigen::MatrixXd& covariates, Eigen::MatrixXd& basis, std::vector<double>& output, data_size_t offset = 0) {
233 PredictInplace(covariates, basis, output, 0, trees_.size(), offset);
234 }
235
236 inline void PredictInplace(Eigen::MatrixXd& covariates, Eigen::MatrixXd& basis, std::vector<double>& output,
237 int tree_begin, int tree_end, data_size_t offset = 0) {
238 double pred;
239 CHECK_EQ(covariates.rows(), basis.rows());
240 CHECK_EQ(output_dimension_, trees_[0]->OutputDimension());
241 CHECK_EQ(output_dimension_, basis.cols());
242 data_size_t n = covariates.rows();
244 if (output.size() < total_output_size + offset) {
245 Log::Fatal("Mismatched size of prediction vector and training data");
246 }
247 for (data_size_t i = 0; i < n; i++) {
248 pred = 0.0;
249 for (size_t j = tree_begin; j < tree_end; j++) {
250 auto& tree = *trees_[j];
251 std::int32_t nidx = EvaluateTree(tree, covariates, i);
252 for (int32_t k = 0; k < output_dimension_; k++) {
253 pred += tree.LeafValue(nidx, k) * basis(i, k);
254 }
255 }
256 if (is_exponentiated_)
257 output[i + offset] = std::exp(pred);
258 else
259 output[i + offset] = pred;
260 }
261 }
262
263 inline void PredictInplace(Eigen::MatrixXd& covariates, std::vector<double>& output, data_size_t offset = 0) {
264 PredictInplace(covariates, output, 0, trees_.size(), offset);
265 }
266
267 inline void PredictInplace(Eigen::MatrixXd& covariates, std::vector<double>& output, int tree_begin, int tree_end, data_size_t offset = 0) {
268 double pred;
269 data_size_t n = covariates.rows();
271 if (output.size() < total_output_size + offset) {
272 Log::Fatal("Mismatched size of prediction vector and training data");
273 }
274 for (data_size_t i = 0; i < n; i++) {
275 pred = 0.0;
276 for (size_t j = tree_begin; j < tree_end; j++) {
277 auto& tree = *trees_[j];
278 std::int32_t nidx = EvaluateTree(tree, covariates, i);
279 pred += tree.LeafValue(nidx, 0);
280 }
281 if (is_exponentiated_)
282 output[i + offset] = std::exp(pred);
283 else
284 output[i + offset] = pred;
285 }
286 }
287
288 inline void PredictRawInplace(ForestDataset& dataset, std::vector<double>& output, data_size_t offset = 0, bool row_major = true) {
289 PredictRawInplace(dataset, output, 0, trees_.size(), offset, row_major);
290 }
291
292 inline void PredictRawInplace(ForestDataset& dataset, std::vector<double>& output,
293 int tree_begin, int tree_end, data_size_t offset = 0, bool row_major = true) {
294 double pred;
295 Eigen::MatrixXd covariates = dataset.GetCovariates();
296 CHECK_EQ(output_dimension_, trees_[0]->OutputDimension());
297 data_size_t n = covariates.rows();
298 data_size_t total_output_size = n * output_dimension_;
299 if (output.size() < total_output_size + offset) {
300 Log::Fatal("Mismatched size of raw prediction vector and training data");
301 }
302 for (data_size_t i = 0; i < n; i++) {
303 for (int32_t k = 0; k < output_dimension_; k++) {
304 pred = 0.0;
305 for (size_t j = tree_begin; j < tree_end; j++) {
306 auto& tree = *trees_[j];
308 pred += tree.LeafValue(nidx, k);
309 }
310 if (row_major) {
311 output[i * output_dimension_ + k + offset] = pred;
312 } else {
313 output[k * n + i + offset] = pred;
314 }
315 }
316 }
317 }
318
319 inline int32_t NumTrees() {
320 return num_trees_;
321 }
322
323 inline int32_t NumLeaves() {
324 int32_t result = 0;
325 for (int i = 0; i < num_trees_; i++) {
326 result += trees_[i]->NumLeaves();
327 }
328 return result;
329 }
330
331 inline double SumLeafSquared() {
332 double result = 0.;
333 for (int i = 0; i < num_trees_; i++) {
334 result += trees_[i]->SumSquaredLeafValues();
335 }
336 return result;
337 }
338
339 inline int32_t OutputDimension() {
340 return output_dimension_;
341 }
342
343 inline bool IsLeafConstant() {
344 return is_leaf_constant_;
345 }
346
347 inline bool IsExponentiated() {
348 return is_exponentiated_;
349 }
350
351 inline int32_t TreeMaxDepth(int tree_num) {
352 return trees_[tree_num]->MaxLeafDepth();
353 }
354
355 inline double AverageMaxDepth() {
356 double numerator = 0.;
357 double denominator = 0.;
358 for (int i = 0; i < num_trees_; i++) {
359 numerator += static_cast<double>(TreeMaxDepth(i));
360 denominator += 1.;
361 }
362 return numerator / denominator;
363 }
364
365 inline bool AllRoots() {
366 for (int i = 0; i < num_trees_; i++) {
367 if (!trees_[i]->IsRoot()) {
368 return false;
369 }
370 }
371 return true;
372 }
373
374 inline void SetLeafValue(double leaf_value) {
375 CHECK_EQ(output_dimension_, 1);
376 for (int i = 0; i < num_trees_; i++) {
377 CHECK(trees_[i]->IsRoot());
378 trees_[i]->SetLeaf(0, leaf_value);
379 }
380 }
381
382 inline void SetLeafVector(std::vector<double>& leaf_vector) {
383 CHECK_EQ(output_dimension_, leaf_vector.size());
384 for (int i = 0; i < num_trees_; i++) {
385 CHECK(trees_[i]->IsRoot());
386 trees_[i]->SetLeafVector(0, leaf_vector);
387 }
388 }
389
396 int max_leaf = 0;
397 for (int j = 0; j < num_trees_; j++) {
398 auto& tree = *trees_[j];
399 max_leaf += tree.NumLeaves();
400 }
401 return max_leaf;
402 }
403
422 PredictLeafIndicesInplace(dataset->GetCovariates(), output, num_trees, n);
423 }
424
442 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) {
443 CHECK_GE(output.size(), num_trees * n);
444 int offset = 0;
445 int max_leaf = 0;
446 for (int j = 0; j < num_trees; j++) {
447 auto& tree = *trees_[j];
448 int num_leaves = tree.NumLeaves();
449 tree.PredictLeafIndexInplace(covariates, output, offset, max_leaf);
450 offset += n;
452 }
453 }
454
473 void PredictLeafIndicesInplace(Eigen::Map<Eigen::Matrix<double, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor>>& covariates,
474 Eigen::Map<Eigen::Matrix<int, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor>>& output,
475 int column_ind, int num_trees, data_size_t n) {
476 CHECK_GE(output.size(), num_trees * n);
477 int offset = 0;
478 int max_leaf = 0;
479 for (int j = 0; j < num_trees; j++) {
480 auto& tree = *trees_[j];
481 int num_leaves = tree.NumLeaves();
482 tree.PredictLeafIndexInplace(covariates, output, column_ind, offset, max_leaf);
483 offset += n;
485 }
486 }
487
505 void PredictLeafIndicesInplace(Eigen::MatrixXd& covariates, std::vector<int32_t>& output, int num_trees, data_size_t n) {
506 CHECK_GE(output.size(), num_trees * n);
507 int offset = 0;
508 int max_leaf = 0;
509 for (int j = 0; j < num_trees; j++) {
510 auto& tree = *trees_[j];
511 int num_leaves = tree.NumLeaves();
512 tree.PredictLeafIndexInplace(covariates, output, offset, max_leaf);
513 offset += n;
515 }
516 }
517
522 std::vector<int32_t> PredictLeafIndices(ForestDataset* dataset) {
523 int num_trees = num_trees_;
524 data_size_t n = dataset->NumObservations();
525 std::vector<int32_t> output(n * num_trees);
527 return output;
528 }
529
531 json to_json() {
532 json result_obj;
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_);
537
538 std::string tree_label;
539 for (int i = 0; i < trees_.size(); i++) {
540 tree_label = "tree_" + std::to_string(i);
541 result_obj.emplace(tree_label, trees_[i]->to_json());
542 }
543
544 return result_obj;
545 }
546
548 void from_json(const json& ensemble_json) {
549 this->num_trees_ = ensemble_json.at("num_trees");
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");
553
554 std::string tree_label;
555 trees_.clear();
556 trees_.resize(this->num_trees_);
557 for (int i = 0; i < this->num_trees_; i++) {
558 tree_label = "tree_" + std::to_string(i);
559 trees_[i] = std::make_unique<Tree>();
560 trees_[i]->from_json(ensemble_json.at(tree_label));
561 }
562 }
563
564 private:
565 std::vector<std::unique_ptr<Tree>> trees_;
566 int num_trees_;
567 int output_dimension_;
568 bool is_leaf_constant_;
569 bool is_exponentiated_;
570};
571
// end of forest_group
573
574} // namespace StochTree
575
576#endif // STOCHTREE_ENSEMBLE_H_
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