use std::collections::{BTreeMap, BTreeSet};
use std::fs;
use std::path::{Path, PathBuf};
use serde::{Deserialize, Serialize};
use super::records::{SampleRecord, StrataSummary, Stratum};
use super::sample::write_json;
use super::stats::{cohen_kappa, wilson_interval, Kappa, Z_95};
use super::{io_err, EvalError, Result};
use crate::core::config::BucketMap;
pub const UNCLEAR: &str = "unclear";
pub const MIXED: &str = "mixed";
pub const RELEASE_MERGE: &str = "release_merge";
pub const NO_ANSWER_LABELS: [&str; 3] = [UNCLEAR, MIXED, RELEASE_MERGE];
fn is_no_answer(label: &str) -> bool {
NO_ANSWER_LABELS.contains(&label)
}
#[derive(Debug, Clone)]
pub struct ScoreParams {
pub sample: PathBuf,
pub strata: Option<PathBuf>,
pub labels: Vec<PathBuf>,
pub adjudicated: Option<PathBuf>,
pub categories: Option<Vec<String>>,
pub db: Option<PathBuf>,
pub out: PathBuf,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct PrecisionRow {
pub key: String,
pub n: u64,
pub correct: u64,
pub precision: Option<f64>,
pub ci_low: Option<f64>,
pub ci_high: Option<f64>,
pub excluded: u64,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct CoveragePoint {
pub threshold: f64,
pub coverage: f64,
pub precision: Option<f64>,
pub n: u64,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct WeightedAccuracy {
pub estimate: f64,
pub ci_low: f64,
pub ci_high: f64,
pub population_covered: f64,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Abstention {
pub catch_all: u64,
pub unknown: u64,
pub population: u64,
pub share: f64,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct ScoreReport {
#[serde(default)]
pub scored_rater: String,
pub sample_size: u64,
#[serde(default)]
pub merges_excluded: u64,
pub labelled: u64,
pub scored: u64,
pub unclear: u64,
pub mixed: u64,
#[serde(default)]
pub release_merge: u64,
#[serde(default)]
pub window_merges_estimated: u64,
#[serde(default)]
pub window_merges_exact: Option<u64>,
pub unresolved_disagreements: u64,
pub per_rule: Vec<PrecisionRow>,
pub per_method: Vec<PrecisionRow>,
pub per_stratum: Vec<PrecisionRow>,
pub weighted_accuracy: Option<WeightedAccuracy>,
pub coverage_curve: Vec<CoveragePoint>,
pub confusion: BTreeMap<String, BTreeMap<String, u64>>,
pub abstention: Abstention,
pub kappa: Option<Kappa>,
#[serde(default)]
pub buckets: super::score_buckets::BucketScores,
}
#[derive(Debug, Deserialize)]
struct LabelRow {
sha: String,
#[serde(default)]
label: String,
}
fn read_labels(path: &Path, valid: &BTreeSet<String>) -> Result<BTreeMap<String, String>> {
let csv_err = |source| EvalError::Csv {
path: path.to_path_buf(),
source,
};
let mut reader = csv::Reader::from_path(path).map_err(csv_err)?;
let mut out = BTreeMap::new();
for row in reader.deserialize::<LabelRow>() {
let row = row.map_err(csv_err)?;
let label = row.label.trim().to_lowercase();
if label.is_empty() {
continue;
}
if !valid.contains(&label) {
return Err(EvalError::Invalid(format!(
"{}: label {label:?} for {} is not a known category, `unclear`, `mixed` or `release_merge`",
path.display(),
row.sha
)));
}
out.insert(row.sha.trim().to_string(), label);
}
Ok(out)
}
pub(crate) fn read_strata(path: &Path) -> Result<StrataSummary> {
let text = fs::read_to_string(path).map_err(io_err(path))?;
serde_json::from_str(&text).map_err(|source| EvalError::Json {
path: path.to_path_buf(),
source,
})
}
pub(crate) fn read_sample(path: &Path) -> Result<Vec<SampleRecord>> {
let text = fs::read_to_string(path).map_err(io_err(path))?;
text.lines()
.filter(|l| !l.trim().is_empty())
.map(|l| {
serde_json::from_str(l).map_err(|source| EvalError::Json {
path: path.to_path_buf(),
source,
})
})
.collect()
}
pub(super) fn precision_rows<'a>(
rows: impl Iterator<Item = (String, Option<bool>)> + 'a,
) -> Vec<PrecisionRow> {
let mut acc: BTreeMap<String, (u64, u64, u64)> = BTreeMap::new();
for (key, outcome) in rows {
let e = acc.entry(key).or_default();
match outcome {
Some(correct) => {
e.0 += 1;
e.1 += u64::from(correct);
}
None => e.2 += 1,
}
}
acc.into_iter()
.map(|(key, (n, correct, excluded))| {
let ci = wilson_interval(correct, n, Z_95);
PrecisionRow {
key,
n,
correct,
precision: (n > 0).then(|| correct as f64 / n as f64),
ci_low: ci.map(|c| c.0),
ci_high: ci.map(|c| c.1),
excluded,
}
})
.collect()
}
pub fn run_score(params: &ScoreParams) -> Result<ScoreReport> {
run_score_with_buckets(params, &BucketMap::default())
}
pub fn run_score_with_buckets(params: &ScoreParams, buckets: &BucketMap) -> Result<ScoreReport> {
if params.labels.is_empty() || params.labels.len() > 2 {
return Err(EvalError::Invalid("pass one or two --labels files".into()));
}
let sample = read_sample(¶ms.sample)?;
let merge_flags = super::merges::resolve_merges(&sample, params.db.as_deref())?;
let strata_path = params.strata.clone().unwrap_or_else(|| {
params
.sample
.parent()
.unwrap_or(Path::new("."))
.join("strata.json")
});
let mut strata = read_strata(&strata_path)?;
let window_merges_estimated =
super::merges::scale_out_merges(&mut strata, &sample, &merge_flags);
let window_merges_exact = params
.db
.as_deref()
.map(|db| super::merges::count_window_merges(db, &strata))
.transpose()?;
let mut valid: BTreeSet<String> = params
.categories
.clone()
.unwrap_or_else(|| strata.categories.clone())
.into_iter()
.chain(sample.iter().map(|r| r.predicted_category.clone()))
.map(|c| c.to_lowercase())
.collect();
valid.extend(NO_ANSWER_LABELS.map(String::from));
let vocabulary: Vec<String> = valid.iter().cloned().collect();
buckets
.check_known(&vocabulary)
.map_err(|e| EvalError::Invalid(e.to_string()))?;
let mut raters: Vec<BTreeMap<String, String>> = params
.labels
.iter()
.map(|p| read_labels(p, &valid))
.collect::<Result<_>>()?;
let mut adjudicated = match ¶ms.adjudicated {
Some(p) => read_labels(p, &valid)?,
None => BTreeMap::new(),
};
let in_sample: BTreeSet<&str> = sample.iter().map(|r| r.sha.as_str()).collect();
for sha in raters.iter().chain([&adjudicated]).flat_map(|m| m.keys()) {
if !in_sample.contains(sha.as_str()) {
return Err(EvalError::Invalid(format!(
"label for {sha}, which is not in the sample; score against a sample holding \
every rater's rows (for a subsample, the source sample.jsonl)"
)));
}
}
let sample_size = sample.len() as u64;
let merge_shas: BTreeSet<String> = sample
.iter()
.zip(&merge_flags)
.filter(|&(_, &m)| m)
.map(|(r, _)| r.sha.clone())
.collect();
let merges_excluded = merge_flags.iter().filter(|&&m| m).count() as u64;
let sample: Vec<SampleRecord> = sample
.into_iter()
.zip(&merge_flags)
.filter(|&(_, &m)| !m)
.map(|(r, _)| r)
.collect();
for labels in raters.iter_mut().chain(std::iter::once(&mut adjudicated)) {
labels.retain(|sha, _| !merge_shas.contains(sha));
}
if let Some(sha) = adjudicated.keys().find(|s| !raters[0].contains_key(*s)) {
return Err(EvalError::Invalid(format!(
"adjudicated label for {sha}, which the first --labels file leaves blank"
)));
}
let kappa = if raters.len() == 2 {
let pairs: Vec<(String, String)> = raters[0]
.iter()
.filter_map(|(sha, a)| raters[1].get(sha).map(|b| (a.clone(), b.clone())))
.collect();
cohen_kappa(&pairs)
} else {
None
};
let finals: Vec<Option<String>> = sample
.iter()
.map(|r| {
adjudicated
.get(&r.sha)
.or_else(|| raters[0].get(&r.sha))
.cloned()
})
.collect();
let unresolved = raters.get(1).map_or(0, |second| {
raters[0]
.iter()
.filter(|(sha, a)| {
!adjudicated.contains_key(*sha) && second.get(*sha).is_some_and(|b| b != *a)
})
.count() as u64
});
let outcomes: Vec<Option<Option<bool>>> = sample
.iter()
.zip(&finals)
.map(|(r, l)| {
l.as_ref().map(|l| {
(!is_no_answer(l)).then(|| l.eq_ignore_ascii_case(&r.predicted_category))
})
})
.collect();
let labelled_rows = || {
sample
.iter()
.zip(&outcomes)
.filter_map(|(r, o)| o.map(|o| (r, o)))
};
let per_rule = precision_rows(labelled_rows().map(|(r, o)| (r.rule_id.clone(), o)));
let per_method = precision_rows(labelled_rows().map(|(r, o)| (r.method.clone(), o)));
let per_stratum =
precision_rows(labelled_rows().map(|(r, o)| (r.stratum.as_str().to_string(), o)));
let weighted_accuracy = weighted_accuracy(&per_stratum, &strata);
let bucket_scores = super::score_buckets::score_buckets(
buckets,
sample
.iter()
.zip(&finals)
.zip(&outcomes)
.filter(|(_, o)| matches!(o, Some(Some(_))))
.filter_map(|((r, l), _)| l.as_deref().map(|l| (r, l))),
&strata,
);
let coverage_curve = coverage_curve(&labelled_weights(&sample, &outcomes, &strata));
let mut confusion: BTreeMap<String, BTreeMap<String, u64>> = BTreeMap::new();
for (r, l) in sample.iter().zip(&finals) {
if let Some(l) = l {
*confusion
.entry(r.predicted_category.to_lowercase())
.or_default()
.entry(l.clone())
.or_default() += 1;
}
}
let catch_all = strata.population_of(Stratum::CatchAll);
let unknown = strata.population_of(Stratum::Unknown);
let abstention = Abstention {
catch_all,
unknown,
population: strata.population,
share: if strata.population == 0 {
0.0
} else {
(catch_all + unknown) as f64 / strata.population as f64
},
};
let count = |f: &dyn Fn(&str) -> bool| finals.iter().flatten().filter(|l| f(l)).count() as u64;
let report = ScoreReport {
scored_rater: params.labels[0]
.file_name()
.map_or_else(String::new, |n| n.to_string_lossy().into_owned()),
sample_size,
merges_excluded,
labelled: finals.iter().flatten().count() as u64,
scored: outcomes
.iter()
.filter(|o| matches!(o, Some(Some(_))))
.count() as u64,
unclear: count(&|l| l == UNCLEAR),
mixed: count(&|l| l == MIXED),
release_merge: count(&|l| l == RELEASE_MERGE),
window_merges_estimated,
window_merges_exact,
unresolved_disagreements: unresolved,
per_rule,
per_method,
per_stratum,
weighted_accuracy,
coverage_curve,
confusion,
abstention,
kappa,
buckets: bucket_scores,
};
fs::create_dir_all(¶ms.out).map_err(io_err(¶ms.out))?;
write_json(¶ms.out.join("report.json"), &report)?;
let md_path = params.out.join("report.md");
fs::write(&md_path, super::report_md::render(&report, &strata)).map_err(io_err(&md_path))?;
Ok(report)
}
pub(super) fn weighted_accuracy(
per_stratum: &[PrecisionRow],
strata: &StrataSummary,
) -> Option<WeightedAccuracy> {
let scored: Vec<(f64, f64, f64)> = per_stratum
.iter()
.filter(|row| row.n > 0)
.map(|row| {
let pop = strata
.strata
.get(&row.key)
.map(|c| c.population)
.unwrap_or(0) as f64;
(pop, row.correct as f64 / row.n as f64, row.n as f64)
})
.collect();
let covered: f64 = scored.iter().map(|s| s.0).sum();
if covered <= 0.0 {
return None;
}
let estimate: f64 = scored.iter().map(|(pop, p, _)| pop / covered * p).sum();
let variance: f64 = scored
.iter()
.map(|(pop, p, n)| (pop / covered).powi(2) * p * (1.0 - p) / n)
.sum();
let half = Z_95 * variance.sqrt();
Some(WeightedAccuracy {
estimate,
ci_low: (estimate - half).max(0.0),
ci_high: (estimate + half).min(1.0),
population_covered: if strata.population == 0 {
0.0
} else {
covered / strata.population as f64
},
})
}
struct WeightedRow {
confidence: f64,
weight: f64,
outcome: Option<bool>,
}
fn labelled_weights(
sample: &[SampleRecord],
outcomes: &[Option<Option<bool>>],
strata: &StrataSummary,
) -> Vec<WeightedRow> {
let mut labelled: BTreeMap<Stratum, u64> = BTreeMap::new();
for (r, o) in sample.iter().zip(outcomes) {
if o.is_some() {
*labelled.entry(r.stratum).or_default() += 1;
}
}
sample
.iter()
.zip(outcomes)
.filter_map(|(r, o)| {
o.map(|outcome| WeightedRow {
confidence: r.confidence,
weight: strata.population_of(r.stratum) as f64
/ labelled.get(&r.stratum).copied().unwrap_or(1).max(1) as f64,
outcome,
})
})
.collect()
}
fn coverage_curve(rows: &[WeightedRow]) -> Vec<CoveragePoint> {
let mut thresholds: Vec<f64> = rows.iter().map(|r| r.confidence).collect();
thresholds.sort_by(f64::total_cmp);
thresholds.dedup_by(|a, b| (*a - *b).abs() < 1e-9);
let total_weight: f64 = rows.iter().map(|r| r.weight).sum();
thresholds
.into_iter()
.map(|t| {
let (mut kept_weight, mut w_scored, mut w_correct, mut n) = (0.0, 0.0, 0.0, 0u64);
for r in rows.iter().filter(|r| r.confidence >= t - 1e-9) {
kept_weight += r.weight;
if let Some(correct) = r.outcome {
w_scored += r.weight;
w_correct += if correct { r.weight } else { 0.0 };
n += 1;
}
}
CoveragePoint {
threshold: t,
coverage: if total_weight > 0.0 {
kept_weight / total_weight
} else {
0.0
},
precision: (w_scored > 0.0).then(|| w_correct / w_scored),
n,
}
})
.collect()
}