hessboost 0.2.2

Fast, deterministic gradient boosting (GBDT) in Rust: conformal intervals, explainable boosting machines, distributional boosting, tree-based diffusion, and XGBoost model interchange
Documentation
//! Metrics of the distributional `dist:*` objectives: the
//! mean negative log-likelihood (`nll`) and the mean continuous ranked
//! probability score (`crps`) of the predicted distributions.
//!
//! Both read predictions as the objective reports them, one row of natural
//! parameters per instance (`[row][parameter]`), of the family they carry
//! (`EvalMetric::Nll` / `EvalMetric::Crps`; XGBoost's flat form takes it
//! from the `dist:*` objective).

use super::{Metric, weighted_mean};
use crate::data::MetaInfo;
use crate::objective::distributional::{Dist, DistFamily};
use rayon::prelude::*;

/// Weighted mean of `score(dist_i, y_i)` over the rows. Per-row scores are
/// computed in parallel and summed sequentially in row order, so the value
/// does not depend on the thread count. Zero-weight rows are skipped, so a
/// row whose score is infinite or NaN (e.g. a label outside the predicted
/// support) cannot turn the mean into `0 · ∞ = NaN`. NaN unless `preds`
/// holds one row of parameters per label and `weights` one weight per
/// label.
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))
}

/// Mean negative log-likelihood `-ln p(y)` of the predicted distributions
/// (`nll`, the `dist:*` objectives' default metric): the log density for the
/// continuous families, the log probability mass for the count families.
#[derive(Debug, Clone, Copy)]
pub(crate) struct DistNll {
    family: DistFamily,
}

impl DistNll {
    /// The metric for `family`.
    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))
    }

    /// One row of the family's natural parameters per label.
    fn prediction_width(&self, _info: &MetaInfo) -> Option<usize> {
        Some(self.family.n_params())
    }

    fn supports_label_matrix(&self) -> bool {
        false
    }
}

/// Mean continuous ranked probability score of the predicted distributions
/// (`crps`), in the label's units; see [`Dist::crps`] for the closed forms
/// and the exact step sums of the count families.
#[derive(Debug, Clone, Copy)]
pub(crate) struct DistCrps {
    family: DistFamily,
}

impl DistCrps {
    /// The metric for `family`.
    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)
    }

    /// One row of the family's natural parameters per label.
    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() {
        // Two Normal rows: N(0, 1) at y = 0 and N(1, 2) at y = 3.
        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);
    }

    /// A zero-weight row whose score overflows (LogNormal with `sigma =
    /// 100`: `e^{sigma²/2} = ∞` in its CRPS) leaves the mean of the others.
    #[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());
            // Not a distributional objective: nothing to score.
            assert!(named(name, 2, &DEFAULT_SOURCE).is_err());
        }
    }
}