use std::collections::{BTreeMap, HashSet};
use crate::params::Params;
use super::EvalMode;
use super::accumulate::{AccumulatedEval, EvalGrouping, accumulate_impl};
use super::catalog::MetricDef;
use super::mode::FreqGroups;
pub(super) fn mean_or_missing(sum: f64, count: usize) -> f64 {
if count == 0 { -1.0 } else { sum / count as f64 }
}
pub(super) fn mean_of_valid(values: impl Iterator<Item = f64>) -> f64 {
let (sum, count) = values
.filter(|&v| crate::metrics::is_computed(v))
.fold((0.0f64, 0usize), |(s, c), v| (s + v, c + 1));
mean_or_missing(sum, count)
}
#[inline]
pub(super) fn metric_delta(a: f64, b: f64) -> f64 {
if a >= 0.0 && b >= 0.0 { b - a } else { 0.0 }
}
pub(super) fn stats_to_map(metric_keys: &[&str], stats: &[f64]) -> BTreeMap<String, f64> {
metric_keys
.iter()
.zip(stats.iter())
.map(|(&k, &v)| (k.to_string(), v))
.collect()
}
pub(super) fn accumulate_and_summarize(
grouping: &EvalGrouping<'_>,
img_filter: Option<&HashSet<u64>>,
metrics: &[MetricDef],
) -> (AccumulatedEval, Vec<f64>) {
let ev = grouping.eval();
let acc = accumulate_impl(grouping, img_filter);
let stats = summarize_impl(&acc, &ev.params, ev.eval_mode, ev.freq_groups(), metrics);
(acc, stats)
}
fn ap_samples(
eval: &AccumulatedEval,
eval_mode: EvalMode,
t_idx: usize,
k_idx: usize,
a_idx: usize,
m_idx: usize,
) -> impl Iterator<Item = f64> + '_ {
let all_points = eval_mode == EvalMode::OpenImages;
let n = if all_points { 1 } else { eval.shape.r };
(0..n)
.map(move |r_idx| {
if all_points {
eval.ap_all_points[eval.recall_idx(t_idx, k_idx, a_idx, m_idx)]
} else {
eval.precision[eval.precision_idx(t_idx, r_idx, k_idx, a_idx, m_idx)]
}
})
.filter(|&v| crate::metrics::is_computed(v))
}
pub(super) fn per_cat_ap_static(
eval: &AccumulatedEval,
params: &Params,
eval_mode: EvalMode,
) -> Vec<f64> {
let a_idx = params.all_area_idx();
let m_idx = params.max_det_idx();
(0..eval.shape.k)
.map(|k_idx| {
mean_of_valid(
(0..eval.shape.t)
.flat_map(|t_idx| ap_samples(eval, eval_mode, t_idx, k_idx, a_idx, m_idx)),
)
})
.collect()
}
pub(super) fn summarize_impl(
eval: &AccumulatedEval,
params: &Params,
eval_mode: EvalMode,
freq_groups: &FreqGroups,
metrics: &[MetricDef],
) -> Vec<f64> {
let summarize_stat = |ap: bool, iou_thr: Option<f64>, area_lbl: &str, max_det: usize| -> f64 {
let Some(a_idx) = params.area_range_idx(area_lbl) else {
return -1.0;
};
let Some(m_idx) = params.max_dets.iter().position(|&d| d == max_det) else {
return -1.0;
};
let t_indices: Vec<usize> = if let Some(thr) = iou_thr {
params.iou_thr_idx(thr).map(|i| vec![i]).unwrap_or_default()
} else {
(0..eval.shape.t).collect()
};
if ap {
mean_of_valid(t_indices.iter().flat_map(|&t_idx| {
(0..eval.shape.k)
.flat_map(move |k_idx| ap_samples(eval, eval_mode, t_idx, k_idx, a_idx, m_idx))
}))
} else {
mean_of_valid(t_indices.iter().flat_map(|&t_idx| {
(0..eval.shape.k)
.map(move |k_idx| eval.recall[eval.recall_idx(t_idx, k_idx, a_idx, m_idx)])
}))
}
};
let per_cat_ap: Vec<f64> = if eval_mode == EvalMode::Lvis {
per_cat_ap_static(eval, params, eval_mode)
} else {
Vec::new()
};
metrics
.iter()
.map(|m| match m.freq_group {
Some(fg) => mean_of_valid(
freq_groups
.get(fg)
.iter()
.filter_map(|&k| per_cat_ap.get(k).copied()),
),
None => summarize_stat(m.ap, m.iou_thr, m.area_lbl, m.max_det),
})
.collect()
}