use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use super::records::{SampleRecord, StrataSummary};
use super::score::{precision_rows, weighted_accuracy, PrecisionRow, WeightedAccuracy};
use crate::core::config::BucketMap;
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct BucketRow {
pub bucket: String,
pub labelled: u64,
pub predicted: u64,
pub primary: PrecisionRow,
pub secondary: Option<PrecisionRow>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct BucketScores {
pub map: BucketMap,
pub primary: Option<WeightedAccuracy>,
pub secondary: Option<WeightedAccuracy>,
pub primary_counts: PrecisionRow,
pub secondary_counts: PrecisionRow,
pub secondary_buckets: Vec<String>,
pub per_bucket: Vec<BucketRow>,
pub unmapped_labels: u64,
pub unmapped_predictions: u64,
}
struct Scored<'a> {
stratum: String,
label_bucket: &'a str,
primary: bool,
secondary: Option<bool>,
}
pub(super) fn score_buckets<'a>(
map: &BucketMap,
rows: impl Iterator<Item = (&'a SampleRecord, &'a str)>,
strata: &StrataSummary,
) -> BucketScores {
let (mut scored, mut unmapped_labels, mut unmapped_predictions) = (Vec::new(), 0, 0);
let mut predicted_in: BTreeMap<&str, u64> = BTreeMap::new();
for (r, label) in rows {
let predicted = map.bucket_of(&r.predicted_category);
match predicted {
Some(b) => *predicted_in.entry(b).or_default() += 1,
None => unmapped_predictions += 1,
}
let Some(label_bucket) = map.bucket_of(label) else {
unmapped_labels += 1;
continue;
};
scored.push(Scored {
stratum: r.stratum.as_str().to_string(),
label_bucket,
primary: predicted == Some(label_bucket),
secondary: map
.has_secondary(label_bucket)
.then(|| label.eq_ignore_ascii_case(r.predicted_category.trim())),
});
}
let primary_strata =
precision_rows(scored.iter().map(|s| (s.stratum.clone(), Some(s.primary))));
let secondary_strata = precision_rows(
scored
.iter()
.filter_map(|s| s.secondary.map(|ok| (s.stratum.clone(), Some(ok)))),
);
let overall = |key: &str, outcomes: Vec<bool>| {
precision_rows(outcomes.into_iter().map(|ok| (key.to_string(), Some(ok))))
.pop()
.unwrap_or_else(|| PrecisionRow {
key: key.to_string(),
..PrecisionRow::default()
})
};
let per_bucket = map
.buckets()
.iter()
.map(|b| {
let mine: Vec<&Scored> = scored.iter().filter(|s| s.label_bucket == b.name).collect();
let secondary: Vec<bool> = mine.iter().filter_map(|s| s.secondary).collect();
BucketRow {
bucket: b.name.clone(),
labelled: mine.len() as u64,
predicted: predicted_in.get(b.name.as_str()).copied().unwrap_or(0),
primary: overall(&b.name, mine.iter().map(|s| s.primary).collect()),
secondary: map
.has_secondary(&b.name)
.then(|| overall(&b.name, secondary)),
}
})
.collect();
BucketScores {
map: map.clone(),
primary: weighted_accuracy(&primary_strata, strata),
secondary: weighted_accuracy(&secondary_strata, strata),
primary_counts: overall("primary", scored.iter().map(|s| s.primary).collect()),
secondary_counts: overall(
"secondary",
scored.iter().filter_map(|s| s.secondary).collect(),
),
secondary_buckets: map
.buckets()
.iter()
.filter(|b| map.has_secondary(&b.name))
.map(|b| b.name.clone())
.collect(),
per_bucket,
unmapped_labels,
unmapped_predictions,
}
}