use crate::config::{BoosterKind, GrowPolicy, Monotone, ProcessType, TrainingParams};
use crate::data::DMatrix;
use crate::error::{HessboostError, Result};
use crate::model::{BoostedModel, ModelObjective};
use crate::objective::Loss;
use crate::training::multi_output;
pub(super) fn resume_model(
init: &BoostedModel,
params: &TrainingParams,
objective: &dyn Loss,
dtrain: &DMatrix,
num_boost_round: usize,
intercepts: impl FnOnce() -> Result<Vec<f32>>,
) -> Result<BoostedModel> {
init.validate_structure().map_err(|e| {
let reason = match e {
HessboostError::ModelFormat(reason) => reason,
other => other.to_string(),
};
HessboostError::model_format(format!("invalid init model: {reason}"))
})?;
if init.shrinkage().is_some() {
return Err(HessboostError::incompatible_model(
"init_model",
"a model trained with model shrinkage cannot be trained further",
));
}
if params.model_shrinkage_on() {
return Err(HessboostError::invalid_param(
"model_shrink_rate",
"model shrinkage is not supported with continued training",
));
}
if matches!(params.booster, BoosterKind::Boulevard(_)) || init.boulevard().is_some() {
return Err(HessboostError::incompatible_model(
"init_model",
"Boulevard models average every round of one run and cannot be trained further, \
and `booster = boulevard` cannot continue another model",
));
}
if matches!(params.booster, BoosterKind::Ebm(_)) || init.ebm().is_some() {
return Err(HessboostError::incompatible_model(
"init_model",
"EBMs are boosted term by term in one run and cannot be trained further, and \
`booster = ebm` cannot continue another model",
));
}
let is_linear = init.linear().is_some();
if is_linear != (params.booster == BoosterKind::GbLinear) {
return Err(HessboostError::incompatible_model(
"booster",
if is_linear {
"a gblinear model can only be trained further with booster=gblinear"
} else {
"a tree model cannot be trained further with booster=gblinear"
},
));
}
if params.objective.name() != init.objective().name() {
return Err(HessboostError::incompatible_model(
"objective",
format!(
"`{}` does not match the model's objective `{}`",
params.objective.name(),
init.objective().name()
),
));
}
if objective.n_outputs() != init.n_outputs() {
return Err(HessboostError::incompatible_model(
"num_class",
format!(
"the objective's {} outputs do not match the model's {} (num_class {})",
objective.n_outputs(),
init.n_outputs(),
init.num_class()
),
));
}
if dtrain.n_cols() != init.n_features() {
return Err(HessboostError::dimension_mismatch(
"continued-training feature count",
init.n_features(),
dtrain.n_cols(),
));
}
if dtrain.n_targets() != init.n_targets() {
return Err(HessboostError::dimension_mismatch(
"continued-training label targets",
init.n_targets(),
dtrain.n_targets(),
));
}
if !is_linear && init.num_trees() > 0 {
if params.num_parallel_tree != init.num_parallel_tree() {
return Err(HessboostError::incompatible_model(
"num_parallel_tree",
format!(
"{} does not match the model's num_parallel_tree {}",
params.num_parallel_tree,
init.num_parallel_tree()
),
));
}
if init.has_vector_leaves() != multi_output::vector_leaf(params, objective.n_outputs()) {
return Err(HessboostError::incompatible_model(
"multi_strategy",
if init.has_vector_leaves() {
"a vector-leaf model can only be trained further with \
`multi_strategy=multi_output_tree`"
} else {
"a one-output-per-tree model cannot be trained further with \
`multi_strategy=multi_output_tree`"
},
));
}
}
if matches!(params.process_type, ProcessType::Update(_)) {
check_update(init, params, num_boost_round)?;
}
let mut model = init.clone();
model.set_best_iteration(None);
model.set_objective(
ModelObjective::trained_with(¶ms.objective),
params.effective_max_delta_step(),
);
model.set_num_parallel_tree(params.num_parallel_tree);
model.materialize_tree_weights();
if params.base_score.is_some() {
model.set_base_scores(intercepts()?);
}
Ok(model)
}
fn check_update(
init: &BoostedModel,
params: &TrainingParams,
num_boost_round: usize,
) -> Result<()> {
if params.booster != BoosterKind::GbTree {
return Err(HessboostError::invalid_param(
"process_type",
"`update` refreshes gbtree models only (booster=gbtree)",
));
}
if init.has_vector_leaves() {
return Err(HessboostError::incompatible_model(
"process_type",
"`update` cannot refresh vector-leaf trees (`multi_output_tree`)",
));
}
if init.has_non_unit_tree_weights() {
return Err(HessboostError::incompatible_model(
"process_type",
"`update` cannot refresh a model with DART tree weights",
));
}
if params
.monotone_constraints
.iter()
.any(|&m| m != Monotone::None)
{
return Err(HessboostError::invalid_param(
"monotone_constraints",
"are not supported by the refresh updater (`process_type=update`)",
));
}
if init
.trees()
.iter()
.any(|tree| tree.linear_leaves().is_some())
{
return Err(HessboostError::incompatible_model(
"process_type",
"`update` cannot refresh linear-leaf trees (`linear_tree`)",
));
}
reject_unused_by_refresh(params)?;
if num_boost_round > init.num_boost_rounds() {
return Err(HessboostError::incompatible_model(
"num_boost_round",
format!(
"{num_boost_round} exceeds the {} iterations `process_type=update` can refresh",
init.num_boost_rounds()
),
));
}
Ok(())
}
fn reject_unused_by_refresh(params: &TrainingParams) -> Result<()> {
let p = params.clone();
let reference = TrainingParams {
booster: p.booster,
nthread: p.nthread,
seed: p.seed,
device: p.device,
objective: p.objective,
base_score: p.base_score,
eval_metric: p.eval_metric,
eta: p.eta,
lambda: p.lambda,
alpha: p.alpha,
max_delta_step: p.max_delta_step,
num_parallel_tree: p.num_parallel_tree,
multi_strategy: p.multi_strategy,
monotone_constraints: p.monotone_constraints,
process_type: p.process_type,
tree_method: p.tree_method,
max_depth: p.max_depth,
max_leaves: p.max_leaves,
min_child_weight: p.min_child_weight,
gamma: p.gamma,
max_bin: p.max_bin,
interaction_constraints: p.interaction_constraints,
grow_policy: match p.grow_policy {
GrowPolicy::Symmetric => GrowPolicy::default(),
policy => policy,
},
..TrainingParams::default()
};
params.refuse_changes_from(
&reference,
"process_type",
"`update` keeps the existing splits and refreshes them from every row, so it applies \
no sampling, growth, or split-search options",
)
}
pub(super) fn require_model_for_update(params: &TrainingParams) -> Result<()> {
if matches!(params.process_type, ProcessType::Update(_)) {
return Err(HessboostError::invalid_param(
"process_type",
"`update` refreshes an existing model; use Trainer::init_model",
));
}
Ok(())
}
#[cfg(test)]
mod tests {
use crate::config::TrainingParams;
use crate::error::HessboostError;
use crate::test_support::labeled_dense;
use crate::training::{Trainer, train};
#[test]
fn invalid_init_models_are_refused() {
let x: Vec<f32> = (0..20).map(|i| i as f32).collect();
let d = labeled_dense(&x, 20, 1, &x);
let params = TrainingParams::default();
let mut model = train(¶ms, &d, 2).unwrap();
model.set_base_scores(Vec::new());
assert!(matches!(
Trainer::new(¶ms, &d, 1).init_model(&model).train(),
Err(HessboostError::ModelFormat(_))
));
}
}