use hessboost::config::Monotone;
use hessboost::data::FeatureType;
use hessboost::metric::EvalMetric;
use hessboost::prelude::*;
mod common;
use common::{fill_random, lcg};
fn main() -> Result<()> {
let n = 400usize;
let mut rng = lcg(1);
let mut x = vec![0f32; n];
let mut y = vec![0f32; n];
for i in 0..n {
x[i] = rng();
y[i] = x[i] + 0.3 * (rng() - 0.5);
}
let d = DMatrix::from_dense(&x, n, 1)?.with_labels(&y)?;
let params = TrainingParams::builder()
.objective(Objective::SquaredError(RegLoss::default()))
.monotone_constraints(vec![Monotone::Increasing]) .max_depth(4)
.eta(0.2)
.build()?;
let model = train(¶ms, &d, 60)?;
let mut idx: Vec<usize> = (0..n).collect();
idx.sort_by(|&a, &b| x[a].partial_cmp(&x[b]).unwrap());
let preds = model.predict(&d, Iterations::Best)?.into_vec(); let monotone = idx.windows(2).all(|w| preds[w[1]] >= preds[w[0]] - 1e-5);
println!("monotone constraint respected: {monotone}");
let (n2, f2) = (600usize, 4usize);
let mut x2 = vec![0f32; n2 * f2];
let mut y2 = vec![0f32; n2];
fill_random(&mut rng, &mut x2);
for i in 0..n2 {
y2[i] = x2[i * f2] * x2[i * f2 + 1] + x2[i * f2 + 2];
}
let d2 = DMatrix::from_dense(&x2, n2, f2)?.with_labels(&y2)?;
let params2 = TrainingParams::builder()
.objective(Objective::SquaredError(RegLoss::default()))
.interaction_constraints(vec![vec![0, 1], vec![2, 3]])
.max_depth(4)
.eta(0.2)
.build()?;
let m2 = train(¶ms2, &d2, 40)?;
println!("interaction-constrained model: {} trees", m2.num_trees());
let cats = [0.0f32, 1.0, 2.0, 3.0];
let mut xc = Vec::new();
let mut yc = Vec::new();
for _ in 0..100 {
for &c in &cats {
xc.push(c);
yc.push(if (c as u32) % 2 == 1 { 1.0 } else { 0.0 }); }
}
let dc = DMatrix::from_dense(&xc, xc.len(), 1)?
.with_labels(&yc)?
.with_feature_types(&[FeatureType::Categorical])?;
let mc = train(
&TrainingParams::builder()
.objective(Objective::SquaredError(RegLoss::default()))
.max_depth(2)
.eta(0.3)
.build()?,
&dc,
30,
)?;
let pc = mc.predict(&dc, Iterations::Best)?;
let rmse = EvalMetric::Rmse.build(1)?.eval(pc.as_slice(), &yc, None);
println!("categorical fit RMSE on non-ordinal pattern: {rmse:.4}");
Ok(())
}