use std::collections::BTreeMap;
use corescout_represent::discover::{differences, mean};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct CellScore {
pub row: usize,
pub col: usize,
pub model_mae: f64,
pub baseline_mae: f64,
pub skill: f64,
pub samples: usize,
}
#[derive(Debug, Clone, PartialEq)]
pub struct PredictionReport {
pub cells: usize,
pub cells_with_skill: usize,
pub median_skill: f64,
pub mean_skill: f64,
pub best: Vec<CellScore>,
pub worst: Vec<CellScore>,
pub accumulator_columns: Vec<usize>,
}
pub fn evaluate(
series: &BTreeMap<(usize, usize), Vec<f64>>,
accumulator_columns: &[usize],
train_fraction: f64,
) -> PredictionReport {
let mut scores: Vec<CellScore> = Vec::new();
for ((row, col), values) in series {
let is_accumulator = accumulator_columns.contains(col);
if let Some(score) = score_cell(*row, *col, values, is_accumulator, train_fraction) {
scores.push(score);
}
}
let skills: Vec<f64> = scores.iter().map(|s| s.skill).collect();
let median_skill = median(&skills);
let mean_skill = mean(&skills);
let mut ranked = scores.clone();
ranked.sort_by(|a, b| {
b.skill
.partial_cmp(&a.skill)
.unwrap_or(std::cmp::Ordering::Equal)
});
let best: Vec<CellScore> = ranked.iter().take(5).copied().collect();
let worst: Vec<CellScore> = ranked.iter().rev().take(5).copied().collect();
let mut accumulator_columns = accumulator_columns.to_vec();
accumulator_columns.sort_unstable();
PredictionReport {
cells: scores.len(),
cells_with_skill: scores.iter().filter(|s| s.skill > 0.01).count(),
median_skill,
mean_skill,
best,
worst,
accumulator_columns,
}
}
fn score_cell(
row: usize,
col: usize,
values: &[f64],
is_accumulator: bool,
train_fraction: f64,
) -> Option<CellScore> {
let deltas = differences(values);
if deltas.len() < 12 {
return None;
}
let split = ((deltas.len() as f64) * train_fraction) as usize;
if split < 6 || deltas.len() - split < 4 {
return None;
}
let train = &deltas[..split];
let test = &deltas[split..];
let baseline_delta = if is_accumulator {
mean(
&train
.iter()
.copied()
.filter(|d| d.is_finite())
.collect::<Vec<f64>>(),
)
} else {
0.0
};
let (a, b) = fit_ar1(train)?;
let mut model_error = 0.0;
let mut baseline_error = 0.0;
let mut samples = 0usize;
let mut previous = train.last().copied().unwrap_or(0.0);
for actual in test {
if !actual.is_finite() {
previous = f64::NAN;
continue;
}
if previous.is_finite() {
let predicted = a * previous + b;
model_error += (predicted - actual).abs();
baseline_error += (baseline_delta - actual).abs();
samples += 1;
}
previous = *actual;
}
if samples < 4 {
return None;
}
let model_mae = model_error / samples as f64;
let baseline_mae = baseline_error / samples as f64;
if baseline_mae <= 1e-12 && model_mae <= 1e-12 {
return None;
}
let skill = if baseline_mae > 1e-12 {
1.0 - (model_mae / baseline_mae)
} else {
-1.0
};
Some(CellScore {
row,
col,
model_mae,
baseline_mae,
skill: skill.clamp(-1.0, 1.0),
samples,
})
}
fn fit_ar1(series: &[f64]) -> Option<(f64, f64)> {
let pairs: Vec<(f64, f64)> = series
.windows(2)
.filter(|w| w[0].is_finite() && w[1].is_finite())
.map(|w| (w[0], w[1]))
.collect();
if pairs.len() < 4 {
return None;
}
let n = pairs.len() as f64;
let mean_x = pairs.iter().map(|(x, _)| *x).sum::<f64>() / n;
let mean_y = pairs.iter().map(|(_, y)| *y).sum::<f64>() / n;
let mut cov = 0.0;
let mut var = 0.0;
for (x, y) in &pairs {
cov += (x - mean_x) * (y - mean_y);
var += (x - mean_x).powi(2);
}
if var <= 1e-12 {
return Some((0.0, mean_y));
}
let a = cov / var;
Some((a, mean_y - a * mean_x))
}
fn median(values: &[f64]) -> f64 {
if values.is_empty() {
return 0.0;
}
let mut sorted: Vec<f64> = values.to_vec();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let mid = sorted.len() / 2;
if sorted.len() % 2 == 0 {
(sorted[mid - 1] + sorted[mid]) / 2.0
} else {
sorted[mid]
}
}
#[cfg(test)]
mod tests {
use super::*;
fn series(values: Vec<f64>) -> BTreeMap<(usize, usize), Vec<f64>> {
[((0usize, 0usize), values)].into_iter().collect()
}
#[test]
fn a_pure_counter_is_predicted_by_the_drift_baseline_not_by_skill() {
let values: Vec<f64> = (0..60).map(|i| (i as f64) * 10.0).collect();
let report = evaluate(&series(values), &[0], 0.7);
assert_eq!(
report.cells, 0,
"a perfectly steady counter is not scorable"
);
}
#[test]
fn momentum_is_learnable_and_shows_as_skill() {
let values: Vec<f64> = (0..120)
.map(|i| ((i as f64) * 0.35).sin() * 100.0)
.collect();
let report = evaluate(&series(values), &[], 0.7);
assert_eq!(report.cells, 1);
assert!(
report.median_skill > 0.2,
"expected real skill on a smooth signal, got {}",
report.median_skill
);
}
#[test]
fn pure_noise_yields_no_skill() {
let mut state = 12345u64;
let values: Vec<f64> = (0..120)
.map(|_| {
state = state.wrapping_mul(6364136223846793005).wrapping_add(1);
((state >> 33) as f64 / (1u64 << 31) as f64) - 0.5
})
.collect();
let report = evaluate(&series(values), &[], 0.7);
assert!(
report.median_skill < 0.15,
"claimed skill {} on noise",
report.median_skill
);
}
#[test]
fn mislabelling_an_accumulator_costs_measurable_skill() {
let values: Vec<f64> = (0..120)
.map(|i| (i as f64) * 10.0 + ((i as f64) * 0.7).sin() * 3.0)
.collect();
let knowing = evaluate(&series(values.clone()), &[0], 0.7);
let not_knowing = evaluate(&series(values), &[], 0.7);
assert_eq!(knowing.cells, 1);
assert_eq!(not_knowing.cells, 1);
assert!(
not_knowing.median_skill > knowing.median_skill,
"an unknown accumulator should flatter the model: {} vs {}",
not_knowing.median_skill,
knowing.median_skill
);
}
#[test]
fn training_and_scoring_windows_do_not_overlap() {
let mut values: Vec<f64> = (0..60).map(|i| ((i as f64) * 0.3).sin()).collect();
values.extend((0..60).map(|_| 0.0));
let report = evaluate(&series(values), &[], 0.5);
assert_eq!(report.cells, 1);
assert!(report.best[0].samples > 0);
assert!(report.best[0].samples <= 60);
assert!(report.best[0].samples < 119);
}
#[test]
fn short_series_are_not_scored() {
let report = evaluate(&series(vec![1.0, 2.0, 3.0]), &[], 0.7);
assert_eq!(report.cells, 0);
}
#[test]
fn ar1_recovers_a_known_coefficient() {
let mut values = vec![1.0];
for _ in 0..50 {
let last = *values.last().unwrap();
values.push(0.5 * last + 2.0);
}
let (a, b) = fit_ar1(&values).unwrap();
assert!((a - 0.5).abs() < 1e-6, "a = {a}");
assert!((b - 2.0).abs() < 1e-6, "b = {b}");
}
}