use crate::metric::{MetricName, Numeric, state::ConfusionStatsState};
use super::{
Metric, MetricAttributes, MetricMetadata, NumericAttributes, NumericEntry, SerializedEntry,
classification::{ClassReduction, ClassificationMetricConfig, DecisionRule},
confusion_stats::{ConfusionStats, ConfusionStatsInput},
state::FormatOptions,
};
use std::{num::NonZeroUsize, sync::Arc};
#[derive(Clone)]
pub struct RecallMetric {
name: MetricName,
state: ConfusionStatsState,
config: ClassificationMetricConfig,
}
impl Default for RecallMetric {
fn default() -> Self {
Self::new(Default::default())
}
}
impl RecallMetric {
fn new(config: ClassificationMetricConfig) -> Self {
let state = Default::default();
let name = Arc::new(format!(
"Recall @ {:?} [{:?}]",
config.decision_rule, config.class_reduction
));
Self {
state,
config,
name,
}
}
#[allow(dead_code)]
pub fn binary(threshold: f64) -> Self {
Self::new(ClassificationMetricConfig {
decision_rule: DecisionRule::Threshold(threshold),
..Default::default()
})
}
#[allow(dead_code)]
pub fn multiclass(top_k: usize, class_reduction: ClassReduction) -> Self {
Self::new(ClassificationMetricConfig {
decision_rule: DecisionRule::TopK(
NonZeroUsize::new(top_k).expect("top_k must be non-zero"),
),
class_reduction,
})
}
#[allow(dead_code)]
pub fn multilabel(threshold: f64, class_reduction: ClassReduction) -> Self {
Self::new(ClassificationMetricConfig {
decision_rule: DecisionRule::Threshold(threshold),
class_reduction,
})
}
}
impl Metric for RecallMetric {
type Input = ConfusionStatsInput;
fn update(&mut self, input: &Self::Input, _metadata: &MetricMetadata) -> SerializedEntry {
let [sample_size, _] = input.predictions.dims();
let stats = ConfusionStats::new(input, &self.config);
let tp = Some(stats.clone().true_positive());
let fn_ = Some(stats.false_negative());
self.state.update(tp, None, fn_, sample_size);
self.state.compute_update(
self.config.class_reduction,
FormatOptions::new(self.name()).unit("%").precision(2),
|tp, _, fn_| {
let (tp, fn_) = (tp.unwrap(), fn_.unwrap());
let denominator = tp.clone() + fn_;
let mask = denominator.clone().equal_elem(0.0);
let actual_positive = denominator.mask_fill(mask, 1.0);
(tp / actual_positive) * 100.0
},
)
}
fn compute(&mut self) -> SerializedEntry {
self.state
.compute_final(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 RecallMetric {
fn value(&self) -> Option<NumericEntry> {
self.state.current_value()
}
fn running_value(&self) -> Option<NumericEntry> {
self.state.running_value()
}
fn final_value(&self) -> NumericEntry {
self.state.final_value()
}
}
#[cfg(test)]
mod tests {
use super::{
ClassReduction::{self, *},
Metric, MetricMetadata, RecallMetric,
};
use crate::metric::{ConfusionStatsInput, Numeric};
use crate::tests::{ClassificationType, THRESHOLD, dummy_classification_input};
use burn_core::{
Tensor,
tensor::{TensorData, Tolerance},
};
use rstest::rstest;
#[rstest]
#[case::binary(THRESHOLD, 0.5)]
fn test_binary_recall(#[case] threshold: f64, #[case] expected: f64) {
let input = dummy_classification_input(&ClassificationType::Binary).into();
let mut metric = RecallMetric::binary(threshold);
let _entry = metric.update(&input, &MetricMetadata::fake());
TensorData::from([metric.value().unwrap().current()])
.assert_approx_eq::<f64>(&TensorData::from([expected * 100.0]), Tolerance::default())
}
#[rstest]
#[case::multiclass_micro_k1(Micro, 1, 3.0/5.0)]
#[case::multiclass_micro_k2(Micro, 2, 4.0/5.0)]
#[case::multiclass_macro_k1(Macro, 1, (0.5 + 1.0 + 0.5)/3.0)]
#[case::multiclass_macro_k2(Macro, 2, (1.0 + 1.0 + 0.5)/3.0)]
fn test_multiclass_recall(
#[case] class_reduction: ClassReduction,
#[case] top_k: usize,
#[case] expected: f64,
) {
let input = dummy_classification_input(&ClassificationType::Multiclass).into();
let mut metric = RecallMetric::multiclass(top_k, class_reduction);
let _entry = metric.update(&input, &MetricMetadata::fake());
TensorData::from([metric.value().unwrap().current()])
.assert_approx_eq::<f64>(&TensorData::from([expected * 100.0]), Tolerance::default())
}
#[rstest]
#[case::multilabel_micro(Micro, THRESHOLD, 5.0/9.0)]
#[case::multilabel_macro(Macro, THRESHOLD, (0.5 + 1.0 + 1.0/3.0)/3.0)]
fn test_multilabel_recall(
#[case] class_reduction: ClassReduction,
#[case] threshold: f64,
#[case] expected: f64,
) {
let input = dummy_classification_input(&ClassificationType::Multilabel).into();
let mut metric = RecallMetric::multilabel(threshold, class_reduction);
let _entry = metric.update(&input, &MetricMetadata::fake());
TensorData::from([metric.value().unwrap().current()])
.assert_approx_eq::<f64>(&TensorData::from([expected * 100.0]), Tolerance::default())
}
#[test]
fn test_parameterized_unique_name() {
let metric_a = RecallMetric::multiclass(1, ClassReduction::Macro);
let metric_b = RecallMetric::multiclass(2, ClassReduction::Macro);
let metric_c = RecallMetric::multiclass(1, ClassReduction::Macro);
assert_ne!(metric_a.name(), metric_b.name());
assert_eq!(metric_a.name(), metric_c.name());
let metric_a = RecallMetric::binary(0.5);
let metric_b = RecallMetric::binary(0.75);
assert_ne!(metric_a.name(), metric_b.name());
}
#[test]
fn test_recall_global_aggregation() {
let mut metric = RecallMetric::binary(THRESHOLD);
let input_batch1 = ConfusionStatsInput {
predictions: Tensor::from([[0.9], [0.1], [0.1]]),
targets: Tensor::from([[1], [1], [0]]),
};
let _ = metric.update(&input_batch1, &MetricMetadata::fake());
let input_batch2 = ConfusionStatsInput {
predictions: Tensor::from([[0.9], [0.1], [0.1], [0.1], [0.1], [0.1]]),
targets: Tensor::from([[1], [1], [1], [1], [1], [0]]),
};
let _ = metric.update(&input_batch2, &MetricMetadata::fake());
let _final_entry = metric.compute();
let global_recall = metric.final_value().current();
let expected_global_recall = (2.0 / 7.0) * 100.0;
TensorData::from([global_recall]).assert_approx_eq::<f32>(
&TensorData::from([expected_global_recall]),
Tolerance::default(),
);
}
}