| 567 | } |
| 568 | |
| 569 | void LoadConfig(Json const& in) override { |
| 570 | // If configuration is loaded, ensure that the model came from the same version |
| 571 | CHECK(IsA<Object>(in)); |
| 572 | auto origin_version = Version::Load(in); |
| 573 | if (std::get<0>(Version::kInvalid) == std::get<0>(origin_version)) { |
| 574 | LOG(WARNING) << "Invalid version string in config"; |
| 575 | } |
| 576 | |
| 577 | if (!Version::Same(origin_version)) { |
| 578 | error::WarnOldSerialization(); |
| 579 | return; // skip configuration if version is not matched |
| 580 | } |
| 581 | |
| 582 | auto const& learner_parameters = get<Object>(in["learner"]); |
| 583 | FromJson(learner_parameters.at("learner_train_param"), &tparam_); |
| 584 | |
| 585 | auto const& gradient_booster = learner_parameters.at("gradient_booster"); |
| 586 | |
| 587 | auto const& objective_fn = learner_parameters.at("objective"); |
| 588 | if (!obj_) { |
| 589 | CHECK_EQ(get<String const>(objective_fn["name"]), tparam_.objective); |
| 590 | obj_.reset(ObjFunction::Create(tparam_.objective, &ctx_)); |
| 591 | } |
| 592 | obj_->LoadConfig(objective_fn); |
| 593 | learner_model_param_.task = obj_->Task(); |
| 594 | |
| 595 | tparam_.booster = CanonicalizeBoosterName(get<String>(gradient_booster["name"])); |
| 596 | if (!gbm_) { |
| 597 | gbm_.reset(GradientBooster::Create(tparam_.booster, &ctx_, &learner_model_param_)); |
| 598 | } |
| 599 | gbm_->LoadConfig(gradient_booster); |
| 600 | |
| 601 | auto const& j_metrics = learner_parameters.at("metrics"); |
| 602 | auto n_metrics = get<Array const>(j_metrics).size(); |
| 603 | metric_names_.resize(n_metrics); |
| 604 | metrics_.resize(n_metrics); |
| 605 | for (size_t i = 0; i < n_metrics; ++i) { |
| 606 | auto old_serialization = IsA<String>(j_metrics[i]); |
| 607 | if (old_serialization) { |
| 608 | error::WarnOldSerialization(); |
| 609 | metric_names_[i] = get<String>(j_metrics[i]); |
| 610 | } else { |
| 611 | metric_names_[i] = get<String>(j_metrics[i]["name"]); |
| 612 | } |
| 613 | metrics_[i] = std::unique_ptr<Metric>(Metric::Create(metric_names_[i], &ctx_)); |
| 614 | if (!old_serialization) { |
| 615 | metrics_[i]->LoadConfig(j_metrics[i]); |
| 616 | } |
| 617 | } |
| 618 | |
| 619 | ctx_.FromJson(learner_parameters.at("generic_param")); |
| 620 | |
| 621 | this->need_configuration_ = true; |
| 622 | } |
| 623 | |
| 624 | void SaveConfig(Json* p_out) const override { |
| 625 | CHECK(!this->need_configuration_) << "Call Configure before saving model."; |
no test coverage detected