(booster: xgboost.Booster, name: str)
| 27 | |
| 28 | |
| 29 | def run_booster_check(booster: xgboost.Booster, name: str) -> None: |
| 30 | config = json.loads(booster.save_config()) |
| 31 | run_model_param_check(name, config) |
| 32 | n_rounds = get_n_rounds(name) |
| 33 | if name.find("cls") != -1: |
| 34 | assert len(booster.get_dump()) == gm.kForests * n_rounds * gm.kClasses |
| 35 | base_score = get_basescore(config) |
| 36 | assert isinstance(base_score, list) |
| 37 | assert all(v == 0.5 for v in base_score) |
| 38 | assert config["learner"]["learner_train_param"]["objective"] == "multi:softmax" |
| 39 | elif name.find("logitraw") != -1: |
| 40 | assert len(booster.get_dump()) == gm.kForests * n_rounds |
| 41 | assert config["learner"]["learner_model_param"]["num_class"] == str(0) |
| 42 | assert ( |
| 43 | config["learner"]["learner_train_param"]["objective"] == "binary:logitraw" |
| 44 | ) |
| 45 | elif name.find("logit") != -1: |
| 46 | assert len(booster.get_dump()) == gm.kForests * n_rounds |
| 47 | assert config["learner"]["learner_model_param"]["num_class"] == str(0) |
| 48 | assert ( |
| 49 | config["learner"]["learner_train_param"]["objective"] == "binary:logistic" |
| 50 | ) |
| 51 | elif name.find("ltr") != -1: |
| 52 | assert config["learner"]["learner_train_param"]["objective"] == "rank:ndcg" |
| 53 | elif name.find("aft") != -1: |
| 54 | assert config["learner"]["learner_train_param"]["objective"] == "survival:aft" |
| 55 | assert ( |
| 56 | config["learner"]["objective"]["aft_loss_param"]["aft_loss_distribution"] |
| 57 | == "normal" |
| 58 | ) |
| 59 | else: |
| 60 | assert name.find("reg") != -1 |
| 61 | assert len(booster.get_dump()) == gm.kForests * n_rounds |
| 62 | assert get_basescore(config) == [0.5] |
| 63 | assert ( |
| 64 | config["learner"]["learner_train_param"]["objective"] == "reg:squarederror" |
| 65 | ) |
| 66 | |
| 67 | |
| 68 | def get_n_rounds(name: str) -> int: |
no test coverage detected