use rayon::prelude::*;
use super::solver::RidgeSolver;
use super::term_kernel::{TermKernel, TermPart, TermScratch};
use super::{
KernelSolver, NoiseVariance, QUERY_BLOCK, build_solver, check_alpha, check_data, fit_noise,
normal_intervals, z_value,
};
use crate::conformal::Interval;
use crate::data::DMatrix;
use crate::ebm::{EbmInfo, TermShape, term_shape};
use crate::error::{HessboostError, Result};
use crate::model::{BoostedModel, Predictions};
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub struct TermBands {
pub shape: TermShape,
pub standard_errors: Vec<f64>,
pub lower: Vec<f64>,
pub upper: Vec<f64>,
}
struct Stage<'a> {
kernel: TermKernel<'a>,
solver: RidgeSolver,
}
pub struct EbmInference<'a> {
model: &'a BoostedModel,
info: &'a EbmInfo,
stages: Vec<Stage<'a>>,
term_slot: Vec<(usize, usize)>,
c: f64,
s: f64,
n: usize,
noise_variance: f64,
}
impl std::fmt::Debug for EbmInference<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("EbmInference")
.field("rows", &self.n)
.field("terms", &self.info.terms.len())
.field("ridge", &self.c)
.field("scale", &self.s)
.field("noise_variance", &self.noise_variance)
.finish_non_exhaustive()
}
}
impl<'a> EbmInference<'a> {
pub fn fit(
model: &'a BoostedModel,
train: &DMatrix,
noise: NoiseVariance<'a>,
solver: KernelSolver,
) -> Result<Self> {
let not_boulevard = || {
HessboostError::incompatible_model(
"model",
"not a Boulevard EBM: train it with `booster = ebm` and `ebm_boulevard`",
)
};
let info = model.ebm().ok_or_else(not_boulevard)?;
let settings = info.boulevard.ok_or_else(not_boulevard)?;
let noise_variance = fit_noise(model, train, noise)?;
let lambda = settings.learning_rate;
let (c, s) = (1.0 / lambda, (1.0 + lambda) / lambda);
let kappa = settings.reg_lambda / settings.subsample;
let n = train.n_rows();
let mut term_slot = vec![(0, 0); info.terms.len()];
let mut stages = Vec::new();
for stage in info.stages()? {
let parts = stage
.terms
.clone()
.map(|t| TermPart::new(info.term_trees(model, t), &info.terms[t], train, kappa))
.collect::<Result<Vec<_>>>()?;
for (k, t) in stage.terms.enumerate() {
term_slot[t] = (stages.len(), k);
}
let kernel = TermKernel::new(parts, n);
let solver = build_solver(&kernel, solver, c)?;
stages.push(Stage { kernel, solver });
}
Ok(EbmInference {
model,
info,
stages,
term_slot,
c,
s,
n,
noise_variance,
})
}
pub fn noise_variance(&self) -> f64 {
self.noise_variance
}
pub fn intercept_standard_error(&self) -> f64 {
(self.noise_variance / self.n as f64).sqrt()
}
fn through_first_stage(&self, u: &mut [f64], m: usize) {
let first = &self.stages[0];
let n = self.n;
let mut v: Vec<f64> = u
.par_chunks(n)
.flat_map_iter(|row| first.kernel.product(row))
.collect();
first.solver.solve_vectors(&mut v, m, self.c);
for (x, y) in u.iter_mut().zip(v) {
*x -= self.s * y;
}
}
fn norms(
&self,
stage: usize,
m: usize,
rhs: impl Fn(usize, &mut TermScratch, &mut [f64]) + Sync,
) -> Vec<f64> {
let n = self.n;
(0..m.div_ceil(QUERY_BLOCK))
.into_par_iter()
.flat_map_iter(|b| {
let range = b * QUERY_BLOCK..((b + 1) * QUERY_BLOCK).min(m);
let k = range.len();
let mut u = vec![0.0; k * n];
let mut scratch = TermScratch::default();
for (out, a) in u.chunks_exact_mut(n).zip(range) {
rhs(a, &mut scratch, out);
}
self.stages[stage].solver.solve_vectors(&mut u, k, self.c);
if stage > 0 {
self.through_first_stage(&mut u, k);
}
u.chunks_exact(n)
.map(|w| w.iter().map(|v| v * v).sum::<f64>())
.collect::<Vec<_>>()
})
.collect()
}
fn slot(&self, term: usize) -> Result<(usize, usize)> {
self.term_slot.get(term).copied().ok_or_else(|| {
HessboostError::incompatible_model(
"term",
format!("the model has {} terms, got {term}", self.term_slot.len()),
)
})
}
fn cell_standard_errors(&self, term: usize, cells: &[usize]) -> Result<Vec<f64>> {
let (stage, part) = self.slot(term)?;
let kernel = &self.stages[stage].kernel;
let scale = self.s * self.noise_variance.sqrt();
Ok(self
.norms(stage, cells.len(), |a, scratch, out| {
kernel.add_query(part, cells[a], scratch, out);
})
.into_iter()
.map(|w2| scale * w2.max(0.0).sqrt())
.collect())
}
pub fn term_bands(&self, term: usize, alpha: f64) -> Result<TermBands> {
check_alpha(alpha)?;
self.slot(term)?;
let shape = term_shape(self.model, term)?;
let cells: Vec<usize> = (0..shape.values().len()).collect();
let standard_errors = self.cell_standard_errors(term, &cells)?;
let z = z_value(alpha);
let (lower, upper) = shape
.values()
.iter()
.zip(&standard_errors)
.map(|(&v, &se)| (v - z * se, v + z * se))
.unzip();
Ok(TermBands {
shape,
standard_errors,
lower,
upper,
})
}
pub fn term_standard_errors(&self, term: usize, data: &DMatrix) -> Result<Predictions<f64>> {
check_data(self.model, data, "data", false)?;
let (stage, part) = self.slot(term)?;
let part = &self.stages[stage].kernel.parts[part];
let cells: Vec<usize> = (0..data.n_rows()).map(|r| part.cell_of(data, r)).collect();
let se = self.cell_standard_errors(term, &cells)?;
Ok(Predictions::new(se, data.n_rows(), 1))
}
fn prediction_norms(&self, data: &DMatrix) -> Result<Vec<f64>> {
check_data(self.model, data, "data", false)?;
Ok(self.joint_norms(data))
}
fn joint_norms(&self, data: &DMatrix) -> Vec<f64> {
let (n, rows) = (self.n, data.n_rows());
let cells: Vec<Vec<Vec<usize>>> = self
.stages
.iter()
.map(|stage| {
stage
.kernel
.parts
.iter()
.map(|p| (0..rows).map(|r| p.cell_of(data, r)).collect())
.collect()
})
.collect();
(0..rows.div_ceil(QUERY_BLOCK))
.into_par_iter()
.flat_map_iter(|b| {
let range = b * QUERY_BLOCK..((b + 1) * QUERY_BLOCK).min(rows);
let k = range.len();
let mut scratch = TermScratch::default();
let mut sides: Vec<Vec<f64>> = self
.stages
.iter()
.zip(&cells)
.map(|(stage, cells)| {
let mut u = vec![0.0; k * n];
for (out, a) in u.chunks_exact_mut(n).zip(range.clone()) {
for (p, cells) in cells.iter().enumerate() {
stage.kernel.add_query(p, cells[a], &mut scratch, out);
}
}
stage.solver.solve_vectors(&mut u, k, self.c);
u
})
.collect();
let mut w = sides.swap_remove(0);
if let Some(mut second) = sides.pop() {
self.through_first_stage(&mut second, k);
for (a, b) in w.iter_mut().zip(second) {
*a += b;
}
}
w.chunks_exact(n)
.map(|w| {
let w2: f64 = w.iter().map(|v| v * v).sum();
1.0 / n as f64 + self.s * self.s * w2
})
.collect::<Vec<_>>()
})
.collect()
}
pub fn standard_errors(&self, data: &DMatrix) -> Result<Predictions<f64>> {
let sigma = self.noise_variance.sqrt();
let se: Vec<f64> = self
.prediction_norms(data)?
.into_iter()
.map(|w2| sigma * w2.max(0.0).sqrt())
.collect();
Ok(Predictions::new(se, data.n_rows(), 1))
}
fn intervals(
&self,
data: &DMatrix,
alpha: f64,
width: impl Fn(f64) -> f64,
) -> Result<Vec<Interval<f64>>> {
normal_intervals(
self.model,
data,
alpha,
|| self.prediction_norms(data),
|w2| width(w2.max(0.0)),
)
}
pub fn confidence_intervals(&self, data: &DMatrix, alpha: f64) -> Result<Vec<Interval<f64>>> {
let sigma2 = self.noise_variance;
self.intervals(data, alpha, |w2| (sigma2 * w2).sqrt())
}
pub fn prediction_intervals(&self, data: &DMatrix, alpha: f64) -> Result<Vec<Interval<f64>>> {
let sigma2 = self.noise_variance;
self.intervals(data, alpha, |w2| (sigma2 * (1.0 + w2)).sqrt())
}
}