diff --git a/src/bin/alpha_id.rs b/src/bin/alpha_id.rs index ad12468..af2cb0a 100644 --- a/src/bin/alpha_id.rs +++ b/src/bin/alpha_id.rs @@ -1,4 +1,3 @@ -use alpha_id::eval::{inject_errors, recall_at_k}; use alpha_id::model::{Config, Filter, Mode}; use alpha_id::pipeline::Pipeline; use clap::{Parser, Subcommand}; @@ -36,6 +35,27 @@ enum Cmd { }, } +fn run_eval(p: &Pipeline, m: Mode, cases: usize, seed: u64) -> (f64, f64) { + use alpha_id::eval::{inject_errors, recall_at_k}; + use alpha_id::model::Filter; + let billable: Vec<_> = p.entries.iter().enumerate() + .filter(|(_, e)| e.valid) + .filter(|(_, e)| alpha_id::corpus::primary_code(e) + .and_then(|c| p.meta.get(c)) + .map(|mt| mt.para295 == "P").unwrap_or(false)) + .take(cases).collect(); + let mut sum = 0.0; + for (k, (_, e)) in billable.iter().enumerate() { + let code = alpha_id::corpus::primary_code(e).unwrap().to_string(); + let noisy = inject_errors(&e.text, 0.15, seed.wrapping_add(k as u64)); + let res = p.suggest(&noisy, m, &Filter::default(), 10); + let preds: Vec = res.suggestions.iter().map(|s| s.icd_code.clone()).collect(); + sum += recall_at_k(&preds, &code, 10); + } + let n = billable.len().max(1) as f64; + (sum / n, sum / n) +} + fn read_input(arg: &str) -> String { if let Ok(v) = std::env::var("ALPHA_ID_STDIN") { return v; } if arg == "-" { @@ -115,31 +135,17 @@ fn main() { } } Cmd::Eval { mode, cases, seed } => { + let modes: Vec = if mode == "both" { + vec![Mode::C, Mode::A] + } else { + vec![mode.parse().expect("mode")] + }; let p = Pipeline::load(&cfg).expect("load"); - let m: Mode = mode.parse().expect("mode"); - let mut sum = 0.0; - let mut billable_sum = 0.0; - let mut billable_n: f64 = 0.0; - let billable: Vec<_> = p.entries.iter().enumerate() - .filter(|(_, e)| e.valid) - .filter(|(_, e)| { - alpha_id::corpus::primary_code(e) - .and_then(|c| p.meta.get(c)) - .map(|mt| mt.para295 == "P").unwrap_or(false) - }) - .take(cases).collect(); - for (k, (_, e)) in billable.iter().enumerate() { - let code = alpha_id::corpus::primary_code(e).unwrap().to_string(); - let noisy = inject_errors(&e.text, 0.15, seed.wrapping_add(k as u64)); - let res = p.suggest(&noisy, m, &Filter::default(), 10); - let preds: Vec = res.suggestions.iter().map(|s| s.icd_code.clone()).collect(); - sum += recall_at_k(&preds, &code, 10); - billable_sum += recall_at_k(&preds, &code, 10); - billable_n += 1.0; + for m in modes { + let (r10, br10) = run_eval(&p, m, cases, seed); + println!("mode={:?} cases={} Recall@10={:.3} BillableRecall@10={:.3}", + m, cases, r10, br10); } - let n = billable.len().max(1) as f64; - println!("mode={} cases={} Recall@10={:.3} BillableRecall@10={:.3}", - mode, billable.len(), sum / n, billable_sum / billable_n.max(1.0)); } } } diff --git a/tests/eval_compare_tests.rs b/tests/eval_compare_tests.rs new file mode 100644 index 0000000..bbf212d --- /dev/null +++ b/tests/eval_compare_tests.rs @@ -0,0 +1,12 @@ +use std::process::Command; + +#[test] +fn eval_mode_c_runs_and_reports_billable_recall() { + let out = Command::new(env!("CARGO_BIN_EXE_alpha-id")) + .args(["eval", "--mode", "c", "--cases", "30"]) + .arg("--config").arg("config/default.toml") + .output().unwrap(); + assert!(out.status.success(), "stderr: {}", String::from_utf8_lossy(&out.stderr)); + let s = String::from_utf8_lossy(&out.stdout); + assert!(s.contains("BillableRecall@10=")); +}