6#ifndef STOCHTREE_TREE_H_
7#define STOCHTREE_TREE_H_
9#include <nlohmann/json.hpp>
10#include <stochtree/data.h>
11#include <stochtree/log.h>
12#include <stochtree/meta.h>
19using json = nlohmann::json;
26 kNumericalSplitNode = 1,
27 kCategoricalSplitNode = 2
46enum FeatureSplitType {
48 kOrderedCategoricalSplit,
49 kUnorderedCategoricalSplit
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};
75 Tree& operator=(
Tree const&) =
delete;
77 Tree& operator=(
Tree&&)
noexcept =
default;
86 std::int32_t num_nodes{0};
87 std::int32_t num_deleted_nodes{0};
111 inline bool IsRoot() {
return leaves_.size() == 1; }
122 void ChangeToLeaf(std::int32_t
nid,
double value) {
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());
151 if (this->
IsLeaf(nid))
return;
161 void ChangeToLeaf(std::int32_t
nid, std::vector<double>
value_vector) {
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());
191 if (this->
IsLeaf(nid))
return;
207 if (output_dimension_ == 1) {
208 for (
int j = 0;
j < leaf_value_.size();
j++) {
212 for (
int j = 0;
j < leaf_vector_.size();
j++) {
224 if (output_dimension_ == 1) {
225 for (
int j = 0;
j < leaf_value_.size();
j++) {
229 for (
int j = 0;
j < leaf_vector_.size();
j++) {
241 template <
typename Func>
243 std::stack<std::int32_t>
nodes;
246 while (!
nodes.empty()) {
254 if (
left != Tree::kInvalidNodeId) {
257 if (
right != Tree::kInvalidNodeId) {
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);
273 return output_dimension_ > 1;
280 return output_dimension_;
287 return is_log_scale_;
327 return split_index_[
nid];
335 return cleft_[
nid] == kInvalidNodeId;
343 return parent_[
nid] == kInvalidNodeId;
351 return node_deleted_[
nid];
359 return leaf_value_[
nid];
369 if (output_dimension_ == 1 &&
dim_id == 0) {
370 return leaf_value_[
nid];
375 Log::Fatal(
"No leaf vector set for node nid");
386 std::stack<std::int32_t>
nodes;
391 while (!
nodes.empty()) {
401 if (
left != Tree::kInvalidNodeId) {
405 if (
right != Tree::kInvalidNodeId) {
423 return std::vector<double>();
435 if (output_dimension_ == 1) {
436 return std::pow(leaf_value_[
nid], 2.0);
442 Log::Fatal(
"No leaf vector set for node nid");
445 result += std::pow(leaf_vector_[
i], 2.0);
456 for (
auto&
leaf : leaves_) {
467 return leaf_vector_begin_[
nid] != leaf_vector_end_[
nid];
475 return threshold_[
nid];
503 return node_type_[
nid];
511 return node_type_[
nid] == TreeNodeType::kNumericalSplitNode;
519 return node_type_[
nid] == TreeNodeType::kCategoricalSplitNode;
526 return has_categorical_split_;
530 [[
nodiscard]] std::int32_t NumLeaves()
const;
531 [[
nodiscard]] std::int32_t NumLeafParents()
const;
532 [[
nodiscard]] std::int32_t NumSplitNodes()
const;
535 [[
nodiscard]]
bool IsLeafParent(std::int32_t
nid)
const {
559 return internal_nodes_;
573 return leaf_parents_;
580 std::vector<std::int32_t>
output;
581 auto const&
self = *
this;
584 output.push_back(nidx);
618 return num_nodes - num_deleted_nodes;
772 Eigen::Map<Eigen::Matrix<int, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor>>&
output,
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_;
790 std::vector<double> leaf_vector_;
791 std::vector<std::uint64_t> leaf_vector_begin_;
792 std::vector<std::uint64_t> leaf_vector_end_;
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_;
799 bool has_categorical_split_{
false};
800 int output_dimension_{1};
801 bool is_log_scale_{
false};
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_));
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));
896 if (
tree.NodeType(
node_id) == StochTree::TreeNodeType::kCategoricalSplitNode) {
922 if (
tree.NodeType(
node_id) == StochTree::TreeNodeType::kCategoricalSplitNode) {
980 bool SplitSet() {
return split_set_; }
1000 bool split_set_{
false};
1002 double split_value_;
1003 std::vector<std::uint32_t> split_categories_;
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.