use std::time::Instant;
use hessboost::prelude::*;
use hessboost::training::online::{OnlineMode, OnlineModel, OnlineParams};
mod common;
use common::lcg;
const COLS: usize = 10;
fn friedman(n: usize, seed: u64) -> Result<DMatrix> {
let mut next = lcg(seed);
let (mut x, mut y) = (Vec::with_capacity(n * COLS), Vec::with_capacity(n));
for _ in 0..n {
let row: Vec<f32> = (0..COLS).map(|_| next()).collect();
let f = 10.0 * (std::f32::consts::PI * row[0] * row[1]).sin()
+ 20.0 * (row[2] - 0.5).powi(2)
+ 10.0 * row[3]
+ 5.0 * row[4];
y.push(f + 2.0 * (next() + next() + next() - 1.5));
x.extend(row);
}
DMatrix::from_dense(&x, n, COLS)?.with_labels(&y)
}
fn rmse(model: &BoostedModel, data: &DMatrix) -> Result<f64> {
let preds = model.predict(data, Iterations::Best)?;
let metric = EvalMetric::Rmse.build(1)?;
Ok(metric.eval(preds.as_slice(), data.labels().unwrap_or_default(), None))
}
fn main() -> Result<()> {
let (data, test, new_rows) = (friedman(20_000, 1)?, friedman(5000, 2)?, friedman(200, 3)?);
let params = TrainingParams::builder()
.tree_method(TreeMethod::Hist)
.max_depth(6)
.eta(0.1)
.build()?;
let rounds = 100;
let deletions: Vec<usize> = (0..200).map(|i| i * 97).collect();
let start = Instant::now();
let mut online = OnlineModel::train(¶ms, &data, rounds, OnlineParams::default())?;
println!("trained with update state in {:.1?}", start.elapsed());
let original = online.model().clone();
let start = Instant::now();
let report = online.update(Some(&new_rows), &deletions)?;
let update_time = start.elapsed();
let start = Instant::now();
let retrained = train(¶ms, online.data(), rounds)?;
let retrain_time = start.elapsed();
println!(
"add 200 + delete 200 rows: update {update_time:.1?} vs retrain {retrain_time:.1?} \
({:.1}x); kept {} nodes, regrew {} subtrees, refreshed {} rows",
retrain_time.as_secs_f64() / update_time.as_secs_f64(),
report.nodes_kept,
report.subtrees_regrown,
report.rows_refreshed,
);
println!(
"test RMSE: original {:.4}, updated {:.4}, retrained {:.4}",
rmse(&original, &test)?,
rmse(online.model(), &test)?,
rmse(&retrained, &test)?,
);
let mut exact = OnlineModel::train(¶ms, &data, rounds, OnlineParams::exact())?;
assert_eq!(exact.online_params().mode(), OnlineMode::Exact);
exact.update(None, &deletions)?;
let reference = train(¶ms, exact.data(), rounds)?;
println!(
"exact mode after deleting 200 rows equals retraining bit for bit: {}",
exact.model().encode(ModelFormat::Json)? == reference.encode(ModelFormat::Json)?
);
Ok(())
}