use super::MetricMetadata;
use super::state::{FormatOptions, PredictionAccumulatorState};
use crate::metric::{
ClassReduction, ConfusionStatsInput, Metric, MetricAttributes, MetricName, Numeric,
NumericAttributes, SerializedEntry,
};
use burn_core::tensor::{Int, Tensor};
use std::sync::Arc;
#[derive(Clone)]
pub struct AucPrMetric {
name: MetricName,
state: PredictionAccumulatorState,
class_reduction: ClassReduction,
}
impl Default for AucPrMetric {
fn default() -> Self {
Self::new(Default::default())
}
}
impl AucPrMetric {
fn new(class_reduction: ClassReduction) -> Self {
let state = Default::default();
let name = Arc::new(format!("AUC-PR [{:?}]", class_reduction));
Self {
state,
class_reduction,
name,
}
}
#[allow(dead_code)]
pub fn binary() -> Self {
Self::new(ClassReduction::default())
}
#[allow(dead_code)]
pub fn multiclass(class_reduction: ClassReduction) -> Self {
Self::new(class_reduction)
}
#[allow(dead_code)]
pub fn multilabel(class_reduction: ClassReduction) -> Self {
Self::new(class_reduction)
}
fn average_precision(scores: Tensor<2>, targets: Tensor<2>) -> Tensor<1> {
let [n, _c] = scores.dims();
let device = scores.device();
let order = scores.argsort_descending(0);
let sorted_targets = targets.clone().gather(0, order);
let tp = sorted_targets.clone().cumsum(0);
let ranks = Tensor::<1, Int>::arange(1..n as i64 + 1, &device)
.float()
.reshape([n, 1]);
let precision = tp / ranks;
let p_total = targets.sum_dim(0);
let delta_recall = sorted_targets / p_total;
(precision * delta_recall)
.sum_dim(0)
.squeeze_dims::<1>(&[0])
}
}
impl Metric for AucPrMetric {
type Input = ConfusionStatsInput;
fn update(
&mut self,
input: &ConfusionStatsInput,
_metadata: &MetricMetadata,
) -> SerializedEntry {
self.state
.accumulate(input.predictions.clone(), input.targets.clone());
self.state
.serialize_placeholder(FormatOptions::new(self.name()).unit("%").precision(2))
}
fn compute(&mut self) -> SerializedEntry {
if self.state.is_empty() {
return self
.state
.serialize_placeholder(FormatOptions::new(self.name()).unit("%").precision(2));
}
let (predictions, targets) = self.state.tensors();
let [n, c] = predictions.dims();
let (scores, targets) = match self.class_reduction {
ClassReduction::Macro => (predictions, targets.float()),
ClassReduction::Micro => (
predictions.reshape([n * c, 1]),
targets.float().reshape([n * c, 1]),
),
};
let ap = Self::average_precision(scores, targets);
let keep = ap
.clone()
.is_nan()
.bool_not()
.argwhere()
.squeeze_dim::<1>(1);
let metric = if keep.dims()[0] == 0 {
log::warn!(
"AUC-PR is undefined (no class has positive samples in the epoch); reporting \
0.5 as a neutral fallback."
);
0.5
} else {
ap.select(0, keep).mean().into_scalar()
};
self.state.compute(
100.0 * metric,
FormatOptions::new(self.name()).unit("%").precision(2),
)
}
fn clear(&mut self) {
self.state.reset()
}
fn name(&self) -> MetricName {
self.name.clone()
}
fn attributes(&self) -> MetricAttributes {
NumericAttributes {
unit: Some("%".to_string()),
higher_is_better: true,
}
.into()
}
}
impl Numeric for AucPrMetric {
fn value(&self) -> Option<super::NumericEntry> {
None }
fn running_value(&self) -> Option<super::NumericEntry> {
None
}
fn final_value(&self) -> super::NumericEntry {
self.state
.value()
.expect("Compute must be called to get final value")
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::metric::ClassReduction::{self, *};
use burn_core::tensor::{TensorData, Tolerance};
use rstest::rstest;
#[derive(Clone, Copy)]
enum Data {
Binary,
Multiclass,
Multilabel,
}
fn input(data: Data) -> ConfusionStatsInput {
let dev = Default::default();
match data {
Data::Binary => ConfusionStatsInput::new(
Tensor::from_data([[0.63], [0.25], [0.71], [0.3], [0.07], [0.66]], &dev),
Tensor::from_data([[0], [1], [0], [0], [0], [0]], &dev),
),
Data::Multiclass => ConfusionStatsInput::new(
Tensor::from_data(
[
[0.45, 0.3, 0.36],
[0.83, 0.24, 0.09],
[0.19, 0.39, 0.29],
[0.3, 0.14, 0.46],
[0.73, 0.74, 0.16],
[0.43, 0.37, 0.88],
],
&dev,
),
Tensor::from_data(
[
[0, 0, 1],
[0, 0, 1],
[0, 0, 1],
[0, 0, 1],
[0, 1, 0],
[1, 0, 0],
],
&dev,
),
),
Data::Multilabel => ConfusionStatsInput::new(
Tensor::from_data(
[
[0.1, 0.73, 0.84],
[0.84, 0.74, 0.24],
[0.13, 0.54, 0.54],
[0.49, 0.48, 0.71],
[0.9, 0.17, 0.43],
[0.11, 0.29, 0.23],
],
&dev,
),
Tensor::from_data(
[
[1, 0, 1],
[0, 0, 1],
[0, 0, 1],
[0, 0, 1],
[1, 0, 0],
[1, 1, 0],
],
&dev,
),
),
}
}
#[rstest]
#[case::binary_macro(Data::Binary, Macro, 0.2)]
#[case::binary_micro(Data::Binary, Micro, 0.2)]
#[case::multiclass_macro(Data::Multiclass, Macro, 0.6319444444444444)]
#[case::multiclass_micro(Data::Multiclass, Micro, 0.379975579975580)]
#[case::multilabel_macro(Data::Multilabel, Macro, 0.5944444444444444)]
#[case::multilabel_micro(Data::Multilabel, Micro, 0.5918017848017848)]
fn test_auc_pr(
#[case] data: Data,
#[case] class_reduction: ClassReduction,
#[case] expected: f64,
) {
let mut metric = AucPrMetric::new(class_reduction);
let _entry = metric.update(&input(data), &MetricMetadata::fake());
let _entry = metric.compute();
TensorData::from([metric.final_value().current()])
.assert_approx_eq::<f64>(&TensorData::from([expected * 100.0]), Tolerance::default());
}
#[test]
fn test_auc_pr_accumulates_across_batches() {
let dev = Default::default();
let mut single = AucPrMetric::binary();
single.update(
&ConfusionStatsInput::new(
Tensor::from_data([[0.9], [0.4], [0.8], [0.2], [0.6], [0.1]], &dev),
Tensor::from_data([[1], [0], [1], [0], [1], [0]], &dev),
),
&MetricMetadata::fake(),
);
single.compute();
let mut split = AucPrMetric::binary();
split.update(
&ConfusionStatsInput::new(
Tensor::from_data([[0.9], [0.4], [0.8]], &dev),
Tensor::from_data([[1], [0], [1]], &dev),
),
&MetricMetadata::fake(),
);
split.update(
&ConfusionStatsInput::new(
Tensor::from_data([[0.2], [0.6], [0.1]], &dev),
Tensor::from_data([[0], [1], [0]], &dev),
),
&MetricMetadata::fake(),
);
split.compute();
TensorData::from([split.final_value().current()]).assert_approx_eq::<f64>(
&TensorData::from([single.final_value().current()]),
Tolerance::default(),
);
}
#[test]
#[should_panic = "Compute must be called to get final value"]
fn test_auc_pr_should_panic_before_compute() {
let dev = Default::default();
let mut split = AucPrMetric::binary();
split.update(
&ConfusionStatsInput::new(
Tensor::from_data([[0.9], [0.4], [0.8]], &dev),
Tensor::from_data([[1], [0], [1]], &dev),
),
&MetricMetadata::fake(),
);
assert!(split.value().is_none());
assert!(split.running_value().is_none());
split.final_value();
}
}