use super::{Metric, argsort_desc_by, stable_argsort};
use crate::data::MetaInfo;
#[derive(Debug, Clone, Copy)]
struct Column<'a> {
values: &'a [f32],
offset: usize,
stride: usize,
}
impl<'a> Column<'a> {
fn whole(values: &'a [f32]) -> Self {
Column {
values,
offset: 0,
stride: 1,
}
}
fn target(values: &'a [f32], offset: usize, stride: usize) -> Self {
Column {
values,
offset,
stride,
}
}
fn len(self) -> usize {
self.values
.len()
.saturating_sub(self.offset)
.div_ceil(self.stride)
}
#[inline]
fn get(self, row: usize) -> f32 {
self.values[row * self.stride + self.offset]
}
}
trait CurveMetric: Metric {
fn eval_column(&self, preds: Column<'_>, labels: Column<'_>, weights: Option<&[f32]>) -> f64;
}
fn eval_columns(
metric: &impl CurveMetric,
preds: Column<'_>,
labels: Column<'_>,
weights: Option<&[f32]>,
) -> f64 {
let n = labels.len();
if preds.len() != n || weights.is_some_and(|w| w.len() != n) {
return f64::NAN;
}
metric.eval_column(preds, labels, weights)
}
fn tie_runs<'a>(
order: &'a [usize],
preds: Column<'a>,
) -> 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.get(order[end]) == preds.get(order[start]) {
end += 1;
}
let run = start..end;
start = end;
Some(run)
})
}
fn macro_average_targets(metric: &impl CurveMetric, preds: &[f32], info: &MetaInfo) -> f64 {
let k = info.n_targets();
if k <= 1 {
return metric.eval_grouped(preds, info.label_values(), info.weights, info.group);
}
let mut total = 0.0;
for target in 0..k {
total += eval_columns(
metric,
Column::target(preds, target, k),
Column::target(info.label_values(), target, k),
info.weights,
);
}
total / k as f64
}
#[derive(Debug, Clone, Copy, Default)]
#[non_exhaustive]
pub(crate) 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 {
eval_columns(self, Column::whole(preds), Column::whole(labels), weights)
}
fn eval_info(&self, preds: &[f32], info: &MetaInfo) -> f64 {
macro_average_targets(self, preds, info)
}
}
impl CurveMetric for Auc {
fn eval_column(&self, preds: Column<'_>, labels: Column<'_>, _weights: Option<&[f32]>) -> f64 {
let n = labels.len();
let order = stable_argsort(n, |&a, &b| preds.get(a).total_cmp(&preds.get(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, &rank) in ranks.iter().enumerate() {
if labels.get(k) > 0.5 {
sum_pos_rank += rank;
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)
}
}
#[derive(Debug, Clone, Copy, Default)]
#[non_exhaustive]
pub(crate) 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 {
eval_columns(self, Column::whole(preds), Column::whole(labels), weights)
}
fn eval_info(&self, preds: &[f32], info: &MetaInfo) -> f64 {
macro_average_targets(self, preds, info)
}
}
impl CurveMetric for AucPr {
fn eval_column(&self, preds: Column<'_>, labels: Column<'_>, weights: Option<&[f32]>) -> f64 {
let n = labels.len();
let w_of = |i: usize| weights.map_or(1.0, |ws| f64::from(ws[i]));
let order = argsort_desc_by(n, |i| preds.get(i));
let mut total_pos = 0.0f64;
let mut total_neg = 0.0f64;
for (i, label) in (0..n).map(|i| (i, labels.get(i))) {
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.get(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
}
}