use super::{BoostedModel, ModelObjective, Predictions, transform_margins_in_place};
use crate::data::DMatrix;
use crate::error::{HessboostError, Result};
use crate::objective::Objective;
use crate::objective::distributional::{Dist, DistFamily};
#[derive(Debug, Clone)]
pub struct VirtualEnsembles {
iterations: Vec<usize>,
n_rows: usize,
margin_width: usize,
width: usize,
margins: Vec<f32>,
predictions: Vec<f32>,
}
impl VirtualEnsembles {
pub fn n_members(&self) -> usize {
self.iterations.len()
}
pub fn iterations(&self) -> &[usize] {
&self.iterations
}
pub fn n_rows(&self) -> usize {
self.n_rows
}
pub fn width(&self) -> usize {
self.width
}
pub fn margin_width(&self) -> usize {
self.margin_width
}
pub fn member_predictions(&self, m: usize) -> Option<&[f32]> {
let len = self.n_rows * self.width;
(m < self.n_members()).then(|| &self.predictions[m * len..(m + 1) * len])
}
pub fn member_margins(&self, m: usize) -> Option<&[f32]> {
let len = self.n_rows * self.margin_width;
(m < self.n_members()).then(|| &self.margins[m * len..(m + 1) * len])
}
pub fn into_predictions(self) -> Vec<f32> {
self.predictions
}
pub fn into_margins(self) -> Vec<f32> {
self.margins
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct Uncertainty {
pub mean: Predictions<f64>,
pub knowledge: Predictions<f64>,
pub data: Option<Predictions<f64>>,
pub total: Option<Predictions<f64>>,
}
enum Decomposition {
Regression,
Distributional(DistFamily),
Binary,
Multiclass,
}
impl Decomposition {
fn of(objective: &ModelObjective) -> Result<Self> {
let refused = || {
Err(HessboostError::incompatible_model(
"objective",
format!(
"uncertainty is defined for regression, `dist:*`, and probabilistic \
classification objectives, not `{}`",
objective.name()
),
))
};
let Some(built_in) = objective.built_in() else {
return refused();
};
Ok(match built_in {
Objective::Dist(_) => match built_in.dist_family() {
Some(family) => Decomposition::Distributional(family),
None => return refused(),
},
Objective::BinaryLogistic(_)
| Objective::BinaryLogitRaw(_)
| Objective::RegLogistic(_) => Decomposition::Binary,
Objective::Softmax(_) | Objective::Softprob(_) => Decomposition::Multiclass,
Objective::SquaredError(_)
| Objective::SquaredLogError
| Objective::PseudoHuber(_)
| Objective::AbsoluteError
| Objective::Quantile(_)
| Objective::Expectile(_)
| Objective::Poisson
| Objective::Gamma(_)
| Objective::Tweedie(_)
| Objective::Cox
| Objective::Aft(_) => Decomposition::Regression,
Objective::BinaryHinge
| Objective::RankPairwise(_)
| Objective::RankNdcg(_)
| Objective::RankMap(_)
| Objective::RankXendcg
| Objective::Custom(_) => return refused(),
})
}
}
fn binary_entropy(p: f64) -> f64 {
let term = |q: f64| if q > 0.0 { -q * q.ln() } else { 0.0 };
term(p) + term(1.0 - p)
}
fn mean_variance(values: impl Iterator<Item = f64> + Clone) -> (f64, f64) {
let n = values.clone().count() as f64;
let mean = values.clone().sum::<f64>() / n;
let variance = values.map(|v| (v - mean) * (v - mean)).sum::<f64>() / n;
(mean, variance)
}
impl BoostedModel {
fn virtual_ensemble_iterations(&self, count: usize) -> Result<Vec<usize>> {
if self.linear.is_some() {
return Err(HessboostError::incompatible_model(
"model",
"gblinear models have no boosting iterations to form virtual ensembles from",
));
}
if self.ebm.is_some() {
return Err(HessboostError::incompatible_model(
"model",
"an EBM's trees are ordered by bag, stage, and term, so its iteration prefixes \
are not the models of fewer rounds",
));
}
if self.boulevard.is_some() {
return Err(HessboostError::incompatible_model(
"model",
"a Boulevard model's leaves carry its average over every round, so its \
iteration prefixes are not the models of fewer rounds; its variance comes from \
`inference::BoulevardInference`",
));
}
if count == 0 {
return Err(HessboostError::invalid_param(
"virtual_ensembles_count",
"must be at least 1",
));
}
let end = self.effective_num_trees() / self.trees_per_iteration();
let period = end / count.saturating_mul(2);
if period == 0 || period * count >= end {
return Err(HessboostError::incompatible_model(
"virtual_ensembles_count",
format!(
"{count} virtual ensembles need a model of at least {} iterations, this one \
has {end}",
count.saturating_mul(2)
),
));
}
let begin = end - period * count;
Ok((1..=count).map(|m| begin + m * period).collect())
}
pub fn predict_virtual_ensembles(
&self,
data: &DMatrix,
count: usize,
) -> Result<VirtualEnsembles> {
let iterations = self.virtual_ensemble_iterations(count)?;
let margins = self.prefix_margins(data, &iterations)?;
let (n_rows, k) = (data.n_rows(), self.n_outputs);
let width = if self
.objective
.built_in()
.is_some_and(Objective::predicts_class_index)
{
1
} else {
k
};
let mut predictions = Vec::with_capacity(iterations.len() * n_rows * width);
let len = n_rows * k;
for m in 0..iterations.len() {
let start = predictions.len();
predictions.extend_from_slice(&margins[m * len..(m + 1) * len]);
let values = &mut predictions[start..];
transform_margins_in_place(
&self.objective,
self.max_delta_step,
self.n_targets,
values,
k,
);
predictions.truncate(start + n_rows * width);
}
Ok(VirtualEnsembles {
iterations,
n_rows,
margin_width: k,
width,
margins,
predictions,
})
}
pub fn predict_uncertainty(&self, data: &DMatrix, count: usize) -> Result<Uncertainty> {
let decomposition = Decomposition::of(&self.objective)?;
let ensembles = self.predict_virtual_ensembles(data, count)?;
let n = ensembles.n_rows;
let k = self.n_outputs;
let members = || 0..ensembles.n_members();
let margin = |m: usize, row: usize, out: usize| {
f64::from(ensembles.margins[(m * n + row) * k + out])
};
Ok(match decomposition {
Decomposition::Regression => {
let width = ensembles.width;
let mut mean = Vec::with_capacity(n * width);
let mut knowledge = Vec::with_capacity(n * width);
for cell in 0..n * width {
let values =
members().map(|m| f64::from(ensembles.predictions[m * n * width + cell]));
let (mu, variance) = mean_variance(values);
mean.push(mu);
knowledge.push(variance);
}
Uncertainty {
mean: Predictions::new(mean, n, width),
knowledge: Predictions::new(knowledge, n, width),
data: None,
total: None,
}
}
Decomposition::Distributional(family) => {
let mut mean = Vec::with_capacity(n);
let mut knowledge = Vec::with_capacity(n);
let mut aleatoric = Vec::with_capacity(n);
for row in 0..n {
let dists: Vec<_> = members()
.map(|m| {
let eta: Vec<f64> = (0..k).map(|out| margin(m, row, out)).collect();
family.dist_from_margins(&eta)
})
.collect();
let (mu, variance) = mean_variance(dists.iter().map(Dist::mean));
mean.push(mu);
knowledge.push(variance);
aleatoric.push(dists.iter().map(Dist::variance).sum::<f64>() / count as f64);
}
let total = knowledge
.iter()
.zip(&aleatoric)
.map(|(a, b)| a + b)
.collect();
Uncertainty {
mean: Predictions::new(mean, n, 1),
knowledge: Predictions::new(knowledge, n, 1),
data: Some(Predictions::new(aleatoric, n, 1)),
total: Some(Predictions::new(total, n, 1)),
}
}
Decomposition::Binary => {
let mut mean = Vec::with_capacity(n * k);
let mut aleatoric = Vec::with_capacity(n * k);
let mut total = Vec::with_capacity(n * k);
for row in 0..n {
for out in 0..k {
let probs = members().map(|m| 1.0 / (1.0 + (-margin(m, row, out)).exp()));
let p = probs.clone().sum::<f64>() / count as f64;
mean.push(p);
aleatoric.push(probs.map(binary_entropy).sum::<f64>() / count as f64);
total.push(binary_entropy(p));
}
}
let knowledge = total.iter().zip(&aleatoric).map(|(t, d)| t - d).collect();
Uncertainty {
mean: Predictions::new(mean, n, k),
knowledge: Predictions::new(knowledge, n, k),
data: Some(Predictions::new(aleatoric, n, k)),
total: Some(Predictions::new(total, n, k)),
}
}
Decomposition::Multiclass => {
let mut mean = vec![0.0; n * k];
let mut aleatoric = Vec::with_capacity(n);
let mut total = Vec::with_capacity(n);
let mut probs = vec![0.0; k];
for row in 0..n {
let row_mean = &mut mean[row * k..(row + 1) * k];
let mut entropy_sum = 0.0;
for m in members() {
let max = (0..k)
.map(|c| margin(m, row, c))
.fold(f64::NEG_INFINITY, f64::max);
for (c, p) in probs.iter_mut().enumerate() {
*p = (margin(m, row, c) - max).exp();
}
let sum: f64 = probs.iter().sum();
for (p, acc) in probs.iter_mut().zip(row_mean.iter_mut()) {
*p /= sum;
*acc += *p;
if *p > 0.0 {
entropy_sum -= *p * p.ln();
}
}
}
let mut entropy_of_mean = 0.0;
for p in row_mean.iter_mut() {
*p /= count as f64;
if *p > 0.0 {
entropy_of_mean -= *p * p.ln();
}
}
aleatoric.push(entropy_sum / count as f64);
total.push(entropy_of_mean);
}
let knowledge = total.iter().zip(&aleatoric).map(|(t, d)| t - d).collect();
Uncertainty {
mean: Predictions::new(mean, n, k),
knowledge: Predictions::new(knowledge, n, 1),
data: Some(Predictions::new(aleatoric, n, 1)),
total: Some(Predictions::new(total, n, 1)),
}
}
})
}
}