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

Function run_scikit_model_check

tests/python/test_model_compatibility.py:76–124  ·  view source on GitHub ↗
(name: str, path: str)

Source from the content-addressed store, hash-verified

74
75
76def run_scikit_model_check(name: str, path: str) -> None:
77 if name.find("reg") != -1:
78 reg = xgboost.XGBRegressor()
79 reg.load_model(path)
80 config = json.loads(reg.get_booster().save_config())
81 assert (
82 config["learner"]["learner_train_param"]["objective"] == "reg:squarederror"
83 )
84 assert len(reg.get_booster().get_dump()) == get_n_rounds(name) * gm.kForests
85 run_model_param_check(name, config)
86 elif name.find("cls") != -1:
87 cls = xgboost.XGBClassifier()
88 cls.load_model(path)
89 n_rounds = get_n_rounds(name)
90 assert (
91 len(cls.get_booster().get_dump()) == n_rounds * gm.kForests * gm.kClasses
92 ), path
93 config = json.loads(cls.get_booster().save_config())
94 assert (
95 config["learner"]["learner_train_param"]["objective"] == "multi:softprob"
96 ), path
97 run_model_param_check(name, config)
98 elif name.find("ltr") != -1:
99 ltr = xgboost.XGBRanker()
100 ltr.load_model(path)
101 assert len(ltr.get_booster().get_dump()) == get_n_rounds(name) * gm.kForests
102 config = json.loads(ltr.get_booster().save_config())
103 assert config["learner"]["learner_train_param"]["objective"] == "rank:ndcg"
104 run_model_param_check(name, config)
105 elif name.find("logitraw") != -1:
106 logit = xgboost.XGBClassifier()
107 logit.load_model(path)
108 assert len(logit.get_booster().get_dump()) == get_n_rounds(name) * gm.kForests
109 config = json.loads(logit.get_booster().save_config())
110 assert (
111 config["learner"]["learner_train_param"]["objective"] == "binary:logitraw"
112 )
113 run_model_param_check(name, config)
114 elif name.find("logit") != -1:
115 logit = xgboost.XGBClassifier()
116 logit.load_model(path)
117 assert len(logit.get_booster().get_dump()) == get_n_rounds(name) * gm.kForests
118 config = json.loads(logit.get_booster().save_config())
119 assert (
120 config["learner"]["learner_train_param"]["objective"] == "binary:logistic"
121 )
122 run_model_param_check(name, config)
123 else:
124 assert False
125
126
127def download(path: str) -> None:

Callers 1

test_model_compatibilityFunction · 0.85

Calls 6

get_n_roundsFunction · 0.85
run_model_param_checkFunction · 0.85
save_configMethod · 0.80
get_dumpMethod · 0.80
load_modelMethod · 0.45
get_boosterMethod · 0.45

Tested by

no test coverage detected