lean_ctx/core/quality_lab/
calibration.rs1use crate::core::tokens::{TokenizerFamily, count_tokens_for, detect_tokenizer};
2
3#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, serde::Serialize, serde::Deserialize)]
5pub enum CalibrationAccuracy {
6 ExactTextEncoding,
8 CorrectedProxyTokenizer,
10 ProxyTokenizer,
12 CharFallback,
14}
15
16impl CalibrationAccuracy {
17 const fn rank(self) -> u8 {
18 match self {
19 Self::ExactTextEncoding => 3,
20 Self::CorrectedProxyTokenizer => 2,
21 Self::ProxyTokenizer => 1,
22 Self::CharFallback => 0,
23 }
24 }
25}
26
27impl Ord for CalibrationAccuracy {
28 fn cmp(&self, other: &Self) -> std::cmp::Ordering {
29 self.rank().cmp(&other.rank())
30 }
31}
32
33impl PartialOrd for CalibrationAccuracy {
34 fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
35 Some(self.cmp(other))
36 }
37}
38
39#[derive(Debug, Clone, Copy, serde::Serialize, serde::Deserialize)]
41pub struct CalibratedCount {
42 pub tokens: u64,
43 #[serde(
44 serialize_with = "serialize_tokenizer_family",
45 deserialize_with = "deserialize_tokenizer_family"
46 )]
47 pub family: TokenizerFamily,
48 pub accuracy: CalibrationAccuracy,
49}
50
51#[derive(Debug, Clone)]
53pub struct CalibrationProfile {
54 pub family: TokenizerFamily,
55 pub correction_factor: f64,
56 pub accuracy: CalibrationAccuracy,
57 pub variance_pct: f64,
58}
59
60pub fn calibration_profile(family: TokenizerFamily) -> CalibrationProfile {
62 match family {
63 TokenizerFamily::O200kBase => CalibrationProfile {
64 family,
65 correction_factor: 1.0,
66 accuracy: CalibrationAccuracy::ExactTextEncoding,
67 variance_pct: 0.0,
68 },
69 TokenizerFamily::Cl100k => CalibrationProfile {
70 family,
71 correction_factor: 1.0,
72 accuracy: CalibrationAccuracy::CorrectedProxyTokenizer,
73 variance_pct: 3.0,
74 },
75 TokenizerFamily::Gemini => CalibrationProfile {
76 family,
77 correction_factor: 1.08,
78 accuracy: CalibrationAccuracy::CorrectedProxyTokenizer,
79 variance_pct: 5.0,
80 },
81 TokenizerFamily::Llama => CalibrationProfile {
82 family,
83 correction_factor: 1.0,
84 accuracy: CalibrationAccuracy::ProxyTokenizer,
85 variance_pct: 8.0,
86 },
87 }
88}
89
90pub fn count_tokens_with_calibration(text: &str, family: TokenizerFamily) -> CalibratedCount {
92 if text.is_empty() {
93 return CalibratedCount {
94 tokens: 0,
95 family,
96 accuracy: CalibrationAccuracy::ExactTextEncoding,
97 };
98 }
99
100 let profile = calibration_profile(family);
101 let raw = count_tokens_for(text, family);
102 let corrected = (raw as f64 * profile.correction_factor).ceil() as u64;
103
104 CalibratedCount {
105 tokens: corrected,
106 family,
107 accuracy: profile.accuracy,
108 }
109}
110
111pub fn count_tokens_calibrated_from_model(text: &str, model: &str) -> CalibratedCount {
113 count_tokens_with_calibration(text, detect_tokenizer(model))
114}
115
116pub fn calibration_variance(count: &CalibratedCount) -> (u64, u64) {
118 let variance_pct = calibration_profile(count.family).variance_pct;
119 let variance = count.tokens as f64 * variance_pct / 100.0;
120 let lower = (count.tokens as f64 - variance).floor().max(0.0) as u64;
121 let upper = (count.tokens as f64 + variance).ceil() as u64;
122 (lower, upper)
123}
124
125pub fn compare_calibration(text: &str) -> Vec<CalibratedCount> {
127 [
128 TokenizerFamily::O200kBase,
129 TokenizerFamily::Cl100k,
130 TokenizerFamily::Gemini,
131 TokenizerFamily::Llama,
132 ]
133 .into_iter()
134 .map(|family| count_tokens_with_calibration(text, family))
135 .collect()
136}
137
138fn serialize_tokenizer_family<S>(family: &TokenizerFamily, serializer: S) -> Result<S::Ok, S::Error>
139where
140 S: serde::Serializer,
141{
142 serializer.serialize_str(match family {
143 TokenizerFamily::O200kBase => "O200kBase",
144 TokenizerFamily::Cl100k => "Cl100k",
145 TokenizerFamily::Gemini => "Gemini",
146 TokenizerFamily::Llama => "Llama",
147 })
148}
149
150fn deserialize_tokenizer_family<'de, D>(deserializer: D) -> Result<TokenizerFamily, D::Error>
151where
152 D: serde::Deserializer<'de>,
153{
154 let value = <String as serde::Deserialize>::deserialize(deserializer)?;
155 match value.as_str() {
156 "O200kBase" => Ok(TokenizerFamily::O200kBase),
157 "Cl100k" => Ok(TokenizerFamily::Cl100k),
158 "Gemini" => Ok(TokenizerFamily::Gemini),
159 "Llama" => Ok(TokenizerFamily::Llama),
160 _ => Err(serde::de::Error::custom("unknown tokenizer family")),
161 }
162}
163
164#[cfg(test)]
165mod tests {
166 use super::{
167 CalibratedCount, CalibrationAccuracy, calibration_profile, calibration_variance,
168 compare_calibration, count_tokens_calibrated_from_model, count_tokens_with_calibration,
169 };
170 use crate::core::tokens::{TokenizerFamily, count_tokens_for};
171
172 const SAMPLE: &str = "fn summarize(records: &[Record]) -> Summary { records.iter().collect() }";
173
174 #[test]
175 fn test_o200k_exact_calibration() {
176 let profile = calibration_profile(TokenizerFamily::O200kBase);
177 let count = count_tokens_with_calibration(SAMPLE, TokenizerFamily::O200kBase);
178
179 assert_eq!(profile.correction_factor, 1.0);
180 assert_eq!(count.accuracy, CalibrationAccuracy::ExactTextEncoding);
181 assert_eq!(
182 count.tokens,
183 count_tokens_for(SAMPLE, TokenizerFamily::O200kBase) as u64
184 );
185 }
186
187 #[test]
188 fn test_gemini_correction_factor() {
189 let raw = count_tokens_for(SAMPLE, TokenizerFamily::Gemini);
190 let count = count_tokens_with_calibration(SAMPLE, TokenizerFamily::Gemini);
191
192 assert_eq!(count.tokens, (raw as f64 * 1.08).ceil() as u64);
193 }
194
195 #[test]
196 fn test_cl100k_corrected_proxy() {
197 let count = count_tokens_with_calibration(SAMPLE, TokenizerFamily::Cl100k);
198
199 assert_eq!(count.accuracy, CalibrationAccuracy::CorrectedProxyTokenizer);
200 }
201
202 #[test]
203 fn test_llama_proxy_accuracy() {
204 let count = count_tokens_with_calibration(SAMPLE, TokenizerFamily::Llama);
205
206 assert_eq!(count.accuracy, CalibrationAccuracy::ProxyTokenizer);
207 }
208
209 #[test]
210 fn test_empty_text_zero_tokens() {
211 let count = count_tokens_with_calibration("", TokenizerFamily::Gemini);
212
213 assert_eq!(count.tokens, 0);
214 assert_eq!(count.accuracy, CalibrationAccuracy::ExactTextEncoding);
215 }
216
217 #[test]
218 fn test_model_detection_claude() {
219 let count = count_tokens_calibrated_from_model(SAMPLE, "claude-sonnet-4-20250514");
220
221 assert_eq!(count.family, TokenizerFamily::Cl100k);
222 }
223
224 #[test]
225 fn test_model_detection_gpt4o() {
226 let count = count_tokens_calibrated_from_model(SAMPLE, "gpt-4o");
227
228 assert_eq!(count.family, TokenizerFamily::O200kBase);
229 }
230
231 #[test]
232 fn test_variance_bounds() {
233 let count = CalibratedCount {
234 tokens: 1_000,
235 family: TokenizerFamily::Gemini,
236 accuracy: CalibrationAccuracy::CorrectedProxyTokenizer,
237 };
238
239 assert_eq!(calibration_variance(&count), (950, 1_050));
240 }
241
242 #[test]
243 fn test_compare_all_families() {
244 let counts = compare_calibration(SAMPLE);
245 let families: Vec<TokenizerFamily> = counts.into_iter().map(|count| count.family).collect();
246
247 assert_eq!(
248 families,
249 vec![
250 TokenizerFamily::O200kBase,
251 TokenizerFamily::Cl100k,
252 TokenizerFamily::Gemini,
253 TokenizerFamily::Llama,
254 ]
255 );
256 }
257
258 #[test]
259 fn test_calibration_ordering() {
260 assert!(
261 CalibrationAccuracy::ExactTextEncoding > CalibrationAccuracy::CorrectedProxyTokenizer
262 );
263 assert!(CalibrationAccuracy::CorrectedProxyTokenizer > CalibrationAccuracy::ProxyTokenizer);
264 assert!(CalibrationAccuracy::ProxyTokenizer > CalibrationAccuracy::CharFallback);
265 }
266}