StochTree 0.5.0.9000
Loading...
Searching...
No Matches
tree.h
1
6#ifndef STOCHTREE_TREE_H_
7#define STOCHTREE_TREE_H_
8
9#include <nlohmann/json.hpp>
10#include <stochtree/data.h>
11#include <stochtree/log.h>
12#include <stochtree/meta.h>
13#include <Eigen/Dense>
14
15#include <cstdint>
16#include <stack>
17#include <string>
18
19using json = nlohmann::json;
20
21namespace StochTree {
22
25 kLeafNode = 0,
26 kNumericalSplitNode = 1,
27 kCategoricalSplitNode = 2
28};
29
30// template<typename T>
31// int enum_to_int(T& input_enum) {
32// return static_cast<int>(input_enum);
33// }
34
35// template<typename T>
36// T json_to_enum(json& input_json) {
37// return static_cast<T>(input_json);
38// }
39
42
45
46enum FeatureSplitType {
47 kNumericSplit,
48 kOrderedCategoricalSplit,
49 kUnorderedCategoricalSplit
50};
51
53class TreeSplit;
54
66class Tree {
67 public:
68 static constexpr std::int32_t kInvalidNodeId{-1};
69 static constexpr std::int32_t kDeletedNodeMarker = std::numeric_limits<node_t>::max();
70 static constexpr std::int32_t kRoot{0};
71
72 Tree() = default;
73 // ~Tree() = default;
74 Tree(Tree const&) = delete;
75 Tree& operator=(Tree const&) = delete;
76 Tree(Tree&&) noexcept = default;
77 Tree& operator=(Tree&&) noexcept = default;
85
86 std::int32_t num_nodes{0};
87 std::int32_t num_deleted_nodes{0};
88
90 void Reset();
92 void Init(int output_dimension = 1, bool is_log_scale = false);
94 int AllocNode();
96 void DeleteNode(std::int32_t nid);
98 void ExpandNode(std::int32_t nid, int split_index, double split_value, double left_value, double right_value);
100 void ExpandNode(std::int32_t nid, int split_index, std::vector<std::uint32_t> const& categorical_indices, double left_value, double right_value);
102 void ExpandNode(std::int32_t nid, int split_index, double split_value, std::vector<double> left_value_vector, std::vector<double> right_value_vector);
104 void ExpandNode(std::int32_t nid, int split_index, std::vector<std::uint32_t> const& categorical_indices, std::vector<double> left_value_vector, std::vector<double> right_value_vector);
106 void ExpandNode(std::int32_t nid, int split_index, TreeSplit& split, double left_value, double right_value);
108 void ExpandNode(std::int32_t nid, int split_index, TreeSplit& split, std::vector<double> left_value_vector, std::vector<double> right_value_vector);
109
111 inline bool IsRoot() { return leaves_.size() == 1; }
112
114 json to_json();
120 void from_json(const json& tree_json);
121
122 void ChangeToLeaf(std::int32_t nid, double value) {
123 CHECK(this->IsLeaf(this->LeftChild(nid)));
124 CHECK(this->IsLeaf(this->RightChild(nid)));
125 this->DeleteNode(this->LeftChild(nid));
126 this->DeleteNode(this->RightChild(nid));
127 this->SetLeaf(nid, value);
128
129 // Add nid to leaves and remove from internal nodes and leaf parents (if it was there)
130 leaves_.push_back(nid);
131 leaf_parents_.erase(std::remove(leaf_parents_.begin(), leaf_parents_.end(), nid), leaf_parents_.end());
132 internal_nodes_.erase(std::remove(internal_nodes_.begin(), internal_nodes_.end(), nid), internal_nodes_.end());
133
134 // Check if the other child of nid's parent node is also a leaf, if so, add parent back to leaf parents
135 // TODO refactor and add this to the multivariate case as well
136 if (!IsRoot(nid)) {
137 int parent_id = Parent(nid);
139 leaf_parents_.push_back(parent_id);
140 }
141 }
142 }
143
149 void CollapseToLeaf(std::int32_t nid, double value) {
150 CHECK_EQ(output_dimension_, 1);
151 if (this->IsLeaf(nid)) return;
152 if (!this->IsLeaf(this->LeftChild(nid))) {
153 CollapseToLeaf(this->LeftChild(nid), value);
154 }
155 if (!this->IsLeaf(this->RightChild(nid))) {
156 CollapseToLeaf(this->RightChild(nid), value);
157 }
158 this->ChangeToLeaf(nid, value);
159 }
160
161 void ChangeToLeaf(std::int32_t nid, std::vector<double> value_vector) {
162 CHECK(this->IsLeaf(this->LeftChild(nid)));
163 CHECK(this->IsLeaf(this->RightChild(nid)));
164 this->DeleteNode(this->LeftChild(nid));
165 this->DeleteNode(this->RightChild(nid));
166 this->SetLeafVector(nid, value_vector);
167
168 // Add nid to leaves and remove from internal nodes and leaf parents (if it was there)
169 leaves_.push_back(nid);
170 leaf_parents_.erase(std::remove(leaf_parents_.begin(), leaf_parents_.end(), nid), leaf_parents_.end());
171 internal_nodes_.erase(std::remove(internal_nodes_.begin(), internal_nodes_.end(), nid), internal_nodes_.end());
172
173 // Check if the other child of nid's parent node is also a leaf, if so, add parent back to leaf parents
174 // TODO refactor and add this to the multivariate case as well
175 if (!IsRoot(nid)) {
176 int parent_id = Parent(nid);
178 leaf_parents_.push_back(parent_id);
179 }
180 }
181 }
182
188 void CollapseToLeaf(std::int32_t nid, std::vector<double> value_vector) {
189 CHECK_GT(output_dimension_, 1);
190 CHECK_EQ(output_dimension_, value_vector.size());
191 if (this->IsLeaf(nid)) return;
192 if (!this->IsLeaf(this->LeftChild(nid))) {
194 }
195 if (!this->IsLeaf(this->RightChild(nid))) {
197 }
198 this->ChangeToLeaf(nid, value_vector);
199 }
200
207 if (output_dimension_ == 1) {
208 for (int j = 0; j < leaf_value_.size(); j++) {
209 leaf_value_[j] += constant_value;
210 }
211 } else {
212 for (int j = 0; j < leaf_vector_.size(); j++) {
213 leaf_vector_[j] += constant_value;
214 }
215 }
216 }
217
224 if (output_dimension_ == 1) {
225 for (int j = 0; j < leaf_value_.size(); j++) {
226 leaf_value_[j] *= constant_multiple;
227 }
228 } else {
229 for (int j = 0; j < leaf_vector_.size(); j++) {
230 leaf_vector_[j] *= constant_multiple;
231 }
232 }
233 }
234
241 template <typename Func>
242 void WalkTree(Func func) const {
243 std::stack<std::int32_t> nodes;
244 nodes.push(kRoot);
245 auto& self = *this;
246 while (!nodes.empty()) {
247 auto nidx = nodes.top();
248 nodes.pop();
249 if (!func(nidx)) {
250 return;
251 }
252 auto left = self.LeftChild(nidx);
253 auto right = self.RightChild(nidx);
254 if (left != Tree::kInvalidNodeId) {
255 nodes.push(left);
256 }
257 if (right != Tree::kInvalidNodeId) {
258 nodes.push(right);
259 }
260 }
261 }
262
263 std::vector<double> PredictFromNodes(std::vector<std::int32_t> node_indices);
264 std::vector<double> PredictFromNodes(std::vector<std::int32_t> node_indices, Eigen::MatrixXd& basis);
265 double PredictFromNode(std::int32_t node_id);
266 double PredictFromNode(std::int32_t node_id, Eigen::MatrixXd& basis, int row_idx);
267
272 bool HasVectorOutput() const {
273 return output_dimension_ > 1;
274 }
275
279 std::int32_t OutputDimension() const {
280 return output_dimension_;
281 }
282
286 bool IsLogScale() const {
287 return is_log_scale_;
288 }
289
294 std::int32_t Parent(std::int32_t nid) const {
295 return parent_[nid];
296 }
297
302 std::int32_t LeftChild(std::int32_t nid) const {
303 return cleft_[nid];
304 }
305
310 std::int32_t RightChild(std::int32_t nid) const {
311 return cright_[nid];
312 }
313
318 std::int32_t DefaultChild(std::int32_t nid) const {
319 return cleft_[nid];
320 }
321
326 std::int32_t SplitIndex(std::int32_t nid) const {
327 return split_index_[nid];
328 }
329
334 bool IsLeaf(std::int32_t nid) const {
335 return cleft_[nid] == kInvalidNodeId;
336 }
337
342 bool IsRoot(std::int32_t nid) const {
343 return parent_[nid] == kInvalidNodeId;
344 }
345
350 bool IsDeleted(std::int32_t nid) const {
351 return node_deleted_[nid];
352 }
353
358 double LeafValue(std::int32_t nid) const {
359 return leaf_value_[nid];
360 }
361
367 double LeafValue(std::int32_t nid, std::int32_t dim_id) const {
368 CHECK_LT(dim_id, output_dimension_);
369 if (output_dimension_ == 1 && dim_id == 0) {
370 return leaf_value_[nid];
371 } else {
372 std::size_t const offset_begin = leaf_vector_begin_[nid];
373 std::size_t const offset_end = leaf_vector_end_[nid];
374 if (offset_begin >= leaf_vector_.size() || offset_end > leaf_vector_.size()) {
375 Log::Fatal("No leaf vector set for node nid");
376 }
377 return leaf_vector_[offset_begin + dim_id];
378 }
379 }
380
384 std::int32_t MaxLeafDepth() const {
385 std::int32_t max_depth = 0;
386 std::stack<std::int32_t> nodes;
387 std::stack<std::int32_t> node_depths;
388 nodes.push(kRoot);
389 node_depths.push(0);
390 auto& self = *this;
391 while (!nodes.empty()) {
392 auto nidx = nodes.top();
393 nodes.pop();
394 auto node_depth = node_depths.top();
395 node_depths.pop();
396 bool valid_node = !self.IsDeleted(nidx);
397 if (valid_node) {
399 auto left = self.LeftChild(nidx);
400 auto right = self.RightChild(nidx);
401 if (left != Tree::kInvalidNodeId) {
402 nodes.push(left);
403 node_depths.push(node_depth + 1);
404 }
405 if (right != Tree::kInvalidNodeId) {
406 nodes.push(right);
407 node_depths.push(node_depth + 1);
408 }
409 }
410 }
411 return max_depth;
412 }
413
418 std::vector<double> LeafVector(std::int32_t nid) const {
419 std::size_t const offset_begin = leaf_vector_begin_[nid];
420 std::size_t const offset_end = leaf_vector_end_[nid];
421 if (offset_begin >= leaf_vector_.size() || offset_end > leaf_vector_.size()) {
422 // Return empty vector, to indicate the lack of leaf vector
423 return std::vector<double>();
424 }
425 return std::vector<double>(&leaf_vector_[offset_begin], &leaf_vector_[offset_end]);
426 // Use unsafe access here, since we may need to take the address of one past the last
427 // element, to follow with the range semantic of std::vector<>.
428 }
429
434 double SumSquaredNodeValues(std::int32_t nid) const {
435 if (output_dimension_ == 1) {
436 return std::pow(leaf_value_[nid], 2.0);
437 } else {
438 double result = 0.;
439 std::size_t const offset_begin = leaf_vector_begin_[nid];
440 std::size_t const offset_end = leaf_vector_end_[nid];
441 if (offset_begin >= leaf_vector_.size() || offset_end > leaf_vector_.size()) {
442 Log::Fatal("No leaf vector set for node nid");
443 }
444 for (std::size_t i = offset_begin; i < offset_end; i++) {
445 result += std::pow(leaf_vector_[i], 2.0);
446 }
447 return result;
448 }
449 }
450
454 double SumSquaredLeafValues() const {
455 double result = 0.;
456 for (auto& leaf : leaves_) {
458 }
459 return result;
460 }
461
466 bool HasLeafVector(std::int32_t nid) const {
467 return leaf_vector_begin_[nid] != leaf_vector_end_[nid];
468 }
469
474 double Threshold(std::int32_t nid) const {
475 return threshold_[nid];
476 }
477
485 std::vector<std::uint32_t> CategoryList(std::int32_t nid) const {
486 std::size_t const offset_begin = category_list_begin_[nid];
487 std::size_t const offset_end = category_list_end_[nid];
488 if (offset_begin >= category_list_.size() || offset_end > category_list_.size()) {
489 // Return empty vector, to indicate the lack of any category list
490 // The node might be a numerical split
491 return {};
492 }
493 return std::vector<std::uint32_t>(&category_list_[offset_begin], &category_list_[offset_end]);
494 // Use unsafe access here, since we may need to take the address of one past the last
495 // element, to follow with the range semantic of std::vector<>.
496 }
497
502 TreeNodeType NodeType(std::int32_t nid) const {
503 return node_type_[nid];
504 }
505
510 bool IsNumericSplitNode(std::int32_t nid) const {
511 return node_type_[nid] == TreeNodeType::kNumericalSplitNode;
512 }
513
518 bool IsCategoricalSplitNode(std::int32_t nid) const {
519 return node_type_[nid] == TreeNodeType::kCategoricalSplitNode;
520 }
521
525 bool HasCategoricalSplit() const {
526 return has_categorical_split_;
527 }
528
529 /* \brief Count number of leaves in tree. */
530 [[nodiscard]] std::int32_t NumLeaves() const;
531 [[nodiscard]] std::int32_t NumLeafParents() const;
532 [[nodiscard]] std::int32_t NumSplitNodes() const;
533
534 /* \brief Determine whether nid is leaf parent */
535 [[nodiscard]] bool IsLeafParent(std::int32_t nid) const {
536 // False until we deduce left and right node are
537 // available and both are leaves
538 bool is_left_leaf = false;
539 bool is_right_leaf = false;
540 // Check if node nidx is a leaf, if so, return false
541 bool is_leaf = this->IsLeaf(nid);
542 if (is_leaf) {
543 return false;
544 } else {
545 // If nidx is not a leaf, it must have left and right nodes
546 // so we check if those are leaves
547 std::int32_t left_node = LeftChild(nid);
548 std::int32_t right_node = RightChild(nid);
551 }
552 return is_left_leaf && is_right_leaf;
553 }
554
558 [[nodiscard]] std::vector<std::int32_t> const& GetInternalNodes() const {
559 return internal_nodes_;
560 }
561
565 [[nodiscard]] std::vector<std::int32_t> const& GetLeaves() const {
566 return leaves_;
567 }
568
572 [[nodiscard]] std::vector<std::int32_t> const& GetLeafParents() const {
573 return leaf_parents_;
574 }
575
579 [[nodiscard]] std::vector<std::int32_t> GetNodes() {
580 std::vector<std::int32_t> output;
581 auto const& self = *this;
582 this->WalkTree([&output, &self](std::int32_t nidx) {
583 if (!self.IsDeleted(nidx)) {
584 output.push_back(nidx);
585 }
586 return true;
587 });
588 return output;
589 }
590
595 [[nodiscard]] std::int32_t GetDepth(std::int32_t nid) const {
596 int depth = 0;
597 while (!IsRoot(nid)) {
598 ++depth;
599 nid = Parent(nid);
600 }
601 return depth;
602 }
603
607 [[nodiscard]] std::int32_t NumNodes() const noexcept { return num_nodes; }
608
612 [[nodiscard]] std::int32_t NumDeletedNodes() const noexcept { return num_deleted_nodes; }
613
618 return num_nodes - num_deleted_nodes;
619 }
620
627 void SetLeftChild(std::int32_t nid, std::int32_t left_child) {
628 cleft_[nid] = left_child;
629 }
630
636 void SetRightChild(std::int32_t nid, std::int32_t right_child) {
637 cright_[nid] = right_child;
638 }
639
646 void SetChildren(std::int32_t nid, std::int32_t left_child, std::int32_t right_child) {
649 }
650
656 void SetParent(std::int32_t child_node, std::int32_t parent_node) {
657 parent_[child_node] = parent_node;
658 }
659
666 void SetParents(std::int32_t nid, std::int32_t left_child, std::int32_t right_child) {
669 }
670
678 std::int32_t nid, std::int32_t split_index, double threshold);
679
688 void SetCategoricalSplit(std::int32_t nid, std::int32_t split_index,
689 std::vector<std::uint32_t> const& category_list);
690
696 void SetLeaf(std::int32_t nid, double value);
697
703 void SetLeafVector(std::int32_t nid, std::vector<double> const& leaf_vector);
704
726
747 void PredictLeafIndexInplace(Eigen::MatrixXd& covariates, std::vector<int32_t>& output, int32_t offset, int32_t max_leaf);
748
769 void PredictLeafIndexInplace(Eigen::Map<Eigen::Matrix<double, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor>>& covariates, std::vector<int32_t>& output, int32_t offset, int32_t max_leaf);
770
771 void PredictLeafIndexInplace(Eigen::Map<Eigen::Matrix<double, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor>>& covariates,
772 Eigen::Map<Eigen::Matrix<int, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor>>& output,
774
775 // Node info
776 std::vector<TreeNodeType> node_type_;
777 std::vector<std::int32_t> parent_;
778 std::vector<std::int32_t> cleft_;
779 std::vector<std::int32_t> cright_;
780 std::vector<std::int32_t> split_index_;
781 std::vector<double> leaf_value_;
782 std::vector<double> threshold_;
783 std::vector<bool> node_deleted_;
784 std::vector<std::int32_t> internal_nodes_;
785 std::vector<std::int32_t> leaves_;
786 std::vector<std::int32_t> leaf_parents_;
787 std::vector<std::int32_t> deleted_nodes_;
788
789 // Leaf vector
790 std::vector<double> leaf_vector_;
791 std::vector<std::uint64_t> leaf_vector_begin_;
792 std::vector<std::uint64_t> leaf_vector_end_;
793
794 // Category list
795 std::vector<std::uint32_t> category_list_;
796 std::vector<std::uint64_t> category_list_begin_;
797 std::vector<std::uint64_t> category_list_end_;
798
799 bool has_categorical_split_{false};
800 int output_dimension_{1};
801 bool is_log_scale_{false};
802};
803
805inline bool operator==(const Tree& lhs, const Tree& rhs) {
806 return (
807 (lhs.has_categorical_split_ == rhs.has_categorical_split_) &&
808 (lhs.output_dimension_ == rhs.output_dimension_) &&
809 (lhs.is_log_scale_ == rhs.is_log_scale_) &&
810 (lhs.node_type_ == rhs.node_type_) &&
811 (lhs.parent_ == rhs.parent_) &&
812 (lhs.cleft_ == rhs.cleft_) &&
813 (lhs.cright_ == rhs.cright_) &&
814 (lhs.split_index_ == rhs.split_index_) &&
815 (lhs.leaf_value_ == rhs.leaf_value_) &&
816 (lhs.threshold_ == rhs.threshold_) &&
817 (lhs.internal_nodes_ == rhs.internal_nodes_) &&
818 (lhs.leaves_ == rhs.leaves_) &&
819 (lhs.leaf_parents_ == rhs.leaf_parents_) &&
820 (lhs.deleted_nodes_ == rhs.deleted_nodes_) &&
821 (lhs.leaf_vector_ == rhs.leaf_vector_) &&
822 (lhs.leaf_vector_begin_ == rhs.leaf_vector_begin_) &&
823 (lhs.leaf_vector_end_ == rhs.leaf_vector_end_) &&
824 (lhs.category_list_ == rhs.category_list_) &&
825 (lhs.category_list_begin_ == rhs.category_list_begin_) &&
826 (lhs.category_list_end_ == rhs.category_list_end_));
827}
828
835inline bool SplitTrueNumeric(double fvalue, double threshold) {
836 return (fvalue <= threshold);
837}
838
845inline bool SplitTrueCategorical(double fvalue, std::vector<std::uint32_t> const& category_list) {
846 bool category_matched;
847 // A valid (integer) category must satisfy two criteria:
848 // 1) it must be exactly representable as double
849 // 2) it must fit into uint32_t
850 auto max_representable_int = std::min(static_cast<double>(std::numeric_limits<std::uint32_t>::max()),
851 static_cast<double>(std::uint64_t(1) << std::numeric_limits<double>::digits));
852 if (fvalue < 0 || std::fabs(fvalue) > max_representable_int) {
853 category_matched = false;
854 } else {
855 auto const category_value = static_cast<std::uint32_t>(fvalue);
856 category_matched = (std::find(category_list.begin(), category_list.end(), category_value) != category_list.end());
857 }
858 return category_matched;
859}
860
867inline int NextNodeNumeric(double fvalue, double threshold, int left_child, int right_child) {
869}
870
877inline int NextNodeCategorical(double fvalue, std::vector<std::uint32_t> const& category_list, int left_child, int right_child) {
879}
880
888inline int EvaluateTree(Tree const& tree, Eigen::MatrixXd& data, int row) {
889 int node_id = 0;
890 while (!tree.IsLeaf(node_id)) {
891 auto const split_index = tree.SplitIndex(node_id);
892 double const fvalue = data(row, split_index);
893 if (std::isnan(fvalue)) {
894 node_id = tree.DefaultChild(node_id);
895 } else {
896 if (tree.NodeType(node_id) == StochTree::TreeNodeType::kCategoricalSplitNode) {
898 tree.LeftChild(node_id), tree.RightChild(node_id));
899 } else {
900 node_id = NextNodeNumeric(fvalue, tree.Threshold(node_id), tree.LeftChild(node_id), tree.RightChild(node_id));
901 }
902 }
903 }
904 return node_id;
905}
906
914inline int EvaluateTree(Tree const& tree, Eigen::Map<Eigen::Matrix<double, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor>>& data, int row) {
915 int node_id = 0;
916 while (!tree.IsLeaf(node_id)) {
917 auto const split_index = tree.SplitIndex(node_id);
918 double const fvalue = data(row, split_index);
919 if (std::isnan(fvalue)) {
920 node_id = tree.DefaultChild(node_id);
921 } else {
922 if (tree.NodeType(node_id) == StochTree::TreeNodeType::kCategoricalSplitNode) {
924 tree.LeftChild(node_id), tree.RightChild(node_id));
925 } else {
926 node_id = NextNodeNumeric(fvalue, tree.Threshold(node_id), tree.LeftChild(node_id), tree.RightChild(node_id));
927 }
928 }
929 }
930 return node_id;
931}
932
939inline bool RowSplitLeft(Eigen::MatrixXd& covariates, int row, int split_index, double split_value) {
940 double const fvalue = covariates(row, split_index);
942}
943
950inline bool RowSplitLeft(Eigen::MatrixXd& covariates, int row, int split_index, std::vector<std::uint32_t> const& category_list) {
951 double const fvalue = covariates(row, split_index);
953}
954
957 public:
958 TreeSplit() {}
965 numeric_ = true;
966 split_value_ = split_value;
967 split_set_ = true;
968 }
974 TreeSplit(std::vector<std::uint32_t>& split_categories) {
975 numeric_ = false;
976 split_categories_ = split_categories;
977 split_set_ = true;
978 }
979 ~TreeSplit() {}
980 bool SplitSet() { return split_set_; }
982 bool NumericSplit() { return numeric_; }
988 bool SplitTrue(double fvalue) {
989 if (numeric_)
990 return SplitTrueNumeric(fvalue, split_value_);
991 else
992 return SplitTrueCategorical(fvalue, split_categories_);
993 }
995 double SplitValue() { return split_value_; }
997 std::vector<std::uint32_t> SplitCategories() { return split_categories_; }
998
999 private:
1000 bool split_set_{false};
1001 bool numeric_;
1002 double split_value_;
1003 std::vector<std::uint32_t> split_categories_;
1004};
1005
// end of tree_group
1007
1008} // namespace StochTree
1009
1010#endif // STOCHTREE_TREE_H_
API for loading and accessing data used to sample tree ensembles The covariates / bases / weights use...
Definition data.h:272
Representation of arbitrary tree split rules, including numeric split rules (X[,i] <= c) and categori...
Definition tree.h:956
bool SplitTrue(double fvalue)
Whether a given covariate value is True or False on the rule defined by a TreeSplit object.
Definition tree.h:988
bool NumericSplit()
Whether or not a TreeSplit rule is numeric.
Definition tree.h:982
std::vector< std::uint32_t > SplitCategories()
Categories defining a TreeSplit object.
Definition tree.h:997
TreeSplit(std::vector< std::uint32_t > &split_categories)
Construct a categorical TreeSplit.
Definition tree.h:974
double SplitValue()
Numeric cutoff value defining a TreeSplit object.
Definition tree.h:995
TreeSplit(double split_value)
Construct a numeric TreeSplit.
Definition tree.h:964
Decision tree data structure.
Definition tree.h:66
void SetNumericSplit(std::int32_t nid, std::int32_t split_index, double threshold)
Create a numerical split.
void SetLeftChild(std::int32_t nid, std::int32_t left_child)
Identify left child node.
Definition tree.h:627
std::int32_t LeftChild(std::int32_t nid) const
Index of the node's left child.
Definition tree.h:302
void PredictLeafIndexInplace(Eigen::Map< Eigen::Matrix< double, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor > > &covariates, std::vector< int32_t > &output, int32_t offset, int32_t max_leaf)
Obtain a 0-based leaf index for each observation in a ForestDataset. Internally, trees are stored as ...
void CloneFromTree(Tree *tree)
Copy the structure and parameters of another tree. If the Tree object calling this method already has...
std::int32_t RightChild(std::int32_t nid) const
Index of the node's right child.
Definition tree.h:310
std::int32_t NumValidNodes() const noexcept
Get the total number of valid nodes in this tree.
Definition tree.h:617
void SetCategoricalSplit(std::int32_t nid, std::int32_t split_index, std::vector< std::uint32_t > const &category_list)
Create a categorical split.
std::int32_t SplitIndex(std::int32_t nid) const
Feature index defining the node's split rule.
Definition tree.h:326
void PredictLeafIndexInplace(Eigen::MatrixXd &covariates, std::vector< int32_t > &output, int32_t offset, int32_t max_leaf)
Obtain a 0-based leaf index for each observation in a ForestDataset. Internally, trees are stored as ...
bool IsLeaf(std::int32_t nid) const
Whether the node is a leaf node.
Definition tree.h:334
bool IsNumericSplitNode(std::int32_t nid) const
Whether the node is a numeric split node.
Definition tree.h:510
void WalkTree(Func func) const
Iterate through all nodes in this tree.
Definition tree.h:242
double SumSquaredLeafValues() const
Sum of squared values for all leaves in a tree.
Definition tree.h:454
std::int32_t DefaultChild(std::int32_t nid) const
Index of the node's "default" child (potentially used in the case of a missing feature at prediction ...
Definition tree.h:318
void AddValueToLeaves(double constant_value)
Add a constant value to every leaf of a tree. If leaves are multi-dimensional, constant_value will be...
Definition tree.h:206
int AllocNode()
Allocate a new node and return the node's ID.
void ExpandNode(std::int32_t nid, int split_index, double split_value, std::vector< double > left_value_vector, std::vector< double > right_value_vector)
Expand a node based on a numeric split rule.
std::vector< double > LeafVector(std::int32_t nid) const
Get vector-valued parameters of a node (typically leaf)
Definition tree.h:418
std::int32_t Parent(std::int32_t nid) const
Index of the node's parent.
Definition tree.h:294
void PredictLeafIndexInplace(ForestDataset *dataset, std::vector< int32_t > &output, int32_t offset, int32_t max_leaf)
Obtain a 0-based leaf index for each observation in a ForestDataset. Internally, trees are stored as ...
bool IsCategoricalSplitNode(std::int32_t nid) const
Whether the node is a numeric split node.
Definition tree.h:518
bool HasCategoricalSplit() const
Query whether this tree contains any categorical splits.
Definition tree.h:525
double LeafValue(std::int32_t nid, std::int32_t dim_id) const
Get parameter value of a node (typically though not necessarily a leaf node) at a given output dimens...
Definition tree.h:367
void from_json(const json &tree_json)
Load from JSON.
std::int32_t NumNodes() const noexcept
Get the total number of nodes including deleted ones in this tree.
Definition tree.h:607
TreeNodeType NodeType(std::int32_t nid) const
Get the type of a node (i.e. numeric split, categorical split, leaf)
Definition tree.h:502
double SumSquaredNodeValues(std::int32_t nid) const
Sum of squared parameter values for a given node (typically though not necessarily a leaf node)
Definition tree.h:434
void DeleteNode(std::int32_t nid)
Deletes node indexed by node ID.
std::int32_t MaxLeafDepth() const
Get maximum depth of all of the leaf nodes.
Definition tree.h:384
void ExpandNode(std::int32_t nid, int split_index, TreeSplit &split, double left_value, double right_value)
Expand a node based on a generic split rule.
void ExpandNode(std::int32_t nid, int split_index, TreeSplit &split, std::vector< double > left_value_vector, std::vector< double > right_value_vector)
Expand a node based on a generic split rule.
std::int32_t OutputDimension() const
Dimension of tree output.
Definition tree.h:279
bool HasVectorOutput() const
Whether or not a tree has vector output.
Definition tree.h:272
void MultiplyLeavesByValue(double constant_multiple)
Multiply every leaf of a tree by a constant value. If leaves are multi-dimensional,...
Definition tree.h:223
std::int32_t NumDeletedNodes() const noexcept
Get the total number of deleted nodes in this tree.
Definition tree.h:612
void SetChildren(std::int32_t nid, std::int32_t left_child, std::int32_t right_child)
Identify two child nodes of the node and the corresponding parent node of the child nodes.
Definition tree.h:646
std::vector< std::uint32_t > CategoryList(std::int32_t nid) const
Get list of all categories belonging to the left child node. Categories are integers ranging from 0 t...
Definition tree.h:485
void SetParent(std::int32_t child_node, std::int32_t parent_node)
Identify parent node.
Definition tree.h:656
void SetLeaf(std::int32_t nid, double value)
Set the leaf value of the node.
bool IsRoot()
Whether or not a tree is a "stump" consisting of a single root node.
Definition tree.h:111
double LeafValue(std::int32_t nid) const
Get parameter value of a node (typically though not necessarily a leaf node)
Definition tree.h:358
void SetParents(std::int32_t nid, std::int32_t left_child, std::int32_t right_child)
Identify parent node of the left and right node ids.
Definition tree.h:666
std::vector< std::int32_t > const & GetInternalNodes() const
Get indices of all internal nodes.
Definition tree.h:558
std::vector< std::int32_t > const & GetLeafParents() const
Get indices of all leaf parent nodes.
Definition tree.h:572
std::vector< std::int32_t > GetNodes()
Get indices of all valid (non-deleted) nodes.
Definition tree.h:579
bool HasLeafVector(std::int32_t nid) const
Tests whether the leaf node has a non-empty leaf vector.
Definition tree.h:466
void Init(int output_dimension=1, bool is_log_scale=false)
Initialize the tree with a single root node.
void CollapseToLeaf(std::int32_t nid, std::vector< double > value_vector)
Collapse an internal node to a leaf node, deleting its children from the tree.
Definition tree.h:188
double Threshold(std::int32_t nid) const
Get split threshold of the node.
Definition tree.h:474
void ExpandNode(std::int32_t nid, int split_index, std::vector< std::uint32_t > const &categorical_indices, double left_value, double right_value)
Expand a node based on a categorical split rule.
bool IsRoot(std::int32_t nid) const
Whether the node is root.
Definition tree.h:342
void Reset()
Reset tree to empty vectors and default values of boolean / integer variables.
bool IsDeleted(std::int32_t nid) const
Whether the node has been deleted.
Definition tree.h:350
void ExpandNode(std::int32_t nid, int split_index, double split_value, double left_value, double right_value)
Expand a node based on a numeric split rule.
void SetRightChild(std::int32_t nid, std::int32_t right_child)
Identify right child node.
Definition tree.h:636
std::vector< std::int32_t > const & GetLeaves() const
Get indices of all leaf nodes.
Definition tree.h:565
void ExpandNode(std::int32_t nid, int split_index, std::vector< std::uint32_t > const &categorical_indices, std::vector< double > left_value_vector, std::vector< double > right_value_vector)
Expand a node based on a categorical split rule.
bool IsLogScale() const
Whether or not tree parameters should be exponentiated at prediction time.
Definition tree.h:286
void CollapseToLeaf(std::int32_t nid, double value)
Collapse an internal node to a leaf node, deleting its children from the tree.
Definition tree.h:149
void SetLeafVector(std::int32_t nid, std::vector< double > const &leaf_vector)
Set the leaf vector of the node; useful for multi-output trees.
std::int32_t GetDepth(std::int32_t nid) const
Get the depth of a node.
Definition tree.h:595
json to_json()
Convert tree to JSON and return JSON in-memory.
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
bool SplitTrueNumeric(double fvalue, double threshold)
Determine whether an observation produces a "true" value in a numeric split node.
Definition tree.h:835
int NextNodeCategorical(double fvalue, std::vector< std::uint32_t > const &category_list, int left_child, int right_child)
Return left or right node id based on a categorical split.
Definition tree.h:877
bool SplitTrueCategorical(double fvalue, std::vector< std::uint32_t > const &category_list)
Determine whether an observation produces a "true" value in a categorical split node.
Definition tree.h:845
int EvaluateTree(Tree const &tree, Eigen::MatrixXd &data, int row)
Definition tree.h:888
int NextNodeNumeric(double fvalue, double threshold, int left_child, int right_child)
Return left or right node id based on a numeric split.
Definition tree.h:867
bool operator==(const Tree &lhs, const Tree &rhs)
Comparison operator for trees.
Definition tree.h:805
bool RowSplitLeft(Eigen::MatrixXd &covariates, int row, int split_index, double split_value)
Determine whether a given observation is "true" at a split proposed by split_index and split_value.
Definition tree.h:939
A collection of random number generation utilities.
Definition bart.h:15
std::string TreeNodeTypeToString(TreeNodeType type)
Get string representation of TreeNodeType.
TreeNodeType
Tree node type.
Definition tree.h:24
TreeNodeType TreeNodeTypeFromString(std::string const &name)
Get NodeType from string.