Construct the Model from a JSON object. parameters ---------- model : A dictionary loaded by json representing a XGBoost boosted tree model.
(self, model: dict)
| 144 | """Gradient boosted tree model.""" |
| 145 | |
| 146 | def __init__(self, model: dict) -> None: |
| 147 | """Construct the Model from a JSON object. |
| 148 | |
| 149 | parameters |
| 150 | ---------- |
| 151 | model : A dictionary loaded by json representing a XGBoost boosted tree model. |
| 152 | """ |
| 153 | # Basic properties of a model |
| 154 | self.learner_model_shape: ParamT = model["learner"]["learner_model_param"] |
| 155 | self.num_output_group = int(self.learner_model_shape["num_class"]) |
| 156 | self.num_feature = int(self.learner_model_shape["num_feature"]) |
| 157 | self.base_score: List[float] = json.loads( |
| 158 | self.learner_model_shape["base_score"] |
| 159 | ) |
| 160 | # A field encoding which output group a tree belongs |
| 161 | self.tree_info = model["learner"]["gradient_booster"]["model"]["tree_info"] |
| 162 | |
| 163 | model_shape: ParamT = model["learner"]["gradient_booster"]["model"][ |
| 164 | "gbtree_model_param" |
| 165 | ] |
| 166 | |
| 167 | # JSON representation of trees |
| 168 | j_trees = model["learner"]["gradient_booster"]["model"]["trees"] |
| 169 | |
| 170 | # Load the trees |
| 171 | self.num_trees = int(model_shape["num_trees"]) |
| 172 | |
| 173 | trees: List[Tree] = [] |
| 174 | for i in range(self.num_trees): |
| 175 | tree: Dict[str, Any] = j_trees[i] |
| 176 | tree_id = int(tree["id"]) |
| 177 | assert tree_id == i, (tree_id, i) |
| 178 | # - properties |
| 179 | left_children: List[int] = tree["left_children"] |
| 180 | right_children: List[int] = tree["right_children"] |
| 181 | parents: List[int] = tree["parents"] |
| 182 | split_conditions: List[float] = tree["split_conditions"] |
| 183 | split_indices: List[int] = tree["split_indices"] |
| 184 | # when ubjson is used, this is a byte array with each element as uint8 |
| 185 | default_left = to_integers(tree["default_left"]) |
| 186 | |
| 187 | # - categorical features |
| 188 | # when ubjson is used, this is a byte array with each element as uint8 |
| 189 | split_types = to_integers(tree["split_type"]) |
| 190 | # categories for each node is stored in a CSR style storage with segment as |
| 191 | # the begin ptr and the `categories' as values. |
| 192 | cat_segments: List[int] = tree["categories_segments"] |
| 193 | cat_sizes: List[int] = tree["categories_sizes"] |
| 194 | # node index for categorical nodes |
| 195 | cat_nodes: List[int] = tree["categories_nodes"] |
| 196 | assert len(cat_segments) == len(cat_sizes) == len(cat_nodes) |
| 197 | cats = tree["categories"] |
| 198 | assert len(left_children) == len(split_types) |
| 199 | |
| 200 | # The storage for categories is only defined for categorical nodes to |
| 201 | # prevent unnecessary overhead for numerical splits, we track the |
| 202 | # categorical node that are processed using a counter. |
| 203 | cat_cnt = 0 |
nothing calls this directly
no test coverage detected