1use crate::extract::admit_seat_write;
8use crate::search::atom_tokens;
9use serde_json::Value;
10use std::collections::BTreeMap;
11
12pub const BENCHES: &[&str] = &["LongMemEval_S", "MemoryAgentBench", "MemConflict"];
14
15#[derive(Debug, Clone, Default, PartialEq)]
16pub struct BenchRow {
17 pub admitted: usize,
18 pub refused: usize,
19 pub asked: usize,
20 pub hit: usize,
21}
22
23impl BenchRow {
24 #[must_use]
25 pub fn hit_rate(&self) -> Option<f64> {
26 if self.asked == 0 {
27 None
28 } else {
29 Some(self.hit as f64 / self.asked as f64)
30 }
31 }
32}
33
34#[derive(Debug, Clone, Default)]
35pub struct PaperATable {
36 pub rows: BTreeMap<String, BenchRow>,
37}
38
39impl PaperATable {
40 #[must_use]
41 pub fn new() -> Self {
42 let mut rows = BTreeMap::new();
43 for name in BENCHES {
44 rows.insert((*name).to_string(), BenchRow::default());
45 }
46 Self { rows }
47 }
48
49 pub fn add(&mut self, rec: &Value) {
51 let bench = rec["bench"].as_str().unwrap_or("").to_string();
52 let row = self.rows.entry(bench).or_default();
53 let write = rec["write"].as_str().unwrap_or("");
54 let admitted = admit_seat_write(write);
55 if admitted.is_none() {
56 row.refused += 1;
57 return;
58 }
59 row.admitted += 1;
60 let question = rec["question"].as_str().unwrap_or("");
61 if question.is_empty() {
62 return;
63 }
64 row.asked += 1;
65 let answers: Vec<String> = rec["answers"]
66 .as_array()
67 .map(|a| {
68 a.iter()
69 .filter_map(|v| v.as_str().map(str::to_string))
70 .collect()
71 })
72 .unwrap_or_default();
73 let claim = match admitted {
74 Some(
75 crate::extract::SeatWrite::Lesson(c) | crate::extract::SeatWrite::Preference(c),
76 ) => c,
77 Some(crate::extract::SeatWrite::Accept(_)) => return,
78 None => return,
79 };
80 let claim_rec = serde_json::json!({"text": claim});
81 let q_rec = serde_json::json!({"text": question});
82 let hay = atom_tokens(claim_rec.as_object().unwrap());
83 let qtoks = atom_tokens(q_rec.as_object().unwrap());
84 let overlap = !qtoks.is_empty() && qtoks.iter().any(|t| hay.contains(t));
85 let span = answers.iter().any(|a| {
86 let a = a.to_ascii_lowercase();
87 !a.is_empty() && claim.to_ascii_lowercase().contains(&a)
88 });
89 if overlap && span {
90 row.hit += 1;
91 }
92 }
93
94 #[must_use]
95 pub fn table(&self) -> String {
96 let mut s = String::from(
97 "| bench | admitted | refused | asked | hit@1 |\n|---|---:|---:|---:|---:|\n",
98 );
99 for name in BENCHES {
100 let row = self.rows.get(*name).cloned().unwrap_or_default();
101 match row.hit_rate() {
102 Some(r) => s.push_str(&format!(
103 "| {name} | {} | {} | {} | {:.3} |\n",
104 row.admitted, row.refused, row.asked, r
105 )),
106 None => s.push_str(&format!(
107 "| {name} | {} | {} | 0 | — |\n",
108 row.admitted, row.refused
109 )),
110 }
111 }
112 s
113 }
114}
115
116#[cfg(test)]
117mod tests {
118 use super::*;
119 use serde_json::json;
120
121 #[test]
122 fn three_named_benches_produce_a_table() {
123 let mut t = PaperATable::new();
124 t.add(&json!({
125 "bench": "LongMemEval_S",
126 "write": "Remember: the user's dog is named Rex",
127 "question": "What is the dog named?",
128 "answers": ["Rex"]
129 }));
130 t.add(&json!({
131 "bench": "MemoryAgentBench",
132 "write": "The user lives in Berlin.",
133 "question": "Where does the user live?",
134 "answers": ["Berlin"]
135 }));
136 t.add(&json!({
137 "bench": "MemConflict",
138 "write": "Remember: the capital of Germany is Berlin now",
139 "question": "What is the capital of Germany?",
140 "answers": ["Berlin"]
141 }));
142 let table = t.table();
143 assert!(table.contains("LongMemEval_S"), "{table}");
144 assert!(table.contains("MemoryAgentBench"), "{table}");
145 assert!(table.contains("MemConflict"), "{table}");
146 assert_eq!(t.rows["LongMemEval_S"].hit, 1);
147 assert_eq!(t.rows["MemoryAgentBench"].refused, 1);
148 assert_eq!(t.rows["MemoryAgentBench"].asked, 0);
149 assert_eq!(t.rows["MemConflict"].hit, 1);
150 }
151}