use std::collections::{HashMap, HashSet};
use crate::analyses::architecture_trend::{
import_graph_from_live_paths, live_paths_at, sampled_commits,
};
use crate::analyses::code_health::{
CloneSource, CodeHealthRow, HealthScanCtx, SMELL_WEIGHTS, run_code_health_scoped,
};
use crate::facts::FactsDb;
use crate::facts::ingest::at_rev::{ingest_complexity_at_rev, materialize_imports_at_rev};
use crate::repo::Repo;
use crate::{Options, Result};
use super::szz::SzzLink;
use super::{TuningDecision, ValidationMetrics};
const CM_AT_REV: &str = "defect_cm_at_rev";
const IMPORTS_AT_REV: &str = "defect_imports_at_rev";
pub fn band_history<R: Repo>(
db: &FactsDb,
repo: &R,
opts: &Options,
) -> Result<Vec<(String, HashMap<String, String>)>> {
let samples = sampled_commits(db)?;
let scan_opts = opts.with_no_row_limit();
let mut out = Vec::with_capacity(samples.len());
for (rev, ts) in &samples {
let date = ts.get(..10).unwrap_or(ts).to_string();
let live = live_paths_at(db, ts)?;
let graph = import_graph_from_live_paths(repo, rev, &live);
ingest_complexity_at_rev(db, repo, rev, &live, CM_AT_REV)?;
materialize_imports_at_rev(db, &graph.resolved_edges(), IMPORTS_AT_REV)?;
let cx = HealthScanCtx {
complexity_source: CM_AT_REV.to_string(),
imports_source: IMPORTS_AT_REV.to_string(),
history_cutoff: Some(ts.clone()),
include_clones: false,
clone_source: CloneSource::WorkingTree,
};
let code_rows = run_code_health_scoped(db, &scan_opts, &cx)?;
let bands: HashMap<String, String> =
code_rows.into_iter().map(|r| (r.path, r.band)).collect();
out.push((date, bands));
}
Ok(out)
}
const BAND_ORDER: [&str; 3] = ["red", "yellow", "green"];
fn band_at_defect<'a, H: std::hash::BuildHasher>(
bands: &'a [(String, HashMap<String, String, H>)],
defect_date: &str,
) -> Option<&'a HashMap<String, String, H>> {
let mut nearest = None;
for (date, map) in bands {
if date.as_str() <= defect_date {
nearest = Some(map);
} else {
break; }
}
nearest.or_else(|| bands.first().map(|(_, map)| map))
}
fn band_for_link<'a>(
link: &SzzLink,
commit_dates: &HashMap<String, String, impl std::hash::BuildHasher>,
bands: &'a [(String, HashMap<String, String, impl std::hash::BuildHasher>)],
) -> Option<&'a str> {
let defect_date = commit_dates.get(&link.defect_rev)?;
let band_map = band_at_defect(bands, defect_date)?;
band_map.get(&link.path).map(String::as_str)
}
#[must_use]
pub fn validate(
links: &[SzzLink],
commit_dates: &HashMap<String, String, impl std::hash::BuildHasher>,
bands: &[(String, HashMap<String, String, impl std::hash::BuildHasher>)],
head_health: &[CodeHealthRow],
) -> ValidationMetrics {
let labels: HashSet<&str> = links.iter().map(|l| l.path.as_str()).collect();
let linked_defects: HashSet<&str> = links.iter().map(|l| l.defect_rev.as_str()).collect();
let mut band_counts: HashMap<&str, u32> = HashMap::new();
let mut excluded_no_data = 0u32;
for link in links {
match band_for_link(link, commit_dates, bands) {
Some(band) => *band_counts.entry(band).or_insert(0) += 1,
None => excluded_no_data += 1,
}
}
let total: u32 = band_counts.values().sum();
let band_table = BAND_ORDER
.iter()
.map(|&band| {
let count = band_counts.get(band).copied().unwrap_or(0);
let share = if total > 0 {
f64::from(count) / f64::from(total)
} else {
0.0
};
(band.to_string(), count, share)
})
.collect();
let scored: Vec<(f64, bool)> = head_health
.iter()
.map(|r| (r.structural_risk, labels.contains(r.path.as_str())))
.collect();
let red_count = head_health.iter().filter(|r| r.band == "red").count();
ValidationMetrics {
band_table,
auc_default: crate::stats::auc(&scored),
precision_at_10: crate::stats::precision_at_k(&scored, 10),
precision_at_red: crate::stats::precision_at_k(&scored, red_count),
implicated_files: u32::try_from(labels.len()).unwrap_or(u32::MAX),
linked_defects: u32::try_from(linked_defects.len()).unwrap_or(u32::MAX),
sample_dates: bands.iter().map(|(date, _)| date.clone()).collect(),
excluded_no_data,
}
}
#[must_use]
pub fn default_weights() -> Vec<(String, f64)> {
crate::analyses::code_health::default_smell_weights()
}
pub fn capture_intensities(db: &FactsDb) -> Result<HashMap<String, [f64; 8]>> {
let rows: Vec<(String, String, f64)> = crate::analyses::query::query_map_collect(
db,
"SELECT path, smell, intensity FROM code_health_biomarkers_v1",
[],
"defect-calibration:capture-intensities",
|r| {
Ok((
r.get::<_, String>(0)?,
r.get::<_, String>(1)?,
r.get::<_, f64>(2)?,
))
},
)?;
let mut out: HashMap<String, [f64; 8]> = HashMap::new();
for (path, smell, intensity) in rows {
let Some(idx) = SMELL_WEIGHTS.iter().position(|&(name, _)| name == smell) else {
continue; };
out.entry(path).or_insert([0.0; 8])[idx] = intensity;
}
Ok(out)
}
#[must_use]
pub fn structural_risk_from_intensities(weights: &[(String, f64)], intensities: &[f64; 8]) -> f64 {
let raw: Vec<f64> = weights.iter().map(|(_, w)| *w).collect();
risk_raw(&raw, intensities)
}
fn risk_raw(weights: &[f64], intensities: &[f64; 8]) -> f64 {
let sum: f64 = weights
.iter()
.zip(intensities.iter())
.map(|(w, i)| w * i)
.sum();
sum.min(1.0)
}
const MIN_LINKED_DEFECTS: usize = 30;
const MIN_IMPLICATED_FILES: usize = 10;
const ACCEPTANCE_MARGIN: f64 = 0.02;
const DISCRIMINATION_FLOOR: f64 = 0.5;
const STEPS: [f64; 5] = [0.5, 0.75, 1.0, 1.25, 1.5];
const PASSES: usize = 2;
fn project_sum_to_one(weights: &mut [f64]) {
let sum: f64 = weights.iter().sum();
if sum > 0.0 {
for w in weights.iter_mut() {
*w /= sum;
}
}
}
fn auc_for(
weights: &[f64],
labels: &[(String, bool)],
intensities: &HashMap<String, [f64; 8], impl std::hash::BuildHasher>,
) -> Option<f64> {
let scored: Vec<(f64, bool)> = labels
.iter()
.filter_map(|(path, label)| {
intensities
.get(path)
.map(|ints| (risk_raw(weights, ints), *label))
})
.collect();
crate::stats::auc(&scored)
}
fn coordinate_descent(
defaults: &[f64],
train: &[(String, bool)],
intensities: &HashMap<String, [f64; 8], impl std::hash::BuildHasher>,
) -> Vec<f64> {
let mut current = defaults.to_vec();
let mut current_objective = auc_for(¤t, train, intensities).unwrap_or(0.0);
for _pass in 0..PASSES {
for i in 0..defaults.len() {
let mut best: Option<(f64, f64)> = None; for &step in &STEPS {
let mut trial = current.clone();
trial[i] = defaults[i] * step;
let objective = auc_for(&trial, train, intensities).unwrap_or(0.0);
let improves_on_best = match best {
None => true,
Some((best_objective, _)) => objective > best_objective,
};
if improves_on_best {
best = Some((objective, trial[i]));
}
}
if let Some((objective, weight)) = best
&& objective > current_objective
{
current[i] = weight;
project_sum_to_one(&mut current);
current_objective =
auc_for(¤t, train, intensities).unwrap_or(current_objective);
}
}
}
current
}
fn linked_defect_count(train: &[(String, bool)], validation: &[(String, bool)]) -> usize {
train
.iter()
.chain(validation)
.filter(|(_, label)| *label)
.count()
}
fn implicated_file_count(train: &[(String, bool)], validation: &[(String, bool)]) -> usize {
train
.iter()
.chain(validation)
.filter(|(_, label)| *label)
.map(|(path, _)| path.as_str())
.collect::<HashSet<_>>()
.len()
}
#[must_use]
pub fn tune_weights(
intensities: &HashMap<String, [f64; 8], impl std::hash::BuildHasher>,
train: &[(String, bool)],
validation: &[(String, bool)],
defaults: &[(String, f64)],
) -> (Vec<(String, f64)>, TuningDecision) {
let defaults_vec = defaults.to_vec();
let kept = |reason: &str,
auc_validation_default: Option<f64>,
auc_validation_tuned: Option<f64>|
-> (Vec<(String, f64)>, TuningDecision) {
(
defaults_vec.clone(),
TuningDecision::DefaultsKept {
reason: reason.to_string(),
auc_validation_default,
auc_validation_tuned,
},
)
};
if linked_defect_count(train, validation) < MIN_LINKED_DEFECTS {
return kept("fewer than 30 linked defect-changes", None, None);
}
if implicated_file_count(train, validation) < MIN_IMPLICATED_FILES {
return kept("fewer than 10 implicated files", None, None);
}
let default_raw: Vec<f64> = defaults.iter().map(|(_, w)| *w).collect();
let Some(auc_validation_default) = auc_for(&default_raw, validation, intensities) else {
return kept(
"validation split has no positive/negative class to score",
None,
None,
);
};
let tuned_raw = coordinate_descent(&default_raw, train, intensities);
let auc_train = auc_for(&tuned_raw, train, intensities).unwrap_or(auc_validation_default);
let auc_validation_tuned =
auc_for(&tuned_raw, validation, intensities).unwrap_or(auc_validation_default);
if auc_validation_tuned < DISCRIMINATION_FLOOR {
return kept(
"tuned weights rank below random on the validation split",
Some(auc_validation_default),
Some(auc_validation_tuned),
);
}
if auc_validation_tuned >= auc_validation_default + ACCEPTANCE_MARGIN {
let weights: Vec<(String, f64)> = defaults
.iter()
.zip(&tuned_raw)
.map(|((name, _), &w)| (name.clone(), w))
.collect();
(
weights,
TuningDecision::Applied {
auc_train,
auc_validation_default,
auc_validation_tuned,
},
)
} else {
kept(
"tuned weights did not beat the default validation AUC by the required margin",
Some(auc_validation_default),
Some(auc_validation_tuned),
)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn health_row(path: &str, structural_risk: f64, band: &str) -> CodeHealthRow {
CodeHealthRow {
path: path.to_string(),
cognitive: 0.0,
score: 100.0 * (1.0 - structural_risk),
structural_risk,
percentile: 0.0,
band: band.to_string(),
corpus_percentile: None,
beyond_corpus: false,
corpus_percentile_ci_low: None,
corpus_percentile_ci_high: None,
}
}
#[test]
fn validate_hand_computed_band_table_and_auc() {
let links = vec![
SzzLink {
defect_rev: "d1".to_string(),
fix_rev: "f1".to_string(),
path: "a.rs".to_string(),
},
SzzLink {
defect_rev: "d2".to_string(),
fix_rev: "f2".to_string(),
path: "b.rs".to_string(),
},
];
let commit_dates = HashMap::from([
("d1".to_string(), "2026-01-05".to_string()),
("d2".to_string(), "2026-03-01".to_string()),
]);
let bands = vec![
(
"2026-01-01".to_string(),
HashMap::from([
("a.rs".to_string(), "red".to_string()),
("b.rs".to_string(), "green".to_string()),
]),
),
(
"2026-02-01".to_string(),
HashMap::from([
("a.rs".to_string(), "yellow".to_string()),
("b.rs".to_string(), "green".to_string()),
]),
),
];
let head_health = vec![
health_row("a.rs", 0.9, "red"),
health_row("b.rs", 0.2, "green"),
health_row("c.rs", 0.5, "yellow"),
];
let metrics = validate(&links, &commit_dates, &bands, &head_health);
assert_eq!(
metrics.band_table,
vec![
("red".to_string(), 1, 0.5),
("yellow".to_string(), 0, 0.0),
("green".to_string(), 1, 0.5),
]
);
assert_eq!(metrics.implicated_files, 2);
assert_eq!(metrics.linked_defects, 2);
assert_eq!(metrics.excluded_no_data, 0);
assert_eq!(
metrics.sample_dates,
vec!["2026-01-01".to_string(), "2026-02-01".to_string()]
);
assert_eq!(metrics.auc_default, Some(0.5));
assert_eq!(metrics.precision_at_10, None);
assert_eq!(metrics.precision_at_red, Some(1.0));
}
#[test]
fn validate_falls_back_to_earliest_sample_when_defect_predates_all_samples() {
let links = vec![SzzLink {
defect_rev: "d1".to_string(),
fix_rev: "f1".to_string(),
path: "a.rs".to_string(),
}];
let commit_dates = HashMap::from([("d1".to_string(), "2020-01-01".to_string())]);
let bands = vec![(
"2026-01-01".to_string(),
HashMap::from([("a.rs".to_string(), "red".to_string())]),
)];
let head_health = vec![health_row("a.rs", 0.9, "red")];
let metrics = validate(&links, &commit_dates, &bands, &head_health);
assert_eq!(metrics.excluded_no_data, 0);
assert_eq!(
metrics.band_table,
vec![
("red".to_string(), 1, 1.0),
("yellow".to_string(), 0, 0.0),
("green".to_string(), 0, 0.0),
]
);
}
#[test]
fn validate_excludes_links_with_no_band_data_at_the_chosen_sample() {
let links = vec![SzzLink {
defect_rev: "d1".to_string(),
fix_rev: "f1".to_string(),
path: "missing.rs".to_string(),
}];
let commit_dates = HashMap::from([("d1".to_string(), "2026-01-05".to_string())]);
let bands = vec![(
"2026-01-01".to_string(),
HashMap::from([("a.rs".to_string(), "red".to_string())]),
)];
let head_health = vec![health_row("a.rs", 0.9, "red")];
let metrics = validate(&links, &commit_dates, &bands, &head_health);
assert_eq!(metrics.excluded_no_data, 1);
assert_eq!(metrics.band_table.iter().map(|(_, n, _)| n).sum::<u32>(), 0);
}
#[test]
fn validate_excludes_links_with_unknown_commit_date() {
let links = vec![SzzLink {
defect_rev: "unknown-rev".to_string(),
fix_rev: "f1".to_string(),
path: "a.rs".to_string(),
}];
let commit_dates: HashMap<String, String> = HashMap::new();
let bands = vec![(
"2026-01-01".to_string(),
HashMap::from([("a.rs".to_string(), "red".to_string())]),
)];
let head_health = vec![health_row("a.rs", 0.9, "red")];
let metrics = validate(&links, &commit_dates, &bands, &head_health);
assert_eq!(metrics.excluded_no_data, 1);
assert_eq!(metrics.implicated_files, 1);
assert_eq!(metrics.auc_default, None); }
#[test]
fn default_weights_matches_smell_weights_order_and_sums_to_one() {
let weights = default_weights();
assert_eq!(weights.len(), 8);
for (got, &(name, weight)) in weights.iter().zip(SMELL_WEIGHTS.iter()) {
assert_eq!(got.0, name);
assert!((got.1 - weight).abs() < 1e-12);
}
let sum: f64 = weights.iter().map(|(_, w)| w).sum();
assert!((sum - 1.0).abs() < 1e-9);
}
#[test]
fn structural_risk_from_intensities_clamps_at_one() {
let weights: Vec<(String, f64)> = vec![
("a".to_string(), 0.6),
("b".to_string(), 0.6),
("c".to_string(), 0.0),
("d".to_string(), 0.0),
("e".to_string(), 0.0),
("f".to_string(), 0.0),
("g".to_string(), 0.0),
("h".to_string(), 0.0),
];
let intensities = [1.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0];
let risk = structural_risk_from_intensities(&weights, &intensities);
assert!(
(risk - 1.0).abs() < 1e-12,
"expected clamp to 1.0, got {risk}"
);
}
fn seed_files(
dest: &mut Vec<(String, bool)>,
intensities: &mut HashMap<String, [f64; 8]>,
prefix: &str,
n: usize,
label: bool,
intensity: [f64; 8],
) {
for i in 0..n {
let path = format!("{prefix}{i}.rs");
intensities.insert(path.clone(), intensity);
dest.push((path, label));
}
}
#[test]
fn tune_weights_keeps_defaults_below_the_linked_defect_floor() {
let mut intensities = HashMap::new();
let mut train = Vec::new();
let mut validation = Vec::new();
seed_files(
&mut train,
&mut intensities,
"pos",
15,
true,
[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
);
seed_files(
&mut validation,
&mut intensities,
"posv",
5,
true,
[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
);
seed_files(&mut train, &mut intensities, "neg", 15, false, [0.0; 8]);
let defaults = default_weights();
let (weights, decision) = tune_weights(&intensities, &train, &validation, &defaults);
assert_eq!(weights, defaults);
match decision {
TuningDecision::DefaultsKept {
reason,
auc_validation_default,
auc_validation_tuned,
} => {
assert_eq!(reason, "fewer than 30 linked defect-changes");
assert_eq!(auc_validation_default, None);
assert_eq!(auc_validation_tuned, None);
}
TuningDecision::Applied { .. } => {
panic!("expected DefaultsKept(floor), got Applied")
}
}
}
#[test]
fn tune_weights_keeps_defaults_below_the_implicated_file_floor() {
let mut intensities = HashMap::new();
let mut train = Vec::new();
let mut validation = Vec::new();
for rep in 0..6 {
seed_files(
&mut train,
&mut intensities,
&format!("dup{rep}_"),
5,
true,
[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
);
}
train.clear();
intensities.clear();
let shared_paths = ["dup0.rs", "dup1.rs", "dup2.rs", "dup3.rs", "dup4.rs"];
for &path in &shared_paths {
intensities.insert(path.to_string(), [1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0]);
}
for i in 0..30 {
let path = shared_paths[i % shared_paths.len()];
train.push((path.to_string(), true));
}
seed_files(
&mut validation,
&mut intensities,
"neg",
10,
false,
[0.0; 8],
);
let defaults = default_weights();
let (weights, decision) = tune_weights(&intensities, &train, &validation, &defaults);
assert_eq!(weights, defaults);
match decision {
TuningDecision::DefaultsKept {
reason,
auc_validation_default,
auc_validation_tuned,
} => {
assert_eq!(reason, "fewer than 10 implicated files");
assert_eq!(auc_validation_default, None);
assert_eq!(auc_validation_tuned, None);
}
TuningDecision::Applied { .. } => {
panic!("expected DefaultsKept(floor), got Applied")
}
}
}
#[test]
fn tune_weights_keeps_defaults_when_margin_is_not_met() {
let mut intensities = HashMap::new();
let mut train = Vec::new();
let mut validation = Vec::new();
let uniform = [0.5; 8];
seed_files(&mut train, &mut intensities, "pos", 18, true, uniform);
seed_files(&mut train, &mut intensities, "neg", 18, false, uniform);
seed_files(&mut validation, &mut intensities, "posv", 12, true, uniform);
seed_files(
&mut validation,
&mut intensities,
"negv",
12,
false,
uniform,
);
let defaults = default_weights();
let (weights, decision) = tune_weights(&intensities, &train, &validation, &defaults);
assert_eq!(weights, defaults);
match decision {
TuningDecision::DefaultsKept {
reason,
auc_validation_default,
auc_validation_tuned,
} => {
assert_eq!(
reason,
"tuned weights did not beat the default validation AUC by the required margin"
);
assert_eq!(auc_validation_default, Some(0.5));
assert_eq!(auc_validation_tuned, Some(0.5));
}
TuningDecision::Applied { .. } => {
panic!("expected DefaultsKept(margin), got Applied")
}
}
}
#[test]
fn tune_weights_applies_when_a_smell_perfectly_separates_and_clears_the_margin() {
let mut intensities = HashMap::new();
let mut train = Vec::new();
let mut validation = Vec::new();
let positive = [1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0];
let confuser_negative = [0.0, 1.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0];
let plain_negative = [0.0; 8];
seed_files(&mut train, &mut intensities, "pos", 18, true, positive);
seed_files(
&mut validation,
&mut intensities,
"posv",
12,
true,
positive,
);
seed_files(
&mut train,
&mut intensities,
"confuser",
9,
false,
confuser_negative,
);
seed_files(
&mut validation,
&mut intensities,
"confuserv",
6,
false,
confuser_negative,
);
seed_files(
&mut train,
&mut intensities,
"neg",
9,
false,
plain_negative,
);
seed_files(
&mut validation,
&mut intensities,
"negv",
6,
false,
plain_negative,
);
let defaults = default_weights();
let (weights, decision) = tune_weights(&intensities, &train, &validation, &defaults);
match decision {
TuningDecision::Applied {
auc_validation_default,
auc_validation_tuned,
..
} => {
assert!(
auc_validation_tuned >= auc_validation_default + ACCEPTANCE_MARGIN,
"tuned {auc_validation_tuned} must clear default {auc_validation_default} \
by at least {ACCEPTANCE_MARGIN}"
);
assert!(auc_validation_default < 0.9, "expected a real inversion");
}
TuningDecision::DefaultsKept { reason, .. } => {
panic!("expected Applied, got DefaultsKept({reason})")
}
}
assert_ne!(weights, defaults);
let tuned_complex_method = weights[0].1;
let default_complex_method = defaults[0].1;
assert!(
tuned_complex_method > default_complex_method,
"expected complex-method weight to increase: tuned={tuned_complex_method} \
default={default_complex_method}"
);
let sum: f64 = weights.iter().map(|(_, w)| w).sum();
assert!(
(sum - 1.0).abs() < 1e-9,
"tuned weights must sum to 1.0, got {sum}"
);
}
#[test]
fn coordinate_descent_steps_from_the_default_weight_not_the_current_one() {
let defaults = [0.5, 0.3, 0.2];
let mut train = vec![("neg_a0".to_string(), false)];
let mut intensities = HashMap::from([(
"neg_a0".to_string(),
[0.5, 0.4, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
)]);
for i in 0..2 {
let path = format!("pos{i}");
intensities.insert(path.clone(), [0.0, 0.1, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0]);
train.push((path, true));
}
for i in 0..3 {
let path = format!("neg_b{i}");
intensities.insert(path.clone(), [0.3, 0.5, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0]);
train.push((path, false));
}
let tuned = coordinate_descent(&defaults, &train, &intensities);
let expected = [4.0 / 9.0, 0.2, 16.0 / 45.0];
for (got, want) in tuned.iter().zip(expected.iter()) {
assert!(
(got - want).abs() < 1e-9,
"coordinate_descent must step from the DEFAULT weight, not the current one \
— expected {expected:?}, got {tuned:?}"
);
}
}
#[test]
fn tune_weights_rejects_a_tuned_set_that_only_improves_training_auc() {
let mut intensities = HashMap::new();
let mut train = Vec::new();
let mut validation = Vec::new();
let positive = [1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0];
let confuser_negative = [0.0, 1.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0];
let plain_negative = [0.0; 8];
seed_files(&mut train, &mut intensities, "pos", 18, true, positive);
seed_files(
&mut train,
&mut intensities,
"confuser",
9,
false,
confuser_negative,
);
seed_files(
&mut train,
&mut intensities,
"neg",
9,
false,
plain_negative,
);
let uniform_validation = [0.0; 8];
seed_files(
&mut validation,
&mut intensities,
"posv",
15,
true,
uniform_validation,
);
seed_files(
&mut validation,
&mut intensities,
"negv",
15,
false,
uniform_validation,
);
let defaults = default_weights();
let default_raw: Vec<f64> = defaults.iter().map(|(_, w)| *w).collect();
let tuned_raw = coordinate_descent(&default_raw, &train, &intensities);
let auc_train_default =
auc_for(&default_raw, &train, &intensities).expect("train auc (default)");
let auc_train_tuned = auc_for(&tuned_raw, &train, &intensities).expect("train auc (tuned)");
assert!(
auc_train_tuned >= auc_train_default + ACCEPTANCE_MARGIN,
"fixture must give the search real training-AUC headroom to justify this test: \
default={auc_train_default} tuned={auc_train_tuned}"
);
let (weights, decision) = tune_weights(&intensities, &train, &validation, &defaults);
assert_eq!(
weights, defaults,
"a training-only AUC improvement must never be adopted over the defaults"
);
match decision {
TuningDecision::DefaultsKept {
reason,
auc_validation_default,
auc_validation_tuned,
} => {
assert_eq!(
reason,
"tuned weights did not beat the default validation AUC by the required margin"
);
assert_eq!(auc_validation_default, Some(0.5));
assert_eq!(auc_validation_tuned, Some(0.5));
}
TuningDecision::Applied { .. } => panic!(
"expected DefaultsKept(margin): a training-only AUC gain must not be adopted"
),
}
}
#[test]
fn tune_weights_keeps_defaults_when_tuned_validation_auc_is_below_random() {
let mut intensities = HashMap::new();
let mut train = Vec::new();
let mut validation = Vec::new();
seed_files(
&mut train,
&mut intensities,
"tpos",
30,
true,
[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0],
);
seed_files(
&mut train,
&mut intensities,
"tneg",
10,
false,
[0.5, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
);
seed_files(
&mut validation,
&mut intensities,
"vpos",
1,
true,
[0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
);
seed_files(
&mut validation,
&mut intensities,
"vnegw0",
1,
false,
[0.87, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
);
seed_files(
&mut validation,
&mut intensities,
"vneghi",
6,
false,
[0.0, 0.0, 0.61, 0.0, 0.61, 0.61, 0.0, 0.0],
);
seed_files(
&mut validation,
&mut intensities,
"vneglo",
3,
false,
[0.0, 0.0, 0.2, 0.0, 0.2, 0.2, 0.0, 0.0],
);
let defaults = default_weights();
let (weights, decision) = tune_weights(&intensities, &train, &validation, &defaults);
assert_eq!(
weights, defaults,
"a below-random tuning must never replace the defaults"
);
match decision {
TuningDecision::DefaultsKept {
reason,
auc_validation_default,
auc_validation_tuned,
} => {
assert!(
reason.contains("below random"),
"the discrimination floor, not another branch, must fire: {reason}"
);
let d = auc_validation_default.expect("default validation AUC recorded");
let t = auc_validation_tuned.expect("tuned validation AUC recorded");
assert!((d - 0.30).abs() < 1e-9, "default validation AUC, got {d}");
assert!((t - 0.40).abs() < 1e-9, "tuned validation AUC, got {t}");
}
other @ TuningDecision::Applied { .. } => {
panic!("expected DefaultsKept via the discrimination floor, got {other:?}")
}
}
}
}