use super::{Metric, weighted_mean};
use crate::data::MetaInfo;
fn alpha_average(
alpha: &[f32],
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
n_rows: usize,
loss: impl Fn(f32, f32, f32) -> f32,
) -> f64 {
if n_rows == 0
|| labels.is_empty()
|| !labels.len().is_multiple_of(n_rows)
|| labels.len().checked_mul(alpha.len()) != Some(preds.len())
|| weights.is_some_and(|w| w.len() != n_rows)
{
return f64::NAN;
}
let n_targets = labels.len() / n_rows;
let (mut total, mut weight) = (0.0f64, 0.0f64);
for (i, row) in preds.chunks_exact(alpha.len() * n_targets).enumerate() {
let w = weights.map_or(1.0, |ws| ws[i]);
let y_row = &labels[i * n_targets..(i + 1) * n_targets];
for (&a, cells) in alpha.iter().zip(row.chunks_exact(n_targets)) {
for (&p, &y) in cells.iter().zip(y_row) {
total += f64::from(loss(a, p, y) * w);
weight += f64::from(w);
}
}
}
weighted_mean((total, weight))
}
#[derive(Debug, Clone)]
pub(crate) struct QuantileError {
alpha: Vec<f32>,
}
impl QuantileError {
pub(crate) fn new(alpha: Vec<f32>) -> Self {
QuantileError { alpha }
}
}
impl Metric for QuantileError {
fn name(&self) -> &'static str {
"quantile"
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
self.eval_info(preds, &MetaInfo::new(labels, weights, None))
}
fn eval_info(&self, preds: &[f32], info: &MetaInfo) -> f64 {
alpha_average(
&self.alpha,
preds,
info.label_values(),
info.weights,
info.n_rows,
|a, p, y| {
let d = y - p;
let sign = f32::from(u8::from(d >= 0.0));
(a * sign * d) - (1.0 - a) * (1.0 - sign) * d
},
)
}
fn prediction_width(&self, info: &MetaInfo) -> Option<usize> {
Some(self.alpha.len() * info.n_targets())
}
}
#[derive(Debug, Clone)]
pub(crate) struct ExpectileError {
alpha: Vec<f32>,
}
impl ExpectileError {
pub(crate) fn new(alpha: Vec<f32>) -> Self {
ExpectileError { alpha }
}
}
impl Metric for ExpectileError {
fn name(&self) -> &'static str {
"expectile"
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
self.eval_info(preds, &MetaInfo::new(labels, weights, None))
}
fn eval_info(&self, preds: &[f32], info: &MetaInfo) -> f64 {
alpha_average(
&self.alpha,
preds,
info.label_values(),
info.weights,
info.n_rows,
|a, p, y| {
let diff = p - y;
let scale = if diff >= 0.0 { 1.0 - a } else { a };
scale * diff * diff
},
)
}
fn prediction_width(&self, info: &MetaInfo) -> Option<usize> {
Some(self.alpha.len() * info.n_targets())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn quantile_averages_pinball_over_alphas() {
let m = QuantileError::new(vec![0.25, 0.75]);
let preds = [0.0, 2.0, 0.0, 0.0];
assert!((m.eval(&preds, &[1.0, 0.0], None) - 0.5 / 4.0).abs() < 1e-12);
let v = m.eval(&preds, &[1.0, 0.0], Some(&[3.0, 1.0]));
assert!((v - 1.5 / 8.0).abs() < 1e-12, "{v}");
assert!(m.eval(&preds[..2], &[1.0, 0.0], None).is_nan());
}
#[test]
fn expectile_weights_residual_sides() {
let m = ExpectileError::new(vec![0.2]);
let v = m.eval(&[2.0, -1.0], &[0.0, 0.0], None);
assert!((v - 1.7).abs() < 1e-6, "{v}");
}
#[test]
fn empty_labels_evaluate_to_nan() {
let info = MetaInfo {
n_rows: 3,
weights: None,
..MetaInfo::unlabeled(0)
};
let q = QuantileError::new(vec![0.5]);
let e = ExpectileError::new(vec![0.2, 0.8]);
assert!(q.eval_info(&[], &info).is_nan());
assert!(e.eval_info(&[], &info).is_nan());
}
}