use super::{Metric, weighted_mean};
use crate::data::MetaInfo;
use crate::objective::distributional::{Dist, DistFamily};
use rayon::prelude::*;
fn mean_score(
family: DistFamily,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
score: impl Fn(&Dist, f64) -> f64 + Sync,
) -> f64 {
let k = family.n_params();
nan_unless_consistent!(preds, labels, weights, k);
let weight = |i: usize| weights.map_or(1.0, |ws| f64::from(ws[i]));
let scores: Vec<f64> = preds
.par_chunks_exact(k)
.zip(labels.par_iter())
.enumerate()
.map(|(i, (row, &y))| {
if weight(i) == 0.0 {
0.0
} else {
score(&Dist::from_row(family, row), f64::from(y))
}
})
.collect();
let (mut total, mut weight_sum) = (0.0f64, 0.0f64);
for (i, s) in scores.into_iter().enumerate() {
let w = weight(i);
if w != 0.0 {
total += w * s;
weight_sum += w;
}
}
weighted_mean((total, weight_sum))
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct DistNll {
family: DistFamily,
}
impl DistNll {
pub(crate) fn new(family: DistFamily) -> Self {
DistNll { family }
}
}
impl Metric for DistNll {
fn name(&self) -> &'static str {
"nll"
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
mean_score(self.family, preds, labels, weights, |d, y| -d.log_prob(y))
}
fn prediction_width(&self, _info: &MetaInfo) -> Option<usize> {
Some(self.family.n_params())
}
fn supports_label_matrix(&self) -> bool {
false
}
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct DistCrps {
family: DistFamily,
}
impl DistCrps {
pub(crate) fn new(family: DistFamily) -> Self {
DistCrps { family }
}
}
impl Metric for DistCrps {
fn name(&self) -> &'static str {
"crps"
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
mean_score(self.family, preds, labels, weights, Dist::crps)
}
fn prediction_width(&self, _info: &MetaInfo) -> Option<usize> {
Some(self.family.n_params())
}
fn supports_label_matrix(&self) -> bool {
false
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::metric::{DEFAULT_SOURCE, XgboostMetricSource, named};
#[test]
fn metrics_average_the_per_row_scores_with_weights() {
let preds = [0.0f32, 1.0, 1.0, 2.0];
let labels = [0.0f32, 3.0];
let d0 = Dist::Normal {
mu: 0.0,
sigma: 1.0,
};
let d1 = Dist::Normal {
mu: 1.0,
sigma: 2.0,
};
let nll = DistNll::new(DistFamily::Normal);
let crps = DistCrps::new(DistFamily::Normal);
let expect = (-d0.log_prob(0.0) - 3.0 * d1.log_prob(3.0)) / 4.0;
assert!((nll.eval(&preds, &labels, Some(&[1.0, 3.0])) - expect).abs() < 1e-12);
let expect = f64::midpoint(d0.crps(0.0), d1.crps(3.0));
assert!((crps.eval(&preds, &labels, None) - expect).abs() < 1e-12);
}
#[test]
fn zero_weight_rows_do_not_poison_the_mean() {
let preds = [0.0f32, 100.0, 0.0, 1.0];
let labels = [1.0f32, 1.0];
let weights = [0.0f32, 1.0];
let d1 = Dist::LogNormal {
mu: 0.0,
sigma: 1.0,
};
let crps = DistCrps::new(DistFamily::LogNormal);
assert!(!crps.eval(&preds[..2], &labels[..1], None).is_finite());
assert_eq!(crps.eval(&preds, &labels, Some(&weights)), d1.crps(1.0));
let nll = DistNll::new(DistFamily::LogNormal);
assert_eq!(nll.eval(&preds, &labels, Some(&weights)), -d1.log_prob(1.0));
}
#[test]
fn factory_takes_the_family_from_the_objective() {
let dist = XgboostMetricSource {
distribution: Some(DistFamily::Gamma),
..DEFAULT_SOURCE
};
for name in ["nll", "crps"] {
let metric = named(name, 2, &dist).unwrap();
assert_eq!(metric.name(), name);
assert!(!metric.maximize());
assert!(named(name, 2, &DEFAULT_SOURCE).is_err());
}
}
}