use hessboost::conformal::{ConformalizedQuantile, Interval, SplitConformal};
use hessboost::objective::{CustomLoss, GradPair};
use hessboost::prelude::*;
mod common;
use common::{lcg, normal};
fn dataset(n: usize, seed: u64) -> Result<DMatrix> {
let mut next = lcg(seed);
let mut x = Vec::with_capacity(2 * n);
let mut y = Vec::with_capacity(n);
for _ in 0..n {
let (x0, x1) = (next(), next());
let eps = normal(&mut next);
x.extend_from_slice(&[x0, x1]);
y.push((std::f32::consts::TAU * x0).sin() + (0.1 + x1) * eps);
}
DMatrix::from_dense(&x, n, 2)?.with_labels(&y)
}
fn summarize(intervals: &[Interval], data: &DMatrix) -> (f64, f64) {
common::coverage(
intervals
.iter()
.map(|i| (f64::from(i.lower), f64::from(i.upper))),
data.labels().unwrap_or_default(),
)
}
fn main() -> Result<()> {
let alpha = 0.1;
let dtrain = dataset(4000, 1)?;
let dcal = dataset(1000, 2)?;
let dtest = dataset(5000, 3)?;
let params = TrainingParams::builder().max_depth(4).eta(0.3).build()?;
let point = train(¶ms, &dtrain, 50)?;
let split = SplitConformal::calibrate(&point, &dcal, alpha)?;
let split_iv = split.predict_interval(&dtest)?;
let (cov, width) = summarize(&split_iv, &dtest);
println!(
"split conformal: half-width Q = {:.3}, coverage {cov:.3}, mean width {width:.3}",
split.half_width()
);
let taus = [(alpha / 2.0) as f32, (1.0 - alpha / 2.0) as f32];
let pinball = CustomLoss::new("pinball", 2, move |p, y, _w, out| {
for (i, &yi) in y.iter().enumerate() {
for (j, tau) in taus.iter().enumerate() {
let g = if p[2 * i + j] > yi { 1.0 - tau } else { -tau };
out[2 * i + j] = GradPair::new(g, 1.0);
}
}
})
.with_default_metric(EvalMetric::Mae);
let mut quantile_params = params;
quantile_params.objective = Objective::custom(pinball);
let quantiles = train(&quantile_params, &dtrain, 200)?;
let preds = quantiles.predict(&dtest, Iterations::Best)?;
let band: Vec<Interval> = preds
.rows()
.map(|row| Interval {
lower: row[0],
upper: row[1],
})
.collect();
let (cov, width) = summarize(&band, &dtest);
println!("raw quantiles: coverage {cov:.3}, mean width {width:.3}");
let cqr = ConformalizedQuantile::calibrate_outputs(&quantiles, 0, 1, &dcal, alpha)?;
let cqr_iv = cqr.predict_interval(&dtest)?;
let (cov, width) = summarize(&cqr_iv, &dtest);
println!(
"CQR: correction Q = {:+.3}, coverage {cov:.3}, mean width {width:.3}",
cqr.correction()
);
let probe = DMatrix::from_dense(&[0.25, 0.05, 0.25, 0.5, 0.25, 0.95], 3, 2)?;
let split_probe = split.predict_interval(&probe)?;
let cqr_probe = cqr.predict_interval(&probe)?;
println!("\n x1 noise sd split width CQR width");
for (i, x1) in [0.05f32, 0.5, 0.95].into_iter().enumerate() {
println!(
"{x1:5.2} {:8.2} {:11.3} {:9.3}",
0.1 + x1,
split_probe[i].upper - split_probe[i].lower,
cqr_probe[i].upper - cqr_probe[i].lower
);
}
let n_cal = dcal.n_rows() as f64;
println!(
"\nBoth guarantee marginal coverage in [{:.2}, {:.4}] (no ties) averaged over \
calibration sets;\na single calibration set of {n_cal} rows fluctuates by about ±{:.3}.",
1.0 - alpha,
1.0 - alpha + 1.0 / (n_cal + 1.0),
(alpha * (1.0 - alpha) / n_cal).sqrt()
);
Ok(())
}