1use cerno_core::{Engine, EngineError, labels};
19use cerno_host::HostKind;
20use cerno_types::{Answer, Calibration, ChoiceSpec, Question, QuestionKind, ScoreSpec};
21use serde::Deserialize;
22use std::collections::BTreeMap;
23use std::time::{Duration, Instant};
24
25#[derive(Debug, Deserialize)]
26struct Dataset {
27 cases: Vec<Case>,
28}
29
30#[derive(Debug, Deserialize)]
31struct Case {
32 id: String,
33 state: String,
34 #[serde(default)]
35 noul: Option<String>,
36 #[serde(default)]
37 expect_yes: Option<bool>,
38 #[serde(default)]
39 choice: Option<ChoiceSpec>,
40 #[serde(default)]
41 score: Option<ScoreSpec>,
42 #[serde(default)]
44 expect: Option<usize>,
45 #[serde(default)]
46 tolerance: Option<usize>,
47}
48
49impl Case {
50 fn problem(&self) -> Option<String> {
54 let named = [
55 self.noul.is_some(),
56 self.choice.is_some(),
57 self.score.is_some(),
58 ];
59 if named.iter().filter(|n| **n).count() != 1 {
60 return Some("needs exactly one of noul, choice or score".into());
61 }
62
63 if self.noul.is_some() {
64 return self
65 .expect_yes
66 .is_none()
67 .then(|| "a noul case needs expect_yes".into());
68 }
69
70 let (offered, first) = match (&self.choice, &self.score) {
71 (Some(choice), _) => (choice.options.len(), 0),
72 (_, Some(score)) => (score.levels.count(), 1),
73 _ => unreachable!("exactly one primitive, and it is not a noul"),
74 };
75 if offered == 0 {
76 return Some(format!("this {} offers nothing to pick", self.primitive()));
77 }
78 let last = offered + first - 1;
79 match self.expect {
80 None => Some(format!("a {} case needs expect", self.primitive())),
81 Some(want) if want < first || want > last => Some(format!(
82 "expect is {want}, outside {first}..={last} for this {}",
83 self.primitive()
84 )),
85 Some(_) => None,
86 }
87 }
88
89 fn primitive(&self) -> &'static str {
90 if self.noul.is_some() {
91 "noul"
92 } else if self.choice.is_some() {
93 "choice"
94 } else {
95 "score"
96 }
97 }
98
99 fn question(&self) -> Question {
100 let kind = if let Some(q) = &self.noul {
101 QuestionKind::Noul(q.clone())
102 } else if let Some(c) = &self.choice {
103 QuestionKind::Choice(c.clone())
104 } else {
105 QuestionKind::Score(self.score.clone().expect("case has no primitive"))
106 };
107 Question {
108 id: self.id.clone(),
109 kind,
110 }
111 }
112
113 fn correct_label(&self) -> String {
115 let index = if self.noul.is_some() {
116 if self.expect_yes.expect("noul case needs expect_yes") {
118 0
119 } else {
120 1
121 }
122 } else if self.choice.is_some() {
123 self.expect.expect("choice case needs expect")
124 } else {
125 self.expect.expect("score case needs expect") - 1
126 };
127 labels::label(index)
128 .expect("expectation within the alphabet")
129 .to_string()
130 }
131
132 fn tolerance(&self) -> usize {
134 self.tolerance.unwrap_or(0)
135 }
136}
137
138struct Outcome {
140 primitive: &'static str,
141 answer: Option<Answer>,
143 faithful: bool,
147 correct: bool,
148 logprobs: BTreeMap<String, f64>,
150 correct_label: String,
151 latency: Duration,
152 truncated: bool,
153 picked: Option<String>,
155}
156
157fn softmax_of(logprobs: &BTreeMap<String, f64>, temperature: f64) -> BTreeMap<String, f64> {
158 let keys: Vec<&String> = logprobs.keys().collect();
159 let values: Vec<f64> = keys.iter().map(|k| logprobs[*k]).collect();
160 let probabilities = cerno_core::math::softmax(&values, temperature);
161 keys.into_iter().cloned().zip(probabilities).collect()
162}
163
164async fn run_case(engine: &Engine, model: &str, case: &Case) -> Outcome {
165 let question = case.question();
166 let correct_label = case.correct_label();
167 let started = Instant::now();
168
169 let result = engine
170 .answer_with_distribution(&case.state, &question, model, Calibration::default())
171 .await;
172 let latency = started.elapsed();
173
174 match result {
175 Ok((answer, distribution)) => {
176 let logprobs = match &answer {
177 Answer::Noul { raw_logprobs, .. }
178 | Answer::Choice { raw_logprobs, .. }
179 | Answer::Score { raw_logprobs, .. } => raw_logprobs.clone(),
180 };
181 let faithful = top_token_is_a_label(&distribution.tokens, logprobs.keys());
183
184 let (picked_index, correct) = match &answer {
185 Answer::Noul { noul, .. } => {
186 let yes = *noul >= 0.5;
187 (if yes { 0 } else { 1 }, yes == case.expect_yes.unwrap())
188 }
189 Answer::Choice { index, .. } => (*index, *index == case.expect.unwrap()),
190 Answer::Score { score, .. } => {
191 let got = *score as usize;
192 let want = case.expect.unwrap();
193 (got - 1, got.abs_diff(want) <= case.tolerance())
194 }
195 };
196
197 Outcome {
198 primitive: case.primitive(),
199 truncated: answer.truncated(),
200 answer: Some(answer),
201 faithful,
202 correct,
203 logprobs,
204 correct_label,
205 latency,
206 picked: labels::label(picked_index).map(str::to_string),
207 }
208 }
209 Err(err) => {
210 eprintln!(" {} {}: {}", case.id, case.primitive(), err);
211 Outcome {
212 primitive: case.primitive(),
213 answer: None,
214 faithful: false,
215 correct: false,
216 logprobs: BTreeMap::new(),
217 correct_label,
218 latency,
219 truncated: false,
220 picked: None,
221 }
222 }
223 }
224}
225
226fn top_token_is_a_label<'a>(
228 tokens: &[(String, f64)],
229 mut offered: impl Iterator<Item = &'a String>,
230) -> bool {
231 tokens
232 .first()
233 .is_some_and(|(top, _)| offered.any(|label| labels::matches(top, label)))
234}
235
236fn percentile(sorted: &[u128], p: f64) -> u128 {
237 if sorted.is_empty() {
238 return 0;
239 }
240 let rank = ((sorted.len() - 1) as f64 * p).round() as usize;
241 sorted[rank]
242}
243
244fn fit_temperature(outcomes: &[Outcome]) -> (f64, f64) {
250 let usable: Vec<&Outcome> = outcomes.iter().filter(|o| !o.logprobs.is_empty()).collect();
251 if usable.is_empty() {
252 return (1.0, f64::NAN);
253 }
254
255 let mut best = (1.0, f64::INFINITY);
256 let mut t = 0.25;
257 while t <= 8.0 {
258 let nll: f64 = usable
259 .iter()
260 .map(|o| {
261 let p = softmax_of(&o.logprobs, t)
262 .get(&o.correct_label)
263 .copied()
264 .unwrap_or(1e-12);
265 -p.max(1e-12).ln()
266 })
267 .sum::<f64>()
268 / usable.len() as f64;
269
270 if nll < best.1 {
271 best = (t, nll);
272 }
273 t += 0.05;
274 }
275 ((best.0 * 100.0).round() / 100.0, best.1)
276}
277
278fn brier(outcomes: &[Outcome], cases: &[Case]) -> f64 {
281 let mut total = 0.0;
282 let mut n = 0;
283 for (outcome, case) in outcomes.iter().zip(cases) {
284 if let Some(Answer::Noul { noul, .. }) = &outcome.answer {
285 let truth = if case.expect_yes.unwrap() { 1.0 } else { 0.0 };
286 total += (noul - truth).powi(2);
287 n += 1;
288 }
289 }
290 if n == 0 { f64::NAN } else { total / n as f64 }
291}
292
293struct Report {
294 model: String,
295 fidelity: f64,
296 accuracy: f64,
297 per_primitive: BTreeMap<&'static str, f64>,
298 p50: u128,
299 p99: u128,
300 truncation: f64,
301 brier: f64,
302 best_t: f64,
303 agreement: Option<f64>,
304 picked: Vec<Option<String>>,
305}
306
307fn summarise(model: &str, outcomes: &[Outcome], cases: &[Case]) -> Report {
308 let n = outcomes.len() as f64;
309
310 let mut latencies: Vec<u128> = outcomes.iter().map(|o| o.latency.as_millis()).collect();
311 latencies.sort_unstable();
312
313 let mut per_primitive = BTreeMap::new();
314 for primitive in ["noul", "choice", "score"] {
315 let subset: Vec<&Outcome> = outcomes
316 .iter()
317 .filter(|o| o.primitive == primitive)
318 .collect();
319 if !subset.is_empty() {
320 let hits = subset.iter().filter(|o| o.correct).count() as f64;
321 per_primitive.insert(primitive, hits / subset.len() as f64);
322 }
323 }
324
325 let (best_t, _) = fit_temperature(outcomes);
326
327 Report {
328 model: model.to_string(),
329 fidelity: outcomes.iter().filter(|o| o.faithful).count() as f64 / n,
330 accuracy: outcomes.iter().filter(|o| o.correct).count() as f64 / n,
331 per_primitive,
332 p50: percentile(&latencies, 0.50),
333 p99: percentile(&latencies, 0.99),
334 truncation: outcomes.iter().filter(|o| o.truncated).count() as f64 / n,
335 brier: brier(outcomes, cases),
336 best_t,
337 agreement: None,
338 picked: outcomes.iter().map(|o| o.picked.clone()).collect(),
339 }
340}
341
342fn arg(args: &[String], name: &str) -> Option<String> {
343 args.iter()
344 .position(|a| a == name)
345 .and_then(|i| args.get(i + 1))
346 .cloned()
347}
348
349#[tokio::main(flavor = "current_thread")]
350async fn main() -> Result<(), Box<dyn std::error::Error>> {
351 let args: Vec<String> = std::env::args().collect();
352
353 let models: Vec<String> = arg(&args, "--models")
354 .unwrap_or_else(|| "gemma4:e2b-it-qat,granite4:3b,phi4-mini:3.8b".into())
355 .split(',')
356 .map(|s| s.trim().to_string())
357 .filter(|s| !s.is_empty())
358 .collect();
359 let reference = arg(&args, "--reference");
360 let dataset_path =
361 arg(&args, "--dataset").unwrap_or_else(|| "crates/cerno-bench/dataset.json".into());
362 let out_path = arg(&args, "--out").unwrap_or_else(|| "docs/model-selection.md".into());
363 let host_kind: HostKind = arg(&args, "--host-kind")
364 .unwrap_or_else(|| "ollama".into())
365 .parse()?;
366 let host_url = arg(&args, "--host").unwrap_or_else(|| host_kind.default_url().into());
367
368 let dataset: Dataset = serde_json::from_str(&std::fs::read_to_string(&dataset_path)?)?;
369 let problems: Vec<String> = dataset
370 .cases
371 .iter()
372 .filter_map(|case| case.problem().map(|p| format!(" {}: {p}", case.id)))
373 .collect();
374 if !problems.is_empty() {
375 eprintln!(
376 "{dataset_path} has cases that cannot run:\n{}",
377 problems.join("\n")
378 );
379 std::process::exit(1);
380 }
381 println!("{} cases from {dataset_path}", dataset.cases.len());
382
383 let host = cerno_host::connect(
384 host_kind,
385 &host_url,
386 std::env::var("CERNO_HOST_API_KEY").ok(),
387 Duration::from_secs(120),
388 )?;
389 let engine = Engine::new(host, Some("5m".to_string()));
392
393 let mut reports = Vec::new();
394 let mut reference_picks: Option<Vec<Option<String>>> = None;
395
396 let all: Vec<String> = reference
397 .iter()
398 .cloned()
399 .chain(models.iter().cloned())
400 .collect();
401
402 for model in &all {
403 println!("\n{model}");
404
405 let warm_up = engine
411 .answer(
412 "warm up",
413 &Question {
414 id: "warmup".into(),
415 kind: QuestionKind::Noul("Is this a warm-up?".into()),
416 },
417 model,
418 Calibration::default(),
419 )
420 .await;
421 if let Err(err @ (EngineError::Host(_) | EngineError::UnknownModel { .. })) = warm_up {
422 eprintln!("{model}: the warm-up failed, so nothing was measured or written: {err}");
423 std::process::exit(1);
424 }
425
426 let mut outcomes = Vec::with_capacity(dataset.cases.len());
427 for case in &dataset.cases {
428 outcomes.push(run_case(&engine, model, case).await);
429 }
430
431 let mut report = summarise(model, &outcomes, &dataset.cases);
432
433 if reference.as_deref() == Some(model.as_str()) {
434 reference_picks = Some(report.picked.clone());
435 } else if let Some(reference_picks) = &reference_picks {
436 let agreed = report
437 .picked
438 .iter()
439 .zip(reference_picks)
440 .filter(|(a, b)| a.is_some() && a == b)
441 .count();
442 report.agreement = Some(agreed as f64 / dataset.cases.len() as f64);
443 }
444
445 println!(
446 " fidelity {:.0}% accuracy {:.0}% p50 {}ms p99 {}ms best T {:.2}",
447 report.fidelity * 100.0,
448 report.accuracy * 100.0,
449 report.p50,
450 report.p99,
451 report.best_t
452 );
453 reports.push(report);
454 }
455
456 let markdown = render(&reports, &dataset, reference.as_deref());
457 if let Some(parent) = std::path::Path::new(&out_path).parent() {
458 std::fs::create_dir_all(parent)?;
459 }
460 std::fs::write(&out_path, &markdown)?;
461 println!("\nwrote {out_path}");
462
463 Ok(())
464}
465
466fn pct(v: f64) -> String {
467 if v.is_nan() {
468 "—".into()
469 } else {
470 format!("{:.0}%", v * 100.0)
471 }
472}
473
474fn pick_winner<'a>(reports: &'a [Report], reference: Option<&str>) -> Option<&'a Report> {
479 reports
480 .iter()
481 .filter(|r| Some(r.model.as_str()) != reference)
482 .filter(|r| r.fidelity >= 1.0)
483 .min_by(|a, b| {
484 b.accuracy
485 .total_cmp(&a.accuracy)
486 .then_with(|| a.p50.cmp(&b.p50))
487 })
488}
489
490fn markdown_table(headers: &[&str], rows: &[Vec<String>]) -> String {
496 let widths: Vec<usize> = headers
499 .iter()
500 .enumerate()
501 .map(|(i, header)| {
502 rows.iter()
503 .filter_map(|row| row.get(i))
504 .map(|cell| cell.chars().count())
505 .chain(std::iter::once(header.chars().count()))
506 .max()
507 .unwrap_or(0)
508 })
509 .collect();
510
511 let pad = |cell: &str, width: usize, right: bool| {
512 let fill = " ".repeat(width.saturating_sub(cell.chars().count()));
513 if right {
514 format!("{fill}{cell}")
515 } else {
516 format!("{cell}{fill}")
517 }
518 };
519
520 let mut out = String::new();
521
522 let header: Vec<String> = headers
523 .iter()
524 .zip(&widths)
525 .enumerate()
526 .map(|(i, (h, w))| pad(h, *w, i > 0))
527 .collect();
528 out.push_str(&format!("| {} |\n", header.join(" | ")));
529
530 let rule: Vec<String> = widths
533 .iter()
534 .enumerate()
535 .map(|(i, w)| {
536 let dashes = "-".repeat(w + 1);
537 if i == 0 {
538 format!(":{dashes}")
539 } else {
540 format!("{dashes}:")
541 }
542 })
543 .collect();
544 out.push_str(&format!("|{}|\n", rule.join("|")));
545
546 for row in rows {
547 let cells: Vec<String> = row
548 .iter()
549 .zip(&widths)
550 .enumerate()
551 .map(|(i, (cell, w))| pad(cell, *w, i > 0))
552 .collect();
553 out.push_str(&format!("| {} |\n", cells.join(" | ")));
554 }
555
556 out
557}
558
559fn calibration_reading(best_t: f64) -> &'static str {
564 if (best_t - 1.0).abs() <= 0.1 {
566 "that is close enough to 1 that the model's own probabilities can be taken as they are"
567 } else if best_t < 1.0 {
568 "a value below 1 means the model is *under*confident and would need sharpening rather \
569 than flattening"
570 } else {
571 "a value above 1 means the model is *over*confident and its probabilities need \
572 flattening before they mean anything"
573 }
574}
575
576fn render(reports: &[Report], dataset: &Dataset, reference: Option<&str>) -> String {
577 let mut out = String::new();
578 out.push_str("# Model selection\n\n");
579 out.push_str(&format!(
580 "Generated by `cargo run -p cerno-bench`. {} labelled cases: {} noul, {} choice, {} score.\n\n",
581 dataset.cases.len(),
582 dataset.cases.iter().filter(|c| c.noul.is_some()).count(),
583 dataset.cases.iter().filter(|c| c.choice.is_some()).count(),
584 dataset.cases.iter().filter(|c| c.score.is_some()).count(),
585 ));
586
587 let headers = [
588 "Model",
589 "Fidelity",
590 "Accuracy",
591 "noul",
592 "choice",
593 "score",
594 "p50 (ms)",
595 "p99 (ms)",
596 "Truncated",
597 "Brier",
598 "Best T",
599 "Agreement",
600 ];
601
602 let mut reference_seen = false;
603 let rows: Vec<Vec<String>> = reports
604 .iter()
605 .map(|r| {
606 let is_reference = Some(r.model.as_str()) == reference;
611 reference_seen |= is_reference;
612
613 vec![
614 format!("`{}`{}", r.model, if is_reference { " ¹" } else { "" }),
615 pct(r.fidelity),
616 pct(r.accuracy),
617 pct(r.per_primitive.get("noul").copied().unwrap_or(f64::NAN)),
618 pct(r.per_primitive.get("choice").copied().unwrap_or(f64::NAN)),
619 pct(r.per_primitive.get("score").copied().unwrap_or(f64::NAN)),
620 r.p50.to_string(),
621 r.p99.to_string(),
622 pct(r.truncation),
623 format!("{:.3}", r.brier),
624 format!("{:.2}", r.best_t),
625 r.agreement.map(pct).unwrap_or_else(|| "—".into()),
626 ]
627 })
628 .collect();
629
630 out.push_str(&markdown_table(&headers, &rows));
631
632 if reference_seen {
633 out.push_str(
634 "\n¹ Reference model — the yardstick the Agreement column is measured against, \
635 not a candidate.\n",
636 );
637 }
638
639 if let Some(winner) = pick_winner(reports, reference) {
640 out.push_str(&format!(
641 "\n## Verdict\n\n\
642 **`{}`** — the most accurate model whose most likely token was an offered letter \
643 in every case, ties broken by median latency.\n\n\
644 Set it as the default with `CERNO_DEFAULT_MODEL={}`. Its best-fit calibration \
645 temperature above is {:.2}; {}. cerno ships a default of 1.0 either way — {} \
646 labelled cases is thin evidence for baking in an adjustment, and callers can always \
647 pass their own.\n",
648 winner.model,
649 winner.model,
650 winner.best_t,
651 calibration_reading(winner.best_t),
652 dataset.cases.len(),
653 ));
654 }
655
656 out.push_str(
657 "\n## Reading the table\n\n\
658 - **Fidelity** — share of cases whose most likely first token was one of the offered\n \
659 letters. This is a gate, not a score: below 100% the model is ignoring the instruction\n \
660 on some inputs, and cerno then reads its answer off tokens it was not going to write.\n\
661 - **Accuracy** — share of cases where the picked label was the labelled one. Score cases\n \
662 allow the per-case tolerance in the dataset.\n\
663 - **p50 / p99** — wall-clock per question, warm model, one question per request.\n\
664 - **Truncated** — share of answers where some label fell outside the host's top-20 window,\n \
665 making its probability an upper bound rather than an observation.\n\
666 - **Brier** — mean squared error of the yes-probability on noul cases; lower is better.\n\
667 - **Best T** — the calibration temperature that best explains the labelled outcomes. Far\n \
668 above 1 means the model is overconfident and needs flattening.\n\
669 - **Agreement** — share of cases where the model picked the same label as the reference.\n",
670 );
671
672 out
673}
674
675#[cfg(test)]
676mod tests {
677 use super::*;
678
679 fn case(json: serde_json::Value) -> Case {
682 let mut base = serde_json::json!({"id": "c", "state": "s"});
683 base.as_object_mut()
684 .unwrap()
685 .extend(json.as_object().unwrap().clone());
686 serde_json::from_value(base).unwrap()
687 }
688
689 #[test]
691 fn every_shipped_case_can_run() {
692 let text = std::fs::read_to_string(
693 std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join("dataset.json"),
694 )
695 .unwrap();
696 let dataset: Dataset = serde_json::from_str(&text).unwrap();
697
698 for case in &dataset.cases {
699 assert_eq!(case.problem(), None, "{}", case.id);
700 }
701 }
702
703 #[test]
704 fn a_case_that_could_not_be_scored_is_found_before_running() {
705 let bad = [
706 serde_json::json!({}),
707 serde_json::json!({"noul": "a?", "choice": {"options": ["x", "y"]}, "expect_yes": true}),
708 serde_json::json!({"noul": "a?"}),
709 serde_json::json!({"choice": {"options": ["x", "y"]}}),
710 serde_json::json!({"choice": {"options": ["x", "y"]}, "expect": 2}),
711 serde_json::json!({"choice": {"options": []}, "expect": 0}),
712 serde_json::json!({"score": {"levels": 5}, "expect": 0}),
714 serde_json::json!({"score": {"levels": 5}, "expect": 6}),
715 ];
716 for json in bad {
717 assert!(
718 case(json.clone()).problem().is_some(),
719 "{json} was accepted"
720 );
721 }
722
723 for good in [
724 serde_json::json!({"noul": "a?", "expect_yes": false}),
725 serde_json::json!({"choice": {"options": ["x", "y"]}, "expect": 1}),
726 serde_json::json!({"score": {"levels": 5}, "expect": 5}),
727 ] {
728 assert_eq!(case(good.clone()).problem(), None, "{good}");
729 }
730 }
731
732 #[test]
733 fn table_columns_line_up_exactly() {
734 let table = markdown_table(
735 &["Model", "p50 (ms)"],
736 &[
737 vec!["`gemma4:e2b-it-qat`".into(), "37".into()],
738 vec!["`x`".into(), "1234".into()],
739 ],
740 );
741
742 assert_eq!(
743 table,
744 "\
745| Model | p50 (ms) |
746|:--------------------|---------:|
747| `gemma4:e2b-it-qat` | 37 |
748| `x` | 1234 |
749"
750 );
751 }
752
753 #[test]
754 fn every_line_of_a_table_is_the_same_width() {
755 let table = markdown_table(
756 &["Model", "Brier", "Agreement"],
757 &[
758 vec![
759 "`gemma4:26b-a4b-it-q4_K_M` ¹".into(),
760 "0.029".into(),
761 "—".into(),
762 ],
763 vec!["`granite4:3b`".into(), "0.168".into(), "75%".into()],
764 ],
765 );
766
767 let widths: Vec<usize> = table.lines().map(|l| l.chars().count()).collect();
768
769 assert!(
770 widths.windows(2).all(|w| w[0] == w[1]),
771 "ragged table, line widths {widths:?}:\n{table}"
772 );
773 }
774
775 #[test]
778 fn only_a_label_at_the_top_is_faithful() {
779 let offered = ["A".to_string(), "B".to_string()];
780 let ranked = |tokens: &[(&str, f64)]| -> Vec<(String, f64)> {
781 tokens.iter().map(|(t, l)| (t.to_string(), *l)).collect()
782 };
783
784 assert!(top_token_is_a_label(
785 &ranked(&[(" b", -0.1), ("A", -3.0)]),
786 offered.iter()
787 ));
788 assert!(!top_token_is_a_label(
789 &ranked(&[("**", -0.01), ("A", -8.0)]),
790 offered.iter()
791 ));
792 assert!(!top_token_is_a_label(&[], offered.iter()));
793 }
794
795 #[test]
797 fn the_calibration_reading_follows_the_temperature() {
798 assert!(calibration_reading(0.8).contains("*under*confident"));
799 assert!(calibration_reading(3.2).contains("*over*confident"));
800 assert!(calibration_reading(1.05).contains("close enough to 1"));
801 }
802
803 #[test]
804 fn a_long_cell_widens_its_column() {
805 let table = markdown_table(&["T"], &[vec!["a-very-long-value".into()]]);
806
807 assert!(table.starts_with("| T |\n"), "{table}");
808 }
809}