use std::num::NonZeroUsize;
use std::ops::ControlFlow;
use hessboost::config::{
BalancedBagging, BoosterKind, Dart, GrowPolicy, Langevin, MaxDeltaStep, ModelShrink,
ModelShrinkMode, QueryBagging,
};
use hessboost::data::FeatureType;
use hessboost::metric::EvalMetric;
use hessboost::objective::{CustomLoss, GradPair, LambdaRank, Objective, RegLoss};
use hessboost::prelude::*;
use hessboost::training::RoundEval;
use hessboost::training::online::{OnlineMode, OnlineModel, OnlineParams};
mod common;
use common::{incompatible_model, invalid_data, invalid_param, lcg, rmse, with_threads};
const COLS: usize = 4;
fn logistic() -> Objective {
Objective::BinaryLogistic(RegLoss::default())
}
fn data(n: usize, seed: u64, binary: bool) -> DMatrix {
let mut next = lcg(seed);
let (mut x, mut y) = (Vec::new(), Vec::new());
for _ in 0..n {
let row = [next(), next(), next(), next()];
let f = 3.0 * row[0] - 2.0 * row[1] + row[2] * row[3] + 0.2 * (next() - 0.5);
x.extend(row);
y.push(if binary { f32::from(f > 0.6) } else { f });
}
DMatrix::from_dense(&x, n, COLS)
.unwrap()
.with_labels(&y)
.unwrap()
}
fn params(objective: Objective) -> TrainingParams {
TrainingParams::builder()
.objective(objective)
.tree_method(TreeMethod::Hist)
.max_depth(4)
.eta(0.3)
.build()
.unwrap()
}
#[test]
fn exact_updates_equal_retraining_bit_for_bit() {
for objective in [Objective::SquaredError(RegLoss::default()), logistic()] {
let binary = matches!(objective, Objective::BinaryLogistic(_));
let name = objective.name().to_owned();
let p = params(objective);
let train_data = data(400, 1, binary);
let mut online = OnlineModel::train(&p, &train_data, 15, OnlineParams::exact()).unwrap();
let added = data(30, 2, binary);
for (additions, deletions) in [
(Some(&added), vec![0, 7, 399]),
(None, vec![3, 4, 5]),
(Some(&added), Vec::new()),
] {
online.update(additions, &deletions).unwrap();
let retrained = train(&p, online.data(), 15).unwrap();
assert_eq!(
online.model().encode(ModelFormat::Json).unwrap(),
retrained.encode(ModelFormat::Json).unwrap(),
"{name}"
);
}
assert_eq!(online.data().n_rows(), 400 - 3 + 30 - 3 + 30);
}
}
#[test]
fn approximate_updates_stay_close_to_retraining() {
let p = params(Objective::SquaredError(RegLoss::default()));
let train_data = data(2000, 3, false);
let test = data(1000, 4, false);
let added = data(40, 5, false);
let deletions: Vec<usize> = (0..40).map(|i| i * 13).collect();
let mut online = OnlineModel::train(&p, &train_data, 30, OnlineParams::default()).unwrap();
let report = online.update(Some(&added), &deletions).unwrap();
assert_eq!(online.data().n_rows(), 2000);
assert!(report.nodes_kept > 0);
let retrained = train(&p, online.data(), 30).unwrap();
let (updated, reference) = (rmse(online.model(), &test), rmse(&retrained, &test));
assert!(
(updated - reference).abs() < 0.05 * reference,
"updated {updated} vs retrained {reference}"
);
let original = OnlineModel::train(&p, &train_data, 30, OnlineParams::default()).unwrap();
assert!(rmse(online.model(), &added) <= rmse(original.model(), &added) * 1.01);
let mut frozen =
OnlineModel::train(&p, &train_data, 30, OnlineParams::approximate(1.0).unwrap()).unwrap();
let tolerant = frozen.update(Some(&added), &deletions).unwrap();
assert!(tolerant.subtrees_regrown <= report.subtrees_regrown);
assert!(tolerant.nodes_kept >= report.nodes_kept);
}
#[test]
fn updates_ignore_the_thread_count() {
let p = params(logistic());
let train_data = data(600, 6, true);
let added = data(20, 7, true);
let run = |threads| {
with_threads(threads, || {
let mut online = OnlineModel::train(
&p,
&train_data,
10,
OnlineParams::approximate(0.05).unwrap(),
)
.unwrap();
online.update(Some(&added), &[1, 2, 3, 50]).unwrap();
online.model().encode(ModelFormat::Json).unwrap()
})
};
assert_eq!(run(1), run(4));
}
#[test]
fn an_interrupted_update_changes_nothing() {
let p = params(Objective::SquaredError(RegLoss::default()));
let train_data = data(300, 8, false);
for mode in [
OnlineParams::exact(),
OnlineParams::approximate(0.1).unwrap(),
] {
let mut online = OnlineModel::train(&p, &train_data, 10, mode).unwrap();
let before = online.model().encode(ModelFormat::Json).unwrap();
let stop = |round: RoundEval| {
if round.iteration() == 3 {
ControlFlow::Break(())
} else {
ControlFlow::Continue(())
}
};
assert_eq!(
invalid_param(online.update_with(None, &[0, 1], stop)),
"on_round"
);
assert_eq!(online.model().encode(ModelFormat::Json).unwrap(), before);
assert_eq!(online.data().n_rows(), 300);
let mut rounds = 0;
let refused = online.update_with_commit(
None,
&[0, 1],
|_| {
rounds += 1;
ControlFlow::Continue(())
},
|| ControlFlow::Break(()),
);
assert_eq!(invalid_param(refused), "on_round");
assert_eq!(rounds, 10);
assert_eq!(online.model().encode(ModelFormat::Json).unwrap(), before);
assert_eq!(online.data().n_rows(), 300);
online.update(None, &[0, 1]).unwrap();
let mut fresh = OnlineModel::train(&p, &train_data, 10, mode).unwrap();
fresh.update(None, &[0, 1]).unwrap();
assert_eq!(
online.model().encode(ModelFormat::Json).unwrap(),
fresh.model().encode(ModelFormat::Json).unwrap()
);
}
}
#[test]
fn updates_refuse_labels_retraining_refuses() {
let p = params(logistic());
let train_data = data(300, 10, true);
let bad = DMatrix::from_dense(&[0.5; COLS], 1, COLS)
.unwrap()
.with_labels(&[2.0])
.unwrap();
let good = data(5, 11, true);
for mode in [
OnlineParams::approximate(0.1).unwrap(),
OnlineParams::exact(),
] {
let mut online = OnlineModel::train(&p, &train_data, 8, mode).unwrap();
let before = online.model().encode(ModelFormat::Json).unwrap();
let retrain_err = invalid_data(train(&p, &bad, 8));
assert_eq!(
invalid_data(online.update(Some(&bad), &[0])),
retrain_err,
"{mode:?}"
);
assert_eq!(online.model().encode(ModelFormat::Json).unwrap(), before);
assert_eq!(online.data().n_rows(), 300);
online.update(Some(&good), &[0]).unwrap();
assert_eq!(online.data().n_rows(), 304);
let mut fresh = OnlineModel::train(&p, &train_data, 8, mode).unwrap();
fresh.update(Some(&good), &[0]).unwrap();
assert_eq!(
online.model().encode(ModelFormat::Json).unwrap(),
fresh.model().encode(ModelFormat::Json).unwrap()
);
}
}
#[test]
fn unsound_configurations_and_changes_are_refused() {
let train_data = data(200, 9, false);
let base = || {
TrainingParams::builder()
.tree_method(TreeMethod::Hist)
.max_depth(3)
};
let online = OnlineParams::default();
for (params, name) in [
(base().subsample(0.8).build().unwrap(), "subsample"),
(base().colsample_bytree(0.5).build().unwrap(), "subsample"),
(
base()
.grow_policy(GrowPolicy::LossGuide)
.max_leaves(8)
.build()
.unwrap(),
"grow_policy",
),
(base().max_leaves(8).build().unwrap(), "grow_policy"),
(base().unlimited_depth().build().unwrap(), "grow_policy"),
(
base()
.booster(BoosterKind::Dart(Dart::default()))
.build()
.unwrap(),
"booster",
),
(
base().num_parallel_tree(2).build().unwrap(),
"num_parallel_tree",
),
(
base().objective(Objective::AbsoluteError).build().unwrap(),
"objective",
),
(
base()
.objective(Objective::custom(CustomLoss::new(
"custom:sqerr",
1,
|p, y, _w, out| {
for (o, (p, y)) in out.iter_mut().zip(p.iter().zip(y)) {
*o = GradPair::new(p - y, 1.0);
}
},
)))
.build()
.unwrap(),
"objective",
),
(
base().langevin(Langevin::default()).build().unwrap(),
"params",
),
(
base()
.model_shrink(ModelShrink::new(0.1, ModelShrinkMode::Constant).unwrap())
.build()
.unwrap(),
"params",
),
(base().posterior_sampling(true).build().unwrap(), "params"),
] {
assert_eq!(
invalid_param(OnlineModel::train(¶ms, &train_data, 3, online)),
name
);
}
let p = base().build().unwrap();
for tolerance in [0.0, -0.1, 1.5, f64::NAN] {
assert_eq!(
invalid_param(OnlineParams::approximate(tolerance)),
"tolerance"
);
}
assert_eq!(OnlineParams::exact().mode(), OnlineMode::Exact);
let weighted = data(200, 9, false).with_weights(&[1.0; 200]).unwrap();
assert_eq!(
invalid_data(OnlineModel::train(&p, &weighted, 3, online)),
("data", None)
);
let categorical = data(200, 9, false)
.with_feature_types(&[FeatureType::Numerical; COLS])
.unwrap();
assert!(OnlineModel::train(&p, &categorical, 3, online).is_ok());
let mut model = OnlineModel::train(&p, &train_data, 3, online).unwrap();
assert_eq!(invalid_param(model.update(None, &[200])), "deletions");
assert_eq!(invalid_param(model.update(None, &[4, 4])), "deletions");
let all: Vec<usize> = (0..200).collect();
assert_eq!(invalid_param(model.update(None, &all)), "deletions");
let unlabelled = DMatrix::from_dense(&[0.0; COLS], 1, COLS).unwrap();
assert_eq!(
invalid_data(model.update(Some(&unlabelled), &[])),
("additions", None)
);
let narrow = DMatrix::from_dense(&[0.0; 2], 1, 2)
.unwrap()
.with_labels(&[1.0])
.unwrap();
assert_eq!(
invalid_data(model.update(Some(&narrow), &[])),
("additions", None)
);
}
#[test]
fn row_bagging_is_refused() {
let balanced = TrainingParams::builder()
.objective(logistic())
.tree_method(TreeMethod::Hist)
.max_depth(3)
.balanced_bagging(BalancedBagging::new(0.5, 0.8).unwrap())
.build()
.unwrap();
let binary = data(200, 9, true);
let ranking = TrainingParams::builder()
.objective(Objective::RankPairwise(LambdaRank::default()))
.tree_method(TreeMethod::Hist)
.max_depth(3)
.bagging_by_query(QueryBagging::new(0.5).unwrap())
.build()
.unwrap();
let queries = data(200, 9, true).with_group_sizes(&[50; 4]).unwrap();
assert!(train(&balanced, &binary, 3).is_ok());
assert!(train(&ranking, &queries, 3).is_ok());
let model = train(¶ms(logistic()), &binary, 3).unwrap();
for online in [
OnlineParams::approximate(0.1).unwrap(),
OnlineParams::exact(),
] {
for (p, d) in [(&balanced, &binary), (&ranking, &queries)] {
assert_eq!(
invalid_param(OnlineModel::train(p, d, 3, online)),
"params",
"train, {online:?}"
);
assert_eq!(
invalid_param(OnlineModel::from_model(model.clone(), p, d, online)),
"params",
"from_model, {online:?}"
);
}
}
}
#[test]
fn from_model_refuses_linear_leaves() {
let path = concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/data/lightgbm-4.7.0-linear.txt"
);
let model = BoostedModel::load(path, ModelFormat::LightgbmText).unwrap();
assert_eq!(model.n_features(), 6);
let mut next = lcg(12);
let x: Vec<f32> = (0..200 * 6).map(|_| next()).collect();
let y: Vec<f32> = x.chunks(6).map(|r| r[0] - r[1]).collect();
let data = DMatrix::from_dense(&x, 200, 6)
.unwrap()
.with_labels(&y)
.unwrap();
let p = params(Objective::SquaredError(RegLoss::default()));
for online in [
OnlineParams::approximate(0.1).unwrap(),
OnlineParams::exact(),
] {
assert_eq!(
incompatible_model(OnlineModel::from_model(model.clone(), &p, &data, online)),
"model"
);
let trained = train(&p, &data, 3).unwrap();
assert!(OnlineModel::from_model(trained, &p, &data, online).is_ok());
}
}
#[test]
fn from_model_refuses_shrunk_models() {
let data = data(200, 5, false);
let p = params(Objective::SquaredError(RegLoss::default()));
let mut shrunk = p.clone();
shrunk.model_shrink = Some(ModelShrink::new(0.1, ModelShrinkMode::Constant).unwrap());
let model = train(&shrunk, &data, 3).unwrap();
for online in [
OnlineParams::approximate(0.1).unwrap(),
OnlineParams::exact(),
] {
assert_eq!(
incompatible_model(OnlineModel::from_model(model.clone(), &p, &data, online)),
"model"
);
}
}
#[test]
fn from_model_refuses_trees_deeper_than_max_depth() {
let d = data(200, 14, false);
let p = params(Objective::SquaredError(RegLoss::default()));
let mut deeper = p.clone();
deeper.max_depth = p.max_depth.and_then(|depth| depth.checked_add(2));
let model = train(&deeper, &d, 3).unwrap();
for online in [
OnlineParams::approximate(0.1).unwrap(),
OnlineParams::exact(),
] {
assert_eq!(
incompatible_model(OnlineModel::from_model(model.clone(), &p, &d, online)),
"model"
);
assert!(OnlineModel::from_model(model.clone(), &deeper, &d, online).is_ok());
}
}
#[test]
fn an_abandoned_update_keeps_the_update_state() {
let p = params(Objective::SquaredError(RegLoss::default()));
let train_data = data(600, 13, false);
let (a, b, c) = (
data(30, 14, false),
data(20, 15, false),
data(25, 16, false),
);
for online in [
OnlineParams::approximate(0.1).unwrap(),
OnlineParams::exact(),
] {
let mut online = OnlineModel::train(&p, &train_data, 10, online).unwrap();
online.update(Some(&a), &[0, 5, 9]).unwrap();
let mut control = online.clone();
let stop = |round: RoundEval| {
if round.iteration() == 4 {
ControlFlow::Break(())
} else {
ControlFlow::Continue(())
}
};
assert_eq!(
invalid_param(online.update_with(Some(&b), &[1, 2], stop)),
"on_round"
);
let refused = online.update_with_commit(
Some(&b),
&[1, 2],
|_| ControlFlow::Continue(()),
|| ControlFlow::Break(()),
);
assert_eq!(invalid_param(refused), "on_round");
assert_eq!(
online.model().encode(ModelFormat::Json).unwrap(),
control.model().encode(ModelFormat::Json).unwrap()
);
for (additions, deletions) in [(Some(&c), vec![3, 7]), (None, vec![0, 1, 40])] {
let report = online.update(additions, &deletions).unwrap();
assert_eq!(report, control.update(additions, &deletions).unwrap());
assert_eq!(
online.model().encode(ModelFormat::Json).unwrap(),
control.model().encode(ModelFormat::Json).unwrap(),
"{online:?}"
);
}
}
}
fn rows(x: &[f32], cols: usize, y: &[f32]) -> DMatrix {
DMatrix::from_dense(x, y.len(), cols)
.unwrap()
.with_labels(y)
.unwrap()
}
fn plain(max_depth: usize) -> TrainingParams {
TrainingParams::builder()
.tree_method(TreeMethod::Hist)
.max_depth(max_depth)
.eta(1.0)
.lambda(0.0)
.base_score(0.0)
.build()
.unwrap()
}
#[test]
fn from_model_refuses_models_and_metrics_training_would_not_give() {
let d = data(300, 11, false);
let valid = data(100, 12, false);
let p = params(Objective::SquaredError(RegLoss::default()));
let stopped = Trainer::new(&p, &d, 200)
.eval(&valid, "valid")
.early_stopping_rounds(NonZeroUsize::new(2).unwrap())
.train()
.unwrap()
.model;
let best = stopped.best_iteration().expect("stops early");
for online in [OnlineParams::default(), OnlineParams::exact()] {
assert_eq!(
incompatible_model(OnlineModel::from_model(stopped.clone(), &p, &d, online)),
"model"
);
let best_slice = stopped.slice(..=best, 1).unwrap();
assert!(OnlineModel::from_model(best_slice, &p, &d, online).is_ok());
}
let counts = rows(&[0.0, 0.0, 1.0, 1.0], 1, &[0.0, 0.0, 1.0, 1.0]);
let poisson = TrainingParams::builder()
.objective(Objective::Poisson)
.tree_method(TreeMethod::Hist)
.max_depth(1)
.build()
.unwrap();
let model = train(&poisson, &counts, 2).unwrap();
let mut unbounded = poisson.clone();
unbounded.max_delta_step = MaxDeltaStep::Unbounded;
assert_eq!(
incompatible_model(OnlineModel::from_model(
model.clone(),
&unbounded,
&counts,
OnlineParams::default()
)),
"model"
);
assert!(OnlineModel::from_model(model, &poisson, &counts, OnlineParams::default()).is_ok());
let binary = data(100, 13, true);
let logistic_params = params(logistic());
let model = train(&logistic_params, &binary, 2).unwrap();
let mut multiclass_metric = logistic_params.clone();
multiclass_metric.eval_metric = vec![EvalMetric::MLogLoss];
assert_eq!(
invalid_param(OnlineModel::from_model(
model,
&multiclass_metric,
&binary,
OnlineParams::default()
)),
"eval_metric"
);
}
#[test]
fn the_exact_mode_updates_categorical_models() {
let d = rows(
&[0.0, 0.0, 1.0, 1.0, 2.0, 2.0],
1,
&[0.0, 0.0, 1.0, 1.0, 3.0, 3.0],
)
.with_feature_types(&[FeatureType::Categorical])
.unwrap();
let p = plain(2);
let mut exact = OnlineModel::train(&p, &d, 2, OnlineParams::exact()).unwrap();
assert!(
exact.model().trees()[0]
.nodes()
.iter()
.any(|n| n.is_categorical)
);
let added = rows(&[1.0], 1, &[2.0])
.with_feature_types(&[FeatureType::Categorical])
.unwrap();
exact.update(Some(&added), &[0]).unwrap();
assert_eq!(
exact.model().trees(),
train(&p, exact.data(), 2).unwrap().trees()
);
assert_eq!(
invalid_param(OnlineModel::train(&p, &d, 2, OnlineParams::default())),
"data"
);
}
#[test]
fn updates_refuse_data_without_a_finite_intercept() {
let poisson = TrainingParams::builder()
.objective(Objective::Poisson)
.tree_method(TreeMethod::Hist)
.max_depth(1)
.build()
.unwrap();
let counts = rows(&[0.0, 1.0], 1, &[0.0, 1.0]);
for online_params in [OnlineParams::default(), OnlineParams::exact()] {
let mut online = OnlineModel::train(&poisson, &counts, 2, online_params).unwrap();
assert_eq!(
invalid_data(train(&poisson, &rows(&[0.0], 1, &[0.0]), 2)),
("labels", None)
);
assert_eq!(invalid_data(online.update(None, &[1])), ("labels", None));
assert_eq!(online.data().n_rows(), 2);
}
}
#[test]
fn updates_run_on_the_configured_threads() {
let mut p = params(Objective::SquaredError(RegLoss::default()));
p.nthread = std::num::NonZeroUsize::new(1);
let d = data(200, 14, false);
with_threads(4, || {
let mut online = OnlineModel::train(&p, &d, 3, OnlineParams::default()).unwrap();
let mut seen = Vec::new();
online
.update_with(None, &[0], |_| {
seen.push(rayon::current_num_threads());
ControlFlow::Continue(())
})
.unwrap();
assert_eq!(seen, vec![1; 3]);
});
}
#[test]
fn updates_reach_splits_the_cached_rows_never_did() {
let nan = f32::NAN;
let p = plain(2);
let full = rows(
&[nan, 0.0, nan, 1.0, 0.0, 0.0, 1.0, 1.0],
2,
&[0.0, 1.0, 10.0, 11.0],
);
let model = train(&p, &full, 1).unwrap();
let observed = rows(&[0.0, 0.0, 1.0, 1.0], 2, &[10.0, 11.0]);
let online = OnlineParams::approximate(1.0).unwrap();
let mut resumed = OnlineModel::from_model(model, &p, &observed, online).unwrap();
let missing = rows(&[nan, 0.0], 2, &[0.0]);
assert!(resumed.update(Some(&missing), &[]).is_ok());
let start = rows(
&[0.0, 0.0, 1.0, 1.0, nan, 0.0, nan, 1.0],
2,
&[0.0, 10.0, 1.0, 11.0],
);
let mut online_model = OnlineModel::train(&p, &start, 1, online).unwrap();
let replacement = rows(
&[1.0, 0.0, 1.0, 1.0, nan, 0.0, nan, 1.0],
2,
&[0.0, 1.0, 10.0, 11.0],
);
online_model
.update(Some(&replacement), &[0, 1, 2, 3])
.unwrap();
assert!(
online_model
.update(Some(&rows(&[0.0, 0.0], 2, &[0.0])), &[])
.is_ok()
);
assert_eq!(online_model.data().n_rows(), 5);
}
#[test]
fn approximate_updates_refuse_values_beyond_the_training_bins() {
let nan = f32::NAN;
let mut p = plain(1);
p.min_child_weight = 2.0;
let d = rows(&[0.0, 1.0, nan, nan], 1, &[0.0, 0.0, 1.0, 1.0]);
let beyond = rows(&[3.0], 1, &[0.0]);
let mut approximate =
OnlineModel::train(&p, &d, 1, OnlineParams::approximate(1.0).unwrap()).unwrap();
let before = approximate.model().clone();
assert_eq!(
invalid_data(approximate.update(Some(&beyond), &[0])),
("additions", None)
);
assert_eq!(approximate.model().trees(), before.trees());
assert_eq!(approximate.data().n_rows(), 4);
assert!(
approximate
.update(Some(&rows(&[0.5], 1, &[0.0])), &[0])
.is_ok()
);
let mut exact = OnlineModel::train(&p, &d, 1, OnlineParams::exact()).unwrap();
exact.update(Some(&beyond), &[0]).unwrap();
assert_eq!(
exact.model().trees(),
train(&p, exact.data(), 1).unwrap().trees()
);
}
#[test]
fn updates_that_overflow_are_refused_and_change_nothing() {
let max = f32::MAX;
let p = TrainingParams::builder()
.tree_method(TreeMethod::Hist)
.base_score(f64::from(max))
.build()
.unwrap();
let mut online =
OnlineModel::train(&p, &rows(&[0.0], 1, &[max]), 2, OnlineParams::default()).unwrap();
let before = online.model().clone();
let err = online
.update(Some(&rows(&[0.0], 1, &[-max])), &[])
.unwrap_err();
assert!(matches!(err, HessboostError::ModelFormat(_)), "{err}");
assert_eq!(online.model().trees(), before.trees());
assert_eq!(online.data().n_rows(), 1);
}
#[test]
fn approximate_from_model_refuses_leaves_other_parameters_made() {
let text = "tree\nversion=v4\nnum_class=1\nnum_tree_per_iteration=1\nlabel_index=0\n\
max_feature_idx=0\nobjective=regression\nfeature_names=x\nfeature_infos=none\n\n\
Tree=0\nnum_leaves=1\nnum_cat=0\nleaf_value=5\nleaf_weight=4\nleaf_count=4\n\
is_linear=0\nshrinkage=1\n\n\nend of trees\n";
let imported = BoostedModel::decode(text, ModelFormat::LightgbmText).unwrap();
let fives = rows(&[0.0, 1.0, 2.0, 3.0], 1, &[5.0; 4]);
let mut p = plain(1);
p.eta = 0.3;
p.base_score = None;
let approximate = OnlineParams::default();
assert_eq!(
incompatible_model(OnlineModel::from_model(
imported.clone(),
&p,
&fives,
approximate
)),
"model"
);
assert!(OnlineModel::from_model(imported, &p, &fives, OnlineParams::exact()).is_ok());
let d = data(300, 21, false);
let trained = params(Objective::SquaredError(RegLoss::default()));
let model = train(&trained, &d, 5).unwrap();
let mut other_eta = trained.clone();
other_eta.eta = 0.5;
assert_eq!(
incompatible_model(OnlineModel::from_model(
model.clone(),
&other_eta,
&d,
approximate
)),
"model"
);
let mut online = OnlineModel::from_model(model, &trained, &d, approximate).unwrap();
online
.update(Some(&data(20, 22, false)), &[0, 5, 9])
.unwrap();
let saved = online.model().clone();
assert!(OnlineModel::from_model(saved, &trained, online.data(), approximate).is_ok());
}
#[test]
fn refreshed_margins_follow_the_prediction_recurrence() {
let mut p = plain(1);
p.base_score = Some(1e8);
let d = rows(&[0.0, 1.0, 2.0, 3.0], 1, &[0.0, 1.0, 2.0, 3.0]);
let mut online =
OnlineModel::train(&p, &d, 3, OnlineParams::approximate(0.01).unwrap()).unwrap();
online.update(Some(&rows(&[0.0], 1, &[10.0])), &[]).unwrap();
let data = online.data();
let prefix = online
.model()
.slice(..2, 1)
.unwrap()
.predict(data, ..)
.unwrap();
let labels: Vec<f32> = data.labels().unwrap().to_vec();
let x: Vec<f32> = (0..data.n_rows())
.map(|i| data.get(i, 0).unwrap())
.collect();
let residuals: Vec<f32> = labels
.iter()
.zip(prefix.as_slice())
.map(|(y, m)| y - m)
.collect();
let mut correction_params = plain(1);
correction_params.base_score = Some(0.0);
let correction = train(&correction_params, &rows(&x, 1, &residuals), 1)
.unwrap()
.predict(&rows(&x, 1, &residuals), ..)
.unwrap();
let expected: Vec<f32> = prefix
.as_slice()
.iter()
.zip(correction.as_slice())
.map(|(m, c)| m + c)
.collect();
assert_eq!(
online.model().predict(data, ..).unwrap().as_slice(),
expected
);
}