Skip to main content

lean_ctx/core/quality_lab/
calibration.rs

1use crate::core::tokens::{TokenizerFamily, count_tokens_for, detect_tokenizer};
2
3/// Accuracy level of a calibrated token count.
4#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, serde::Serialize, serde::Deserialize)]
5pub enum CalibrationAccuracy {
6    /// Exact BPE encoding match (e.g., o200k_base for GPT-4o)
7    ExactTextEncoding,
8    /// Proxy tokenizer with empirical correction factor applied
9    CorrectedProxyTokenizer,
10    /// Proxy tokenizer without correction (approximate)
11    ProxyTokenizer,
12    /// Character-based fallback (4 chars/token estimate)
13    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/// Token count with calibration metadata.
40#[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/// Calibration data for a tokenizer family.
52#[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
60/// Returns the built-in calibration profile for a tokenizer family.
61pub 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
90/// Counts tokens and applies the tokenizer family's calibration profile.
91pub 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
111/// Detects the tokenizer family from a model name and returns a calibrated count.
112pub fn count_tokens_calibrated_from_model(text: &str, model: &str) -> CalibratedCount {
113    count_tokens_with_calibration(text, detect_tokenizer(model))
114}
115
116/// Returns inclusive lower and upper bounds for a calibrated count.
117pub 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
125/// Returns calibrated counts for every supported tokenizer family.
126pub 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}