use std::num::NonZeroUsize;
use std::time::Instant;
use hessboost::diffusion::DiffusionFormat;
use hessboost::diffusion::{DiffusionModel, DiffusionParams, SampleOptions, Samples};
use hessboost::objective::distributional::{DistFamily, Distributional};
use hessboost::prelude::*;
mod common;
use common::lcg;
fn normal(next: &mut impl FnMut() -> f32) -> f32 {
let (u1, u2) = (1.0 - next(), next());
(-2.0 * u1.ln()).sqrt() * (std::f32::consts::TAU * u2).cos()
}
fn bimodal(n: usize, seed: u64) -> Result<DMatrix> {
let mut next = lcg(seed);
let (mut x, mut y) = (Vec::with_capacity(n), Vec::with_capacity(n));
for _ in 0..n {
let xi = next();
let sign = if next() < 0.5 { -1.0 } else { 1.0 };
x.push(xi);
y.push(sign * (1.0 + xi) + 0.1 * normal(&mut next));
}
DMatrix::from_dense(&x, n, 1)?.with_labels(&y)
}
fn heteroscedastic(n: usize, seed: u64) -> Result<DMatrix> {
let mut next = lcg(seed);
let (mut x, mut y) = (Vec::with_capacity(2 * n), Vec::with_capacity(n));
for _ in 0..n {
let (x0, x1) = (next(), next());
let e = -(1.0 - next()).ln();
x.extend_from_slice(&[x0, x1]);
y.push((std::f32::consts::TAU * x0).sin() + (0.1 + x1) * (e - 1.0));
}
DMatrix::from_dense(&x, n, 2)?.with_labels(&y)
}
fn mean(values: &[f64]) -> f64 {
values.iter().sum::<f64>() / values.len() as f64
}
fn normal_crps(train: &DMatrix, test: &DMatrix) -> Result<f64> {
let params = TrainingParams::builder()
.objective(Objective::Dist(Distributional::new(DistFamily::Normal)))
.tree_method(TreeMethod::Hist)
.eta(0.05)
.max_depth(4)
.build()?;
let model = Trainer::new(¶ms, train, 1000)
.eval(test, "test")
.early_stopping_rounds(NonZeroUsize::new(50).unwrap())
.train()?
.model;
let labels = test.labels().unwrap_or_default();
let crps: Vec<f64> = model
.predict_distribution(test, Iterations::Best)?
.iter()
.zip(labels)
.map(|(d, &y)| d.crps(f64::from(y)))
.collect();
Ok(mean(&crps))
}
fn coverage90(samples: &Samples, labels: &[f32]) -> Result<f64> {
let q = samples.quantiles(&[0.05, 0.95])?;
let inside = labels
.iter()
.enumerate()
.filter(|&(row, &y)| match (q.get(row, 0), q.get(row, 1)) {
(Some(lo), Some(hi)) => lo[0] <= f64::from(y) && f64::from(y) <= hi[0],
_ => false,
})
.count();
Ok(inside as f64 / labels.len() as f64)
}
fn histogram(samples: &Samples, row: usize) -> String {
const BINS: usize = 40;
let mut counts = [0usize; BINS];
for &v in samples.row(row).unwrap_or_default() {
let bin = ((f64::from(v) + 2.5) / 5.0 * BINS as f64).floor();
if (0.0..BINS as f64).contains(&bin) {
counts[bin as usize] += 1;
}
}
let max = counts.iter().copied().max().unwrap_or(1).max(1);
counts
.iter()
.map(|&c| [' ', '.', ':', '|', '#'][(c * 4).div_ceil(max)])
.collect()
}
fn report(name: &str, params: &DiffusionParams, train: &DMatrix, test: &DMatrix) -> Result<()> {
let start = Instant::now();
let model = DiffusionModel::fit(params, train)?;
let fit_time = start.elapsed();
let start = Instant::now();
let samples = model.sample(test, 100, &SampleOptions::seeded(1))?;
let sample_time = start.elapsed();
let labels = test.labels().unwrap_or_default();
println!(
" {name:<14} CRPS {:.4} 90% coverage {:.3} ({} rounds, fit {:.1?}, 100 samples × {} rows {:.1?})",
mean(samples.crps(labels)?.as_slice()),
coverage90(&samples, labels)?,
model.regressor().num_boost_rounds(),
fit_time,
test.n_rows(),
sample_time,
);
Ok(())
}
fn main() -> Result<()> {
let (train, test) = (bimodal(2000, 1)?, bimodal(500, 2)?);
println!("bimodal: y = ±(1 + x) + 0.1 ε");
println!(" dist:normal CRPS {:.4}", normal_crps(&train, &test)?);
report("score (EDM)", &DiffusionParams::default(), &train, &test)?;
report("treeffuser", &DiffusionParams::treeffuser(), &train, &test)?;
report(
"flow matching",
&DiffusionParams::flow_matching(),
&train,
&test,
)?;
let model = DiffusionModel::fit(&DiffusionParams::default(), &train)?;
let probes = DMatrix::from_dense(&[0.1, 0.9], 2, 1)?;
let samples = model.sample(&probes, 2000, &SampleOptions::seeded(3))?;
let q = samples.quantiles(&[0.1, 0.25, 0.5, 0.75, 0.9])?;
for (row, x) in [0.1, 0.9].into_iter().enumerate() {
println!(
" x = {x}: quantiles 10/25/50/75/90% = {:?}, modes at ±{:.1}",
(0..q.n_levels())
.filter_map(|level| q.get(row, level))
.map(|v| (v[0] * 100.0).round() / 100.0)
.collect::<Vec<_>>(),
1.0 + x,
);
println!(" [-2.5 {} 2.5]", histogram(&samples, row));
}
let (train, test) = (heteroscedastic(2000, 3)?, heteroscedastic(500, 4)?);
println!("heteroscedastic, right-skewed: y = sin(2π x0) + (0.1 + x1)(E - 1)");
println!(" dist:normal CRPS {:.4}", normal_crps(&train, &test)?);
report("score (EDM)", &DiffusionParams::default(), &train, &test)?;
report(
"flow matching",
&DiffusionParams::flow_matching(),
&train,
&test,
)?;
let n = 1500;
let mut next = lcg(5);
let (mut x, mut y) = (Vec::with_capacity(n), Vec::with_capacity(2 * n));
for _ in 0..n {
let xi = next();
let u = xi + 0.5 * normal(&mut next);
x.push(xi);
y.extend_from_slice(&[u, u * u + 0.05 * normal(&mut next)]);
}
let data = DMatrix::from_dense(&x, n, 1)?.with_label_matrix(&y, 2)?;
let model = DiffusionModel::fit(&DiffusionParams::treeffuser(), &data)?;
let probe = DMatrix::from_dense(&[0.5], 1, 1)?;
let draws = model.sample(&probe, 2000, &SampleOptions::seeded(4))?;
let pairs: Vec<(f64, f64)> = draws
.as_slice()
.as_chunks::<2>()
.0
.iter()
.map(|p| (f64::from(p[0]), f64::from(p[1])))
.collect();
let median_gap = |shift: usize| {
let mut gaps: Vec<f64> = (0..pairs.len())
.map(|i| (pairs[(i + shift) % pairs.len()].1 - pairs[i].0.powi(2)).abs())
.collect();
gaps.sort_by(f64::total_cmp);
gaps[gaps.len() / 2]
};
println!(
"2-D label at x = 0.5: mean {:?} (true [0.5, 0.5]), median |y2 - y1²| {:.3} within draws \
vs {:.3} across draws (true noise alone: 0.034)",
draws
.mean()
.as_slice()
.iter()
.map(|v| (v * 100.0).round() / 100.0)
.collect::<Vec<_>>(),
median_gap(0),
median_gap(1),
);
let from_bytes = DiffusionModel::decode(
&model.encode(DiffusionFormat::Binary)?,
DiffusionFormat::Binary,
)?;
let from_json =
DiffusionModel::decode(&model.encode(DiffusionFormat::Json)?, DiffusionFormat::Json)?;
assert_eq!(
from_bytes.sample(&probe, 50, &SampleOptions::seeded(9))?,
model.sample(&probe, 50, &SampleOptions::seeded(9))?
);
assert_eq!(
from_json.sample(&probe, 50, &SampleOptions::seeded(9))?,
model.sample(&probe, 50, &SampleOptions::seeded(9))?
);
println!(
"saved: {} bytes native, {} bytes JSON; reloaded samples match",
model.encode(DiffusionFormat::Binary)?.len(),
model.encode(DiffusionFormat::Json)?.len()
);
Ok(())
}