gemini_memory_rs/evals/
metrics.rs

1//! Retrieval quality metrics.
2//!
3//! Ordinary information-retrieval measures, defined here rather than pulled in
4//! so the eval harness has no dependency of its own and the definitions are
5//! visible next to the thresholds they are judged against.
6
7/// Fraction of returned results that were relevant.
8///
9/// Returning nothing is treated as perfect precision — it is a recall failure,
10/// and counting it twice would hide which of the two actually went wrong.
11pub fn precision(returned: &[String], relevant: &[String]) -> f32 {
12    if returned.is_empty() {
13        return 1.0;
14    }
15    let hits = returned.iter().filter(|r| relevant.contains(r)).count();
16    hits as f32 / returned.len() as f32
17}
18
19/// Fraction of relevant results that were returned.
20pub fn recall(returned: &[String], relevant: &[String]) -> f32 {
21    if relevant.is_empty() {
22        return 1.0;
23    }
24    let hits = relevant.iter().filter(|r| returned.contains(r)).count();
25    hits as f32 / relevant.len() as f32
26}
27
28/// Reciprocal of the rank of the first relevant result.
29pub fn reciprocal_rank(returned: &[String], relevant: &[String]) -> f32 {
30    returned
31        .iter()
32        .position(|r| relevant.contains(r))
33        .map(|idx| 1.0 / (idx as f32 + 1.0))
34        .unwrap_or(0.0)
35}
36
37/// Normalized discounted cumulative gain over binary relevance.
38pub fn ndcg(returned: &[String], relevant: &[String]) -> f32 {
39    if relevant.is_empty() {
40        return 1.0;
41    }
42    let dcg: f32 = returned
43        .iter()
44        .enumerate()
45        .map(|(idx, id)| {
46            if relevant.contains(id) {
47                1.0 / ((idx as f32 + 2.0).log2())
48            } else {
49                0.0
50            }
51        })
52        .sum();
53    let ideal: f32 = (0..relevant.len().min(returned.len().max(1)))
54        .map(|idx| 1.0 / ((idx as f32 + 2.0).log2()))
55        .sum();
56    if ideal == 0.0 {
57        0.0
58    } else {
59        (dcg / ideal).min(1.0)
60    }
61}
62
63/// Mean of a set of per-case scores.
64pub fn mean(scores: &[f32]) -> f32 {
65    if scores.is_empty() {
66        return 0.0;
67    }
68    scores.iter().sum::<f32>() / scores.len() as f32
69}
70
71/// The `p`th percentile of a sample, by nearest rank.
72pub fn percentile(samples: &mut [f32], p: f32) -> f32 {
73    if samples.is_empty() {
74        return 0.0;
75    }
76    samples.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
77    let rank = ((p / 100.0) * samples.len() as f32).ceil().max(1.0) as usize;
78    samples[rank.min(samples.len()) - 1]
79}
80
81#[cfg(test)]
82mod tests {
83    use super::*;
84
85    fn ids(values: &[&str]) -> Vec<String> {
86        values.iter().map(|v| (*v).to_string()).collect()
87    }
88
89    #[test]
90    fn precision_and_recall_pull_apart_the_two_failure_modes() {
91        let returned = ids(&["a", "b", "c", "d"]);
92        let relevant = ids(&["a", "b"]);
93        assert_eq!(precision(&returned, &relevant), 0.5);
94        assert_eq!(recall(&returned, &relevant), 1.0);
95
96        let stingy = ids(&["a"]);
97        assert_eq!(precision(&stingy, &relevant), 1.0);
98        assert_eq!(recall(&stingy, &relevant), 0.5);
99    }
100
101    #[test]
102    fn returning_nothing_is_a_recall_failure_not_a_precision_one() {
103        let relevant = ids(&["a"]);
104        assert_eq!(precision(&[], &relevant), 1.0);
105        assert_eq!(recall(&[], &relevant), 0.0);
106    }
107
108    #[test]
109    fn reciprocal_rank_rewards_ranking_the_answer_first() {
110        let relevant = ids(&["c"]);
111        assert_eq!(reciprocal_rank(&ids(&["c", "a", "b"]), &relevant), 1.0);
112        assert!((reciprocal_rank(&ids(&["a", "c", "b"]), &relevant) - 0.5).abs() < 1e-6);
113        assert_eq!(reciprocal_rank(&ids(&["a", "b"]), &relevant), 0.0);
114    }
115
116    #[test]
117    fn ndcg_prefers_relevant_results_earlier() {
118        let relevant = ids(&["a", "b"]);
119        let good = ndcg(&ids(&["a", "b", "x"]), &relevant);
120        let worse = ndcg(&ids(&["x", "a", "b"]), &relevant);
121        assert!(good > worse);
122        assert!(good <= 1.0);
123    }
124
125    #[test]
126    fn percentile_uses_nearest_rank() {
127        let mut samples = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
128        assert_eq!(percentile(&mut samples, 50.0), 5.0);
129        assert_eq!(percentile(&mut samples, 95.0), 10.0);
130        assert_eq!(percentile(&mut [], 95.0), 0.0);
131    }
132
133    #[test]
134    fn mean_of_nothing_is_zero() {
135        assert_eq!(mean(&[]), 0.0);
136        assert_eq!(mean(&[1.0, 3.0]), 2.0);
137    }
138}