use super::curve::{Auc, AucPr};
use super::distributional::{DistCrps, DistNll};
use super::elementwise::{Mape, PseudoHuberError, Rmsle};
use super::quantile::{ExpectileError, QuantileError};
use super::ranking::{MeanAveragePrecision, Ndcg, Precision};
use super::survival::{AftNLogLik, CoxNLogLik, IntervalRegressionAccuracy};
use super::{
ErrorRate, GammaNLogLik, LogLoss, MError, MLogLoss, Mae, Metric, PoissonNLogLik, Rmse,
TweedieNLogLik, tweedie_name,
};
use crate::error::{HessboostError, Result};
use crate::objective::distributional::DistFamily;
use crate::objective::{Aft, AftDistribution, Expectiles, PseudoHuber, Quantiles, Tweedie};
use std::borrow::Cow;
use std::num::NonZeroUsize;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct Cutoff {
top_k: Option<NonZeroUsize>,
}
impl Cutoff {
pub fn all() -> Self {
Cutoff { top_k: None }
}
pub fn top(k: usize) -> Result<Self> {
NonZeroUsize::new(k).map(Cutoff::from).ok_or_else(|| {
HessboostError::invalid_param("eval_metric", "the `@k` cutoff must be at least 1")
})
}
pub fn top_k(&self) -> Option<NonZeroUsize> {
self.top_k
}
fn k(self) -> Option<usize> {
self.top_k.map(NonZeroUsize::get)
}
}
impl From<NonZeroUsize> for Cutoff {
fn from(k: NonZeroUsize) -> Self {
Cutoff { top_k: Some(k) }
}
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub enum EvalMetric {
Rmse,
Rmsle,
Mae,
Mape,
Mphe(PseudoHuber),
LogLoss,
Error,
Auc,
AucPr,
MLogLoss,
MError,
PoissonNLogLik,
GammaNLogLik,
TweedieNLogLik(Tweedie),
Ndcg(Cutoff),
Map(Cutoff),
Precision(Cutoff),
Quantile(Quantiles),
Expectile(Expectiles),
CoxNLogLik,
AftNLogLik(Aft),
IntervalRegressionAccuracy,
Nll(DistFamily),
Crps(DistFamily),
}
pub(super) fn cutoff_name(base: &str, k: Option<usize>) -> String {
k.map_or_else(|| base.to_string(), |k| format!("{base}@{k}"))
}
impl EvalMetric {
pub(crate) fn flat_name(&self) -> Cow<'static, str> {
match self {
EvalMetric::TweedieNLogLik(tweedie) => {
Cow::Owned(format!("tweedie-nloglik@{}", tweedie.variance_power()))
}
_ => self.name(),
}
}
pub fn name(&self) -> Cow<'static, str> {
Cow::Borrowed(match self {
EvalMetric::Rmse => "rmse",
EvalMetric::Rmsle => "rmsle",
EvalMetric::Mae => "mae",
EvalMetric::Mape => "mape",
EvalMetric::Mphe(_) => "mphe",
EvalMetric::LogLoss => "logloss",
EvalMetric::Error => "error",
EvalMetric::Auc => "auc",
EvalMetric::AucPr => "aucpr",
EvalMetric::MLogLoss => "mlogloss",
EvalMetric::MError => "merror",
EvalMetric::PoissonNLogLik => "poisson-nloglik",
EvalMetric::GammaNLogLik => "gamma-nloglik",
EvalMetric::TweedieNLogLik(tweedie) => {
return Cow::Owned(tweedie_name(tweedie.variance_power()));
}
EvalMetric::Ndcg(cutoff) => return Cow::Owned(cutoff_name("ndcg", cutoff.k())),
EvalMetric::Map(cutoff) => return Cow::Owned(cutoff_name("map", cutoff.k())),
EvalMetric::Precision(cutoff) => return Cow::Owned(cutoff_name("pre", cutoff.k())),
EvalMetric::Quantile(_) => "quantile",
EvalMetric::Expectile(_) => "expectile",
EvalMetric::CoxNLogLik => "cox-nloglik",
EvalMetric::AftNLogLik(_) => "aft-nloglik",
EvalMetric::IntervalRegressionAccuracy => "interval-regression-accuracy",
EvalMetric::Nll(_) => "nll",
EvalMetric::Crps(_) => "crps",
})
}
pub fn build(&self, n_outputs: usize) -> Result<Box<dyn Metric>> {
let classes = || {
if n_outputs >= 2 {
Ok(n_outputs)
} else {
Err(HessboostError::invalid_param(
"eval_metric",
format!(
"`{}` scores one probability per class and needs a multiclass model, \
got {n_outputs} output(s)",
self.name()
),
))
}
};
Ok(match self {
EvalMetric::Rmse => Box::new(Rmse),
EvalMetric::Rmsle => Box::new(Rmsle),
EvalMetric::Mae => Box::new(Mae),
EvalMetric::Mape => Box::new(Mape),
EvalMetric::Mphe(huber) => Box::new(PseudoHuberError::new(huber.slope() as f32)),
EvalMetric::LogLoss => Box::new(LogLoss),
EvalMetric::Error => Box::new(ErrorRate),
EvalMetric::Auc => Box::new(Auc),
EvalMetric::AucPr => Box::new(AucPr),
EvalMetric::MLogLoss => Box::new(MLogLoss {
num_class: classes()?,
}),
EvalMetric::MError => Box::new(MError {
num_class: classes()?,
}),
EvalMetric::PoissonNLogLik => Box::new(PoissonNLogLik),
EvalMetric::GammaNLogLik => Box::new(GammaNLogLik),
EvalMetric::TweedieNLogLik(tweedie) => {
Box::new(TweedieNLogLik::new(tweedie.variance_power()))
}
EvalMetric::Ndcg(cutoff) => Box::new(Ndcg::new(cutoff.k())),
EvalMetric::Map(cutoff) => Box::new(MeanAveragePrecision::new(cutoff.k())),
EvalMetric::Precision(cutoff) => Box::new(Precision::new(cutoff.k())),
EvalMetric::Quantile(quantiles) => Box::new(QuantileError::new(quantiles.alpha_f32())),
EvalMetric::Expectile(expectiles) => {
Box::new(ExpectileError::new(expectiles.alpha_f32()))
}
EvalMetric::CoxNLogLik => Box::new(CoxNLogLik),
EvalMetric::AftNLogLik(aft) => {
Box::new(AftNLogLik::new(aft.distribution(), aft.scale() as f32))
}
EvalMetric::IntervalRegressionAccuracy => Box::new(IntervalRegressionAccuracy),
EvalMetric::Nll(family) => Box::new(DistNll::new(*family)),
EvalMetric::Crps(family) => Box::new(DistCrps::new(*family)),
})
}
pub(crate) fn borrowed_keys(name: &str) -> &'static [&'static str] {
match name {
"mphe" => &["huber_slope"],
"quantile" => &["quantile_alpha"],
"expectile" => &["expectile_alpha"],
"aft-nloglik" => &["aft_loss_distribution", "aft_loss_distribution_scale"],
_ => &[],
}
}
pub(crate) fn from_xgboost(name: &str, source: &XgboostMetricSource<'_>) -> Result<Self> {
let (base, suffix) = match name.split_once('@') {
Some((b, s)) => (b, Some(s)),
None => (name, None),
};
if suffix.is_some() && !matches!(base, "tweedie-nloglik" | "ndcg" | "map" | "pre") {
return Err(invalid_metric(
name,
&format!("`{base}` takes no `@` suffix"),
));
}
let cutoff = || rank_cutoff(name, suffix);
Ok(match base {
"rmse" => EvalMetric::Rmse,
"rmsle" => EvalMetric::Rmsle,
"mae" => EvalMetric::Mae,
"mape" => EvalMetric::Mape,
"mphe" => EvalMetric::Mphe(source.huber()?),
"logloss" => EvalMetric::LogLoss,
"error" => EvalMetric::Error,
"auc" => EvalMetric::Auc,
"aucpr" => EvalMetric::AucPr,
"mlogloss" => EvalMetric::MLogLoss,
"merror" => EvalMetric::MError,
"poisson-nloglik" => EvalMetric::PoissonNLogLik,
"gamma-nloglik" => EvalMetric::GammaNLogLik,
"tweedie-nloglik" => EvalMetric::TweedieNLogLik(tweedie_power(name, suffix)?),
"ndcg" => EvalMetric::Ndcg(cutoff()?),
"map" => EvalMetric::Map(cutoff()?),
"pre" => EvalMetric::Precision(cutoff()?),
"quantile" => EvalMetric::Quantile(source.quantiles()?),
"expectile" => EvalMetric::Expectile(source.expectiles()?),
"cox-nloglik" => EvalMetric::CoxNLogLik,
"aft-nloglik" => EvalMetric::AftNLogLik(source.aft()?),
"interval-regression-accuracy" => EvalMetric::IntervalRegressionAccuracy,
"nll" | "crps" => {
let family = source.distribution.ok_or_else(|| {
HessboostError::invalid_param(
"eval_metric",
format!(
"`{name}` scores predicted distributions and needs a `dist:*` objective"
),
)
})?;
if base == "nll" {
EvalMetric::Nll(family)
} else {
EvalMetric::Crps(family)
}
}
other => return Err(HessboostError::unknown("metric", other)),
})
}
}
pub(crate) struct XgboostMetricSource<'a> {
pub(crate) huber_slope: f64,
pub(crate) quantile_alpha: &'a [f64],
pub(crate) expectile_alpha: &'a [f64],
pub(crate) aft_loss_distribution: AftDistribution,
pub(crate) aft_loss_distribution_scale: f64,
pub(crate) distribution: Option<DistFamily>,
}
impl XgboostMetricSource<'_> {
fn huber(&self) -> Result<PseudoHuber> {
PseudoHuber::new(self.huber_slope)
}
fn quantiles(&self) -> Result<Quantiles> {
Quantiles::new(self.quantile_alpha.iter().copied())
}
fn expectiles(&self) -> Result<Expectiles> {
Expectiles::new(self.expectile_alpha.iter().copied())
}
fn aft(&self) -> Result<Aft> {
Aft::new(self.aft_loss_distribution, self.aft_loss_distribution_scale)
}
}
fn invalid_metric(name: &str, reason: &str) -> HessboostError {
HessboostError::invalid_param("eval_metric", format!("`{name}`: {reason}"))
}
fn rank_cutoff(name: &str, suffix: Option<&str>) -> Result<Cutoff> {
match suffix {
None => Ok(Cutoff::all()),
Some(s) if s.ends_with('-') => Err(invalid_metric(
name,
"the `-` variants of the ranking metrics are not implemented",
)),
Some(s) => match s.parse::<usize>().ok().and_then(NonZeroUsize::new) {
Some(k) if s.bytes().all(|b| b.is_ascii_digit()) => Ok(Cutoff::from(k)),
_ => Err(invalid_metric(
name,
"the `@k` cutoff must be a positive integer",
)),
},
}
}
fn tweedie_power(name: &str, suffix: Option<&str>) -> Result<Tweedie> {
match suffix {
None => Ok(Tweedie::default()),
Some(s) => s
.parse::<f64>()
.ok()
.and_then(|rho| Tweedie::new(rho).ok())
.ok_or_else(|| invalid_metric(name, "the variance power `@rho` must be in [1, 2)")),
}
}
#[cfg(test)]
pub(crate) fn named(
name: &str,
n_outputs: usize,
source: &XgboostMetricSource<'_>,
) -> Result<Box<dyn Metric>> {
EvalMetric::from_xgboost(name, source)?.build(n_outputs)
}
#[cfg(test)]
pub(crate) const DEFAULT_SOURCE: XgboostMetricSource<'static> = XgboostMetricSource {
huber_slope: 1.0,
quantile_alpha: &[],
expectile_alpha: &[],
aft_loss_distribution: AftDistribution::Normal,
aft_loss_distribution_scale: 1.0,
distribution: None,
};