use super::{CellMetric, Metric, cells_consistent, weighted_mean};
use crate::simd::RowWeights;
fn weighted_sum(
preds: &[f32],
labels: &[f32],
weights: Option<RowWeights<'_>>,
loss: impl Fn(f32, f32) -> f32,
) -> (f64, f64) {
if !cells_consistent(preds, labels, weights) {
return (f64::NAN, 1.0);
}
let mut total = 0.0f64;
let mut weight = 0.0f64;
for (i, (&p, &y)) in preds.iter().zip(labels).enumerate() {
let w = weights.map_or(1.0, |ws| ws.get(i));
total += f64::from(loss(y, p) * w);
weight += f64::from(w);
}
(total, weight)
}
#[derive(Debug, Clone, Copy, Default)]
#[non_exhaustive]
pub(crate) struct Rmsle;
impl Metric for Rmsle {
fn name(&self) -> &'static str {
"rmsle"
}
cell_metric_eval!();
}
impl CellMetric for Rmsle {
fn eval_cells(&self, preds: &[f32], labels: &[f32], weights: Option<RowWeights<'_>>) -> f64 {
weighted_mean(weighted_sum(preds, labels, weights, |y, p| {
let diff = y.ln_1p() - p.ln_1p();
diff * diff
}))
.sqrt()
}
}
#[derive(Debug, Clone, Copy, Default)]
#[non_exhaustive]
pub(crate) struct Mape;
impl Metric for Mape {
fn name(&self) -> &'static str {
"mape"
}
cell_metric_eval!();
}
impl CellMetric for Mape {
fn eval_cells(&self, preds: &[f32], labels: &[f32], weights: Option<RowWeights<'_>>) -> f64 {
weighted_mean(weighted_sum(preds, labels, weights, |y, p| {
((y - p) / y).abs()
}))
}
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct PseudoHuberError {
slope: f32,
}
impl PseudoHuberError {
pub(super) fn new(slope: f32) -> Self {
PseudoHuberError { slope }
}
}
impl Metric for PseudoHuberError {
fn name(&self) -> &'static str {
"mphe"
}
cell_metric_eval!();
}
impl CellMetric for PseudoHuberError {
fn eval_cells(&self, preds: &[f32], labels: &[f32], weights: Option<RowWeights<'_>>) -> f64 {
let slope = self.slope;
weighted_mean(weighted_sum(preds, labels, weights, |y, p| {
let scaled = (y - p) / slope;
slope * slope * ((1.0 + scaled * scaled).sqrt() - 1.0)
}))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rmsle_is_rmse_in_log1p_space() {
let e1 = std::f32::consts::E - 1.0;
let v = Rmsle.eval(&[e1, 0.0], &[0.0, e1], None);
assert!((v - 1.0).abs() < 1e-6, "{v}");
assert_eq!(Rmsle.eval(&[3.0, 5.0], &[3.0, 5.0], None), 0.0);
}
#[test]
fn mape_is_relative_to_label_and_weighted() {
let v = Mape.eval(&[1.0, 5.0], &[2.0, 4.0], Some(&[3.0, 1.0]));
assert!((v - (0.5 * 3.0 + 0.25) / 4.0).abs() < 1e-7, "{v}");
}
#[test]
fn mphe_uses_slope_without_factor_two() {
let v = PseudoHuberError::new(2.0).eval(&[0.0], &[1.5], None);
assert!((v - 1.0).abs() < 1e-6, "{v}");
let unit = PseudoHuberError::new(1.0).eval(&[0.0], &[0.75], None);
assert!((unit - 0.25).abs() < 1e-6, "{unit}");
}
}