A tree built by XGBoost.
| 52 | |
| 53 | |
| 54 | class Tree: |
| 55 | """A tree built by XGBoost.""" |
| 56 | |
| 57 | def __init__(self, tree_id: int, nodes: Sequence[Node]) -> None: |
| 58 | self.tree_id = tree_id |
| 59 | self.nodes = nodes |
| 60 | |
| 61 | def loss_change(self, node_id: int) -> float: |
| 62 | """Loss gain of a node.""" |
| 63 | return self.nodes[node_id].loss_chg |
| 64 | |
| 65 | def sum_hessian(self, node_id: int) -> float: |
| 66 | """Sum Hessian of a node.""" |
| 67 | return self.nodes[node_id].sum_hess |
| 68 | |
| 69 | def base_weight(self, node_id: int) -> float: |
| 70 | """Base weight of a node.""" |
| 71 | return self.nodes[node_id].base_weight |
| 72 | |
| 73 | def split_index(self, node_id: int) -> int: |
| 74 | """Split feature index of node.""" |
| 75 | return self.nodes[node_id].split_idx |
| 76 | |
| 77 | def split_condition(self, node_id: int) -> float: |
| 78 | """Split value of a node.""" |
| 79 | return self.nodes[node_id].split_cond |
| 80 | |
| 81 | def split_categories(self, node_id: int) -> List[int]: |
| 82 | """Categories in a node.""" |
| 83 | return self.nodes[node_id].categories |
| 84 | |
| 85 | def is_categorical(self, node_id: int) -> bool: |
| 86 | """Whether a node has categorical split.""" |
| 87 | return self.nodes[node_id].split_type == SplitType.categorical |
| 88 | |
| 89 | def is_numerical(self, node_id: int) -> bool: |
| 90 | return not self.is_categorical(node_id) |
| 91 | |
| 92 | def parent(self, node_id: int) -> int: |
| 93 | """Parent ID of a node.""" |
| 94 | return self.nodes[node_id].parent |
| 95 | |
| 96 | def left_child(self, node_id: int) -> int: |
| 97 | """Left child ID of a node.""" |
| 98 | return self.nodes[node_id].left |
| 99 | |
| 100 | def right_child(self, node_id: int) -> int: |
| 101 | """Right child ID of a node.""" |
| 102 | return self.nodes[node_id].right |
| 103 | |
| 104 | def is_leaf(self, node_id: int) -> bool: |
| 105 | """Whether a node is leaf.""" |
| 106 | return self.nodes[node_id].left == -1 |
| 107 | |
| 108 | def is_deleted(self, node_id: int) -> bool: |
| 109 | """Whether a node is deleted.""" |
| 110 | return self.split_index(node_id) == np.iinfo(np.uint32).max |
| 111 |