use crate::error::{HessboostError, Result};
pub trait Metric: Send + Sync {
fn name(&self) -> &str;
fn maximize(&self) -> bool {
false
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64;
fn eval_grouped(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
_group: Option<&crate::data::GroupInfo>,
) -> f64 {
self.eval(preds, labels, weights)
}
}
#[inline]
fn weighted_mean((total, weight): (f64, f64)) -> f64 {
if weight > 0.0 { total / weight } else { 0.0 }
}
macro_rules! simple_metric {
($(#[$m:meta])* $ty:ident, $name:literal, $simd:path) => {
$(#[$m])*
#[derive(Debug, Clone, Copy, Default)]
pub struct $ty;
impl Metric for $ty {
fn name(&self) -> &str {
$name
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
weighted_mean($simd(preds, labels, weights))
}
}
};
($(#[$m:meta])* $ty:ident, $name:literal, $simd:path => sqrt) => {
$(#[$m])*
#[derive(Debug, Clone, Copy, Default)]
pub struct $ty;
impl Metric for $ty {
fn name(&self) -> &str {
$name
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
weighted_mean($simd(preds, labels, weights)).sqrt()
}
}
};
($(#[$m:meta])* $ty:ident, $name:literal, $field:ident: $field_ty:ty, $simd:path) => {
$(#[$m])*
#[derive(Debug, Clone, Copy)]
pub struct $ty {
$field: $field_ty,
}
impl Metric for $ty {
fn name(&self) -> &str {
$name
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
weighted_mean($simd(preds, labels, weights, self.$field))
}
}
};
}
simple_metric!(
Rmse, "rmse", crate::simd::squared_error_sum => sqrt
);
simple_metric!(
Mae, "mae", crate::simd::absolute_error_sum
);
simple_metric!(
LogLoss, "logloss", crate::simd::log_loss_sum
);
simple_metric!(
ErrorRate, "error", crate::simd::classification_error_sum
);
fn tie_runs<'a>(
order: &'a [usize],
preds: &'a [f32],
) -> impl Iterator<Item = std::ops::Range<usize>> + 'a {
let mut start = 0;
std::iter::from_fn(move || {
if start >= order.len() {
return None;
}
let mut end = start + 1;
while end < order.len() && preds[order[end]] == preds[order[start]] {
end += 1;
}
let run = start..end;
start = end;
Some(run)
})
}
#[derive(Debug, Clone, Copy, Default)]
pub struct Auc;
impl Metric for Auc {
fn name(&self) -> &'static str {
"auc"
}
fn maximize(&self) -> bool {
true
}
fn eval(&self, preds: &[f32], labels: &[f32], _weights: Option<&[f32]>) -> f64 {
let n = preds.len();
let mut order: Vec<usize> = (0..n).collect();
order.sort_by(|&a, &b| preds[a].total_cmp(&preds[b]));
let mut ranks = vec![0.0f64; n];
for run in tie_runs(&order, preds) {
let avg = ((run.start + 1 + run.end) as f64) / 2.0; for &idx in &order[run] {
ranks[idx] = avg;
}
}
let mut sum_pos_rank = 0.0f64;
let mut n_pos = 0.0f64;
let mut n_neg = 0.0f64;
for k in 0..n {
if labels[k] > 0.5 {
sum_pos_rank += ranks[k];
n_pos += 1.0;
} else {
n_neg += 1.0;
}
}
if n_pos == 0.0 || n_neg == 0.0 {
return 0.5; }
(sum_pos_rank - n_pos * (n_pos + 1.0) / 2.0) / (n_pos * n_neg)
}
}
simple_metric!(
MLogLoss, "mlogloss", num_class: usize, crate::simd::multiclass_log_loss_sum
);
simple_metric!(
MError, "merror", num_class: usize, crate::simd::multiclass_error_sum
);
simple_metric!(
PoissonNLogLik, "poisson-nloglik", crate::simd::positive_nloglik_sum::<false>
);
simple_metric!(
GammaNLogLik, "gamma-nloglik", crate::simd::positive_nloglik_sum::<true>
);
simple_metric!(
TweedieNLogLik, "tweedie-nloglik", rho: f64, crate::simd::tweedie_nloglik_sum
);
pub(crate) fn group_ranges(
n: usize,
group: Option<&crate::data::GroupInfo>,
) -> Vec<(usize, usize)> {
match group {
Some(g) if g.num_rows() == n => g.iter_ranges().collect(),
_ => vec![(0, n)],
}
}
pub(crate) fn argsort_desc(values: &[f32]) -> Vec<usize> {
let mut order: Vec<usize> = (0..values.len()).collect();
order.sort_by(|&a, &b| {
values[b]
.partial_cmp(&values[a])
.unwrap_or_else(|| values[b].total_cmp(&values[a]))
});
order
}
fn grouped_average(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
group: Option<&crate::data::GroupInfo>,
mut score: impl FnMut(&[f32], &[f32]) -> f64,
) -> f64 {
let ranges = group_ranges(preds.len(), group);
if ranges.is_empty() {
return 0.0;
}
let mut sum = 0.0;
let mut weight_sum = 0.0;
for &(start, end) in &ranges {
let weight = weights.map_or(1.0, |values| f64::from(values[start]));
sum += weight * score(&preds[start..end], &labels[start..end]);
weight_sum += weight;
}
weighted_mean((sum, weight_sum))
}
#[derive(Debug, Clone, Copy, Default)]
pub struct Ndcg {
k: Option<usize>,
}
impl Ndcg {
pub fn new(k: Option<usize>) -> Self {
Ndcg { k }
}
fn group_ndcg(&self, preds: &[f32], labels: &[f32]) -> f64 {
let m = preds.len();
let cut = self.k.map_or(m, |k| k.min(m));
let order = argsort_desc(preds);
let dcg: f64 = order[..cut]
.iter()
.enumerate()
.map(|(p, &i)| ndcg_gain(f64::from(labels[i])) * ndcg_discount(p))
.sum();
let labels_f64: Vec<f64> = labels.iter().map(|&l| f64::from(l)).collect();
let idcg = ideal_dcg(&labels_f64, cut);
if idcg <= 0.0 { 0.0 } else { dcg / idcg }
}
}
impl Metric for Ndcg {
fn name(&self) -> &'static str {
"ndcg"
}
fn maximize(&self) -> bool {
true
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
self.eval_grouped(preds, labels, weights, None)
}
fn eval_grouped(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
group: Option<&crate::data::GroupInfo>,
) -> f64 {
grouped_average(preds, labels, weights, group, |p, l| self.group_ndcg(p, l))
}
}
#[inline]
pub(crate) fn ndcg_gain(rel: f64) -> f64 {
(2.0f64).powf(rel) - 1.0
}
#[inline]
pub(crate) fn ndcg_discount(p: usize) -> f64 {
1.0 / ((p + 2) as f64).log2()
}
pub(crate) fn ideal_dcg(labels: &[f64], cut: usize) -> f64 {
let mut ideal: Vec<f64> = labels.to_vec();
ideal.sort_by(|a, b| b.total_cmp(a));
ideal[..cut]
.iter()
.enumerate()
.map(|(p, &l)| ndcg_gain(l) * ndcg_discount(p))
.sum()
}
#[derive(Debug, Clone, Copy, Default)]
pub struct MeanAveragePrecision {
k: Option<usize>,
}
impl MeanAveragePrecision {
pub fn new(k: Option<usize>) -> Self {
MeanAveragePrecision { k }
}
fn group_ap(&self, preds: &[f32], labels: &[f32]) -> f64 {
let m = preds.len();
let cut = self.k.map_or(m, |k| k.min(m));
let order = argsort_desc(preds);
let num_rel = labels.iter().filter(|&&l| l > 0.0).count();
if num_rel == 0 {
return 0.0;
}
let mut hits = 0usize;
let mut ap = 0.0f64;
for (p, &i) in order[..cut].iter().enumerate() {
if labels[i] > 0.0 {
hits += 1;
ap += hits as f64 / (p + 1) as f64;
}
}
ap / num_rel as f64
}
}
impl Metric for MeanAveragePrecision {
fn name(&self) -> &'static str {
"map"
}
fn maximize(&self) -> bool {
true
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
self.eval_grouped(preds, labels, weights, None)
}
fn eval_grouped(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
group: Option<&crate::data::GroupInfo>,
) -> f64 {
grouped_average(preds, labels, weights, group, |p, l| self.group_ap(p, l))
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct AucPr;
impl Metric for AucPr {
fn name(&self) -> &'static str {
"aucpr"
}
fn maximize(&self) -> bool {
true
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
let w_of = |i: usize| weights.map_or(1.0, |ws| f64::from(ws[i]));
let order = argsort_desc(preds);
let mut total_pos = 0.0f64;
let mut total_neg = 0.0f64;
for (i, &label) in labels.iter().enumerate() {
if label > 0.5 {
total_pos += w_of(i);
} else {
total_neg += w_of(i);
}
}
if total_pos <= 0.0 || total_neg <= 0.0 {
return 0.0;
}
let mut area = 0.0f64;
let (mut tp, mut fp) = (0.0f64, 0.0f64);
let (mut tp_prev, mut fp_prev) = (0.0f64, 0.0f64);
for run in tie_runs(&order, preds) {
for &idx in &order[run] {
if labels[idx] > 0.5 {
tp += w_of(idx);
} else {
fp += w_of(idx);
}
}
if tp + fp > 0.0 {
let recall = tp / total_pos;
let recall_prev = tp_prev / total_pos;
let prec = tp / (tp + fp);
let prec_prev = if tp_prev + fp_prev > 0.0 {
tp_prev / (tp_prev + fp_prev)
} else {
prec
};
area += (recall - recall_prev) * (prec + prec_prev) / 2.0;
}
tp_prev = tp;
fp_prev = fp;
}
area
}
}
type MetricFn = dyn Fn(&[f32], &[f32], Option<&[f32]>) -> f64 + Send + Sync;
pub struct CustomMetric {
name: String,
maximize: bool,
f: Box<MetricFn>,
}
impl CustomMetric {
pub fn new(
name: impl Into<String>,
maximize: bool,
f: impl Fn(&[f32], &[f32], Option<&[f32]>) -> f64 + Send + Sync + 'static,
) -> Self {
CustomMetric {
name: name.into(),
maximize,
f: Box::new(f),
}
}
}
impl Metric for CustomMetric {
fn name(&self) -> &str {
&self.name
}
fn maximize(&self) -> bool {
self.maximize
}
fn eval(&self, preds: &[f32], labels: &[f32], weights: Option<&[f32]>) -> f64 {
(self.f)(preds, labels, weights)
}
}
pub fn create_metric(name: &str, num_class: usize) -> Result<Box<dyn Metric>> {
let (base, rho) = match name.split_once('@') {
Some((b, r)) => (b, r.parse::<f64>().ok()),
None => (name, None),
};
match base {
"rmse" => Ok(Box::new(Rmse)),
"mae" => Ok(Box::new(Mae)),
"logloss" => Ok(Box::new(LogLoss)),
"error" => Ok(Box::new(ErrorRate)),
"auc" => Ok(Box::new(Auc)),
"aucpr" => Ok(Box::new(AucPr)),
"mlogloss" => Ok(Box::new(MLogLoss {
num_class: num_class.max(2),
})),
"merror" => Ok(Box::new(MError {
num_class: num_class.max(2),
})),
"poisson-nloglik" => Ok(Box::new(PoissonNLogLik)),
"gamma-nloglik" => Ok(Box::new(GammaNLogLik)),
"tweedie-nloglik" => Ok(Box::new(TweedieNLogLik {
rho: rho.unwrap_or(1.5),
})),
"ndcg" => Ok(Box::new(Ndcg::new(rho.map(|r| r as usize)))),
"map" => Ok(Box::new(MeanAveragePrecision::new(rho.map(|r| r as usize)))),
other => Err(HessboostError::unknown("metric", other)),
}
}
pub fn create_metrics(
eval_metric: &[String],
default_name: &str,
num_class: usize,
) -> Result<Vec<Box<dyn Metric>>> {
if eval_metric.is_empty() {
Ok(vec![create_metric(default_name, num_class)?])
} else {
eval_metric
.iter()
.map(|n| create_metric(n, num_class))
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
#[test]
fn rmse_basic() {
let m = Rmse;
assert_relative_eq!(m.eval(&[2.0, 0.0], &[1.0, 1.0], None), 1.0, epsilon = 1e-9);
}
#[test]
fn logloss_perfect_and_wrong() {
let m = LogLoss;
let loss = m.eval(&[0.999_999, 0.000_001], &[1.0, 0.0], None);
assert!(loss < 1e-4);
let loss = m.eval(&[0.5, 0.5], &[1.0, 0.0], None);
assert_relative_eq!(loss, 2.0f64.ln(), epsilon = 1e-6);
}
#[test]
fn error_rate_counts_misclassified() {
let m = ErrorRate;
assert_relative_eq!(
m.eval(&[0.9, 0.4, 0.2], &[1.0, 1.0, 0.0], None),
1.0 / 3.0,
epsilon = 1e-9
);
}
#[test]
fn factory_defaults_to_objective_metric() {
let ms = create_metrics(&[], "rmse", 0).unwrap();
assert_eq!(ms.len(), 1);
assert_eq!(ms[0].name(), "rmse");
assert!(create_metrics(&["nope".to_string()], "rmse", 0).is_err());
}
#[test]
fn auc_ranks_perfectly_separable() {
let m = Auc;
let auc = m.eval(&[0.1, 0.2, 0.8, 0.9], &[0.0, 0.0, 1.0, 1.0], None);
assert!((auc - 1.0).abs() < 1e-9);
assert!(m.maximize());
}
#[test]
fn ndcg_perfect_and_reversed() {
use crate::data::GroupInfo;
let m = Ndcg::new(None);
let labels = [3.0f32, 2.0, 0.0];
let perfect = [0.9f32, 0.5, 0.1];
let g = GroupInfo::from_sizes(&[3]);
assert_relative_eq!(
m.eval_grouped(&perfect, &labels, None, Some(&g)),
1.0,
epsilon = 1e-6
);
let reversed = [0.1f32, 0.5, 0.9];
let got = m.eval_grouped(&reversed, &labels, None, Some(&g));
assert_relative_eq!(got, 5.392_789 / 8.892_789, epsilon = 1e-5);
assert!(m.maximize());
}
#[test]
fn ndcg_truncation_at_k() {
let m = Ndcg::new(Some(1));
let labels = [3.0f32, 2.0, 0.0];
assert_relative_eq!(m.eval(&[0.9, 0.5, 0.1], &labels, None), 1.0, epsilon = 1e-6);
assert_relative_eq!(m.eval(&[0.1, 0.5, 0.9], &labels, None), 0.0, epsilon = 1e-6);
}
#[test]
fn map_hand_computed() {
let m = MeanAveragePrecision::new(None);
let labels = [1.0f32, 0.0, 1.0, 0.0]; assert_relative_eq!(
m.eval(&[0.9, 0.1, 0.8, 0.2], &labels, None),
1.0,
epsilon = 1e-9
);
assert_relative_eq!(
m.eval(&[0.8, 0.9, 0.1, 0.7], &labels, None),
0.5,
epsilon = 1e-9
);
assert!(m.maximize());
}
#[test]
fn ranking_metrics_average_over_groups() {
use crate::data::GroupInfo;
let labels = [1.0f32, 0.0, 1.0, 0.0];
let preds = [0.9f32, 0.1, 0.2, 0.8]; let g = GroupInfo::from_sizes(&[2, 2]);
let m = MeanAveragePrecision::new(None);
assert_relative_eq!(
m.eval_grouped(&preds, &labels, None, Some(&g)),
0.75,
epsilon = 1e-9
);
}
#[test]
fn factory_parses_ranking_metrics_with_k() {
assert_eq!(create_metric("ndcg", 0).unwrap().name(), "ndcg");
assert_eq!(create_metric("map", 0).unwrap().name(), "map");
assert_eq!(create_metric("ndcg@5", 0).unwrap().name(), "ndcg");
assert_eq!(create_metric("map@10", 0).unwrap().name(), "map");
}
#[test]
fn aucpr_perfect_and_ranks_better_than_random() {
let m = AucPr;
assert!(m.maximize());
assert_eq!(create_metric("aucpr", 0).unwrap().name(), "aucpr");
let perfect = m.eval(&[0.1, 0.2, 0.8, 0.9], &[0.0, 0.0, 1.0, 1.0], None);
assert_relative_eq!(perfect, 1.0, epsilon = 1e-9);
let labels = [1.0f32, 1.0, 1.0, 0.0, 0.0, 0.0];
let good = m.eval(&[0.9, 0.8, 0.7, 0.3, 0.2, 0.1], &labels, None);
let poor = m.eval(&[0.9, 0.2, 0.7, 0.8, 0.1, 0.3], &labels, None);
assert_relative_eq!(good, 1.0, epsilon = 1e-9);
assert!(
good > poor,
"better ranking should score higher: {good} vs {poor}"
);
assert!(poor > 0.5, "poor ranking still beats nothing: {poor}");
}
#[test]
fn aucpr_degenerate_returns_zero() {
let m = AucPr;
assert_eq!(m.eval(&[0.3, 0.6, 0.9], &[1.0, 1.0, 1.0], None), 0.0);
assert_eq!(m.eval(&[0.3, 0.6, 0.9], &[0.0, 0.0, 0.0], None), 0.0);
}
#[test]
fn mlogloss_and_merror() {
let ml = MLogLoss { num_class: 3 };
let me = MError { num_class: 3 };
let preds = [0.8, 0.1, 0.1, 0.05, 0.9, 0.05];
let labels = [0.0, 1.0];
assert!(ml.eval(&preds, &labels, None) < 0.3);
assert_eq!(me.eval(&preds, &labels, None), 0.0);
let labels_wrong = [1.0, 1.0];
assert_eq!(me.eval(&preds, &labels_wrong, None), 0.5);
}
}