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

Method LoadConfig

src/learner.cc:569–622  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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.";

Callers 4

LoadModelMethod · 0.45
LoadMethod · 0.45
SliceMethod · 0.45
EvalOneIterMethod · 0.45

Calls 8

WarnOldSerializationFunction · 0.85
FromJsonFunction · 0.85
CanonicalizeBoosterNameFunction · 0.85
resizeMethod · 0.80
resetMethod · 0.45
TaskMethod · 0.45
sizeMethod · 0.45
FromJsonMethod · 0.45

Tested by

no test coverage detected