MCPcopy Create free account
hub / github.com/dmlc/xgboost / __init__

Method __init__

demo/guide-python/model_parser.py:146–254  ·  view source on GitHub ↗

Construct the Model from a JSON object. parameters ---------- model : A dictionary loaded by json representing a XGBoost boosted tree model.

(self, model: dict)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 4

to_integersFunction · 0.85
TreeClass · 0.85
NodeClass · 0.70
SplitTypeClass · 0.70

Tested by

no test coverage detected