#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum CharClass {
Whitespace,
AsciiLetter,
AsciiDigit,
AsciiPunct,
Cjk,
OtherLetter,
OtherSymbol,
}
impl CharClass {
fn weight(self) -> f64 {
match self {
CharClass::Whitespace => 0.25,
CharClass::AsciiLetter => 0.25,
CharClass::AsciiDigit => 0.40,
CharClass::AsciiPunct => 0.50,
CharClass::Cjk => 0.75,
CharClass::OtherLetter => 0.50,
CharClass::OtherSymbol => 1.00,
}
}
}
fn is_cjk(c: char) -> bool {
matches!(c as u32,
0x3000..=0x303F | 0x3040..=0x309F | 0x30A0..=0x30FF | 0x3400..=0x4DBF | 0x4E00..=0x9FFF | 0xAC00..=0xD7AF | 0xF900..=0xFAFF | 0xFF00..=0xFFEF | 0x20000..=0x2FA1F )
}
fn classify(c: char) -> CharClass {
if c.is_whitespace() {
CharClass::Whitespace
} else if c.is_ascii() {
if c.is_ascii_alphabetic() {
CharClass::AsciiLetter
} else if c.is_ascii_digit() {
CharClass::AsciiDigit
} else {
CharClass::AsciiPunct
}
} else if is_cjk(c) {
CharClass::Cjk
} else if c.is_alphabetic() {
CharClass::OtherLetter
} else {
CharClass::OtherSymbol
}
}
fn raw_cost(text: &str) -> f64 {
text.chars().map(|c| classify(c).weight()).sum()
}
#[must_use]
pub fn estimate_tokens(text: &str) -> usize {
raw_cost(text).round() as usize
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct LineEstimate {
pub line: usize,
pub tokens: usize,
}
#[must_use]
pub fn estimate_tokens_by_line(text: &str) -> Vec<LineEstimate> {
text.split('\n')
.enumerate()
.map(|(i, line)| LineEstimate {
line: i + 1,
tokens: estimate_tokens(line),
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_string_is_zero() {
assert_eq!(estimate_tokens(""), 0);
}
#[test]
fn deterministic_same_input_same_count() {
let s = "The quick brown fox jumps over the lazy dog. 你好世界 123!";
let first = estimate_tokens(s);
for _ in 0..5 {
assert_eq!(estimate_tokens(s), first);
}
}
#[test]
fn monotonic_longer_text_never_fewer_tokens() {
let base = "Lorem ipsum dolor sit amet";
let mut acc = String::new();
let mut prev = 0usize;
for word in base.split(' ') {
acc.push_str(word);
acc.push(' ');
let now = estimate_tokens(&acc);
assert!(now >= prev, "estimate dropped from {prev} to {now} after adding {word:?}");
prev = now;
}
}
#[test]
fn english_paragraph_within_documented_band() {
let text = "The quick brown fox jumps over the lazy dog.";
let chars = text.chars().count() as f64; let est = estimate_tokens(text) as f64;
let chars_per_token = chars / est;
assert!(
(3.3..=5.0).contains(&chars_per_token),
"English estimate {est} => {chars_per_token:.2} chars/token, want 3.3-5.0 \
(real BPE ~10 tokens for {chars} chars)"
);
}
#[test]
fn longer_english_prose_within_band() {
let text = "Tokenization turns text into the discrete units a language model \
consumes. Estimating that count cheaply lets an agent decide whether \
to read a file whole or ask only for its outline, spending its budget \
where it matters most.";
let chars = text.chars().count() as f64;
let est = estimate_tokens(text) as f64;
let chars_per_token = chars / est;
assert!(
(3.3..=5.0).contains(&chars_per_token),
"prose estimate {est} => {chars_per_token:.2} chars/token, want 3.3-5.0"
);
}
#[test]
fn cjk_string_estimated_sensibly_not_4x_off() {
let text = "机器学习模型需要把文本转换成离散的标记单元";
let chars = text.chars().count();
let est = estimate_tokens(text);
let naive_latin = chars / 4; assert!(
est >= 2 * naive_latin,
"CJK estimate {est} is too close to the Latin heuristic {naive_latin} (4x under)"
);
let tokens_per_char = est as f64 / chars as f64;
assert!(
(0.5..=1.0).contains(&tokens_per_char),
"CJK estimate {est} => {tokens_per_char:.2} tokens/char, want 0.5-1.0"
);
}
#[test]
fn mixed_script_between_regimes() {
let text = "The model 模型 reads text 文本 as tokens 标记.";
let est = estimate_tokens(text);
let chars = text.chars().count();
assert!(est > chars / 4, "mixed estimate {est} should exceed the pure-Latin chars/4");
assert!(est < chars, "mixed estimate {est} should stay below 1 token/char");
}
#[test]
fn whitespace_only_is_cheap_but_nonzero_for_runs() {
assert_eq!(estimate_tokens(""), 0);
assert_eq!(estimate_tokens(" "), 0);
assert!(estimate_tokens(&" ".repeat(40)) > 0);
}
#[test]
fn digits_denser_than_letters() {
let letters = estimate_tokens(&"a".repeat(20));
let digits = estimate_tokens(&"1".repeat(20));
assert!(digits > letters, "digits {digits} should exceed letters {letters}");
}
#[test]
fn by_line_numbers_lines_from_one() {
let text = "alpha beta\ngamma delta epsilon\n";
let lines = estimate_tokens_by_line(text);
assert_eq!(lines.len(), 3);
assert_eq!(lines[0].line, 1);
assert_eq!(lines[1].line, 2);
assert_eq!(lines[2].line, 3);
assert_eq!(lines[2].tokens, 0, "trailing empty line costs nothing");
assert!(lines[1].tokens >= lines[0].tokens);
}
#[test]
fn by_line_empty_input_is_one_zero_line() {
let lines = estimate_tokens_by_line("");
assert_eq!(lines.len(), 1);
assert_eq!(lines[0], LineEstimate { line: 1, tokens: 0 });
}
#[test]
fn by_line_breakdown_roughly_sums_to_total() {
let text = "one two three\nfour five six seven\neight nine";
let total = estimate_tokens(text);
let sum: usize = estimate_tokens_by_line(text).iter().map(|l| l.tokens).sum();
let diff = total.abs_diff(sum);
assert!(diff <= 2, "per-line sum {sum} vs total {total} drifted by {diff}");
}
}