Skip to content

Eval Metrics — Rust source

The standard LLM-eval numbers with exact math — unbiased pass@k, precision/recall/F1 from confusion counts, exact-match and label-set micro-F1. 100% client-side.

This is the Rust implementation — the same logic the interactive tool runs, in a shareable, citable form.

//! Eval Metrics — the standard LLM-eval metrics with exact, testable formulas.
//!
//! Language: Rust (edition 2021, standard library only)
//! Port of src/lib/evalMetrics.ts (the canonical TypeScript implementation).
//! Tool page: https://dev.cosmolabs.org/tools/eval-metrics
//!
//! pass@k is the unbiased estimator from the Codex paper (Chen et al. 2021),
//! identical to the HumanEval implementation's combinatorial form. Precision,
//! recall, and F1 run over confusion counts and over label sets; exact-match
//! rate runs over paired strings (case-sensitive).

/// Precision / recall / F1 triple.
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Prf1 {
    pub precision: f64,
    pub recall: f64,
    pub f1: f64,
}

/// Confusion counts; `tn` is accepted but unused by P/R/F1.
#[derive(Debug, Clone, Copy)]
pub struct ConfusionCounts {
    pub tp: i64,
    pub fp: i64,
    pub fn_: i64,
    pub tn: Option<i64>,
}

impl ConfusionCounts {
    pub fn new(tp: i64, fp: i64, fn_: i64) -> Self {
        Self { tp, fp, fn_, tn: None }
    }
}

/// Errors mirroring the TS suite's RangeError contract.
#[derive(Debug, PartialEq)]
pub enum EvalError {
    NMustBePositive,
    COutOfRange,
    KOutOfRange,
    CountsNegative,
    LengthMismatch,
}

impl std::fmt::Display for EvalError {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        let msg = match self {
            EvalError::NMustBePositive => "n must be > 0",
            EvalError::COutOfRange => "c must be in [0, n]",
            EvalError::KOutOfRange => "k must be in [1, n]",
            EvalError::CountsNegative => "counts must be >= 0",
            EvalError::LengthMismatch => "predictions and references must have the same length",
        };
        f.write_str(msg)
    }
}

impl std::error::Error for EvalError {}

/// Unbiased pass@k: probability that at least one of k samples drawn without
/// replacement from n (of which c are correct) passes.
///
///   1                                       when n - c < k  (a wrong draw is impossible)
///   1 - Π_{i=0..k-1} (n - c - i) / (n - i)  otherwise
pub fn pass_at_k(n: i64, c: i64, k: i64) -> Result<f64, EvalError> {
    if n <= 0 {
        return Err(EvalError::NMustBePositive);
    }
    if c < 0 || c > n {
        return Err(EvalError::COutOfRange);
    }
    if k <= 0 || k > n {
        return Err(EvalError::KOutOfRange);
    }
    if n - c < k {
        return Ok(1.0);
    }
    let mut product = 1.0;
    for i in 0..k {
        product *= (n - c - i) as f64 / (n - i) as f64;
    }
    Ok(1.0 - product)
}

/// Precision/recall/F1 over confusion counts. Zero denominators score 0.
/// The zero-value `Prf1` returns alongside an error, mirroring TS.
pub fn precision_recall(counts: ConfusionCounts) -> Result<Prf1, EvalError> {
    let (tp, fp, fn_) = (counts.tp, counts.fp, counts.fn_);
    if tp < 0 || fp < 0 || fn_ < 0 {
        return Err(EvalError::CountsNegative);
    }
    let precision = if tp + fp > 0 { tp as f64 / (tp + fp) as f64 } else { 0.0 };
    let recall = if tp + fn_ > 0 { tp as f64 / (tp + fn_) as f64 } else { 0.0 };
    let f1 = if precision + recall > 0.0 {
        (2.0 * precision * recall) / (precision + recall)
    } else {
        0.0
    };
    Ok(Prf1 { precision, recall, f1 })
}

/// Micro-averaged P/R/F1 across per-class confusion counts.
pub fn micro_average(per_class: &[ConfusionCounts]) -> Result<Prf1, EvalError> {
    let mut sums = ConfusionCounts::new(0, 0, 0);
    for c in per_class {
        sums.tp += c.tp;
        sums.fp += c.fp;
        sums.fn_ += c.fn_;
    }
    precision_recall(sums)
}

/// Exact-match rate over paired predictions/references (case-sensitive).
/// Empty input scores 0; length mismatch errors.
pub fn exact_match_rate(predictions: &[str], references: &[str]) -> Result<f64, EvalError> {
    if predictions.len() != references.len() {
        return Err(EvalError::LengthMismatch);
    }
    if predictions.is_empty() {
        return Ok(0.0);
    }
    let hits = predictions
        .iter()
        .zip(references)
        .filter(|(p, r)| p == r)
        .count();
    Ok(hits as f64 / predictions.len() as f64)
}

/// P/R/F1 over label SETS — the standard multi-label / extraction metric.
pub fn set_match(prediction: &[&str], reference: &[&str]) -> Result<Prf1, EvalError> {
    use std::collections::HashSet;
    let p: HashSet<&str> = prediction.iter().copied().collect();
    let r: HashSet<&str> = reference.iter().copied().collect();
    let tp = r.iter().filter(|l| p.contains(*l)).count() as i64;
    let fp = p.difference(&r).count() as i64;
    let fn_ = r.difference(&p).count() as i64;
    precision_recall(ConfusionCounts::new(tp, fp, fn_))
}

Also available in 13 other languages

Every CosmoDev tool ships its pure logic in TypeScript (web) and Go (CLI), with authored implementations in a dozen-plus languages — the same contract, ported. Compare all languages side by side →