(model_path: str)
| 218 | |
| 219 | |
| 220 | def save_load_model(model_path: str) -> None: |
| 221 | from sklearn.datasets import load_digits |
| 222 | from sklearn.model_selection import KFold |
| 223 | |
| 224 | rng = np.random.RandomState(1994) |
| 225 | |
| 226 | digits = load_digits(n_class=2) |
| 227 | y = digits["target"] |
| 228 | X = digits["data"] |
| 229 | kf = KFold(n_splits=2, shuffle=True, random_state=rng) |
| 230 | for train_index, test_index in kf.split(X, y): |
| 231 | xgb_model = xgb.XGBClassifier().fit(X[train_index], y[train_index]) |
| 232 | xgb_model.save_model(model_path) |
| 233 | |
| 234 | xgb_model = xgb.XGBClassifier() |
| 235 | xgb_model.load_model(model_path) |
| 236 | |
| 237 | assert isinstance(xgb_model.classes_, np.ndarray) |
| 238 | np.testing.assert_equal(xgb_model.classes_, np.array([0, 1])) |
| 239 | assert isinstance(xgb_model._Booster, xgb.Booster) |
| 240 | |
| 241 | preds = xgb_model.predict(X[test_index]) |
| 242 | labels = y[test_index] |
| 243 | err = sum( |
| 244 | 1 for i in range(len(preds)) if int(preds[i] > 0.5) != labels[i] |
| 245 | ) / float(len(preds)) |
| 246 | assert err < 0.1 |
| 247 | assert xgb_model.get_booster().attr("scikit_learn") is None |
| 248 | |
| 249 | # test native booster |
| 250 | preds = xgb_model.predict(X[test_index], output_margin=True) |
| 251 | booster = xgb.Booster(model_file=model_path) |
| 252 | predt_1 = booster.predict(xgb.DMatrix(X[test_index]), output_margin=True) |
| 253 | assert np.allclose(preds, predt_1) |
| 254 | |
| 255 | with pytest.raises(TypeError): |
| 256 | xgb_model = xgb.XGBModel() |
| 257 | xgb_model.load_model(model_path) |
| 258 | |
| 259 | clf = xgb.XGBClassifier(booster="gblinear", early_stopping_rounds=1) |
| 260 | clf.fit(X, y, eval_set=[(X, y)]) |
| 261 | best_iteration = clf.best_iteration |
| 262 | best_score = clf.best_score |
| 263 | predt_0 = clf.predict(X) |
| 264 | clf.save_model(model_path) |
| 265 | clf.load_model(model_path) |
| 266 | assert clf.booster == "gblinear" |
| 267 | predt_1 = clf.predict(X) |
| 268 | np.testing.assert_allclose(predt_0, predt_1) |
| 269 | assert clf.best_iteration == best_iteration |
| 270 | assert clf.best_score == best_score |
| 271 | |
| 272 | clfpkl = pickle.dumps(clf) |
| 273 | clf = pickle.loads(clfpkl) |
| 274 | predt_2 = clf.predict(X) |
| 275 | np.testing.assert_allclose(predt_0, predt_2) |
| 276 | assert clf.best_iteration == best_iteration |
| 277 | assert clf.best_score == best_score |
no test coverage detected