use std::cmp::Reverse;
use std::collections::{BinaryHeap, HashMap};
use std::sync::OnceLock;
use regex::Regex;
const VOCAB: &[u8] = include_bytes!("../../assets/o200k_base.bin");
const MAGIC: &[u8] = b"O200K\x01";
type Rank = u32;
const PIECE_PATTERN: &str = concat!(
r"[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}]*[\p{Ll}\p{Lm}\p{Lo}\p{M}]+(?i:'s|'t|'re|'ve|'m|'ll|'d)?",
"|",
r"[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}]+[\p{Ll}\p{Lm}\p{Lo}\p{M}]*(?i:'s|'t|'re|'ve|'m|'ll|'d)?",
"|",
r"\p{N}{1,3}",
"|",
r" ?[^\s\p{L}\p{N}]+[\r\n/]*",
"|",
r"\s*[\r\n]+",
"|",
r"\s+",
);
fn ranks() -> &'static HashMap<&'static [u8], Rank> {
static RANKS: OnceLock<HashMap<&'static [u8], Rank>> = OnceLock::new();
RANKS.get_or_init(|| {
let (magic, rest) = VOCAB.split_at(MAGIC.len());
assert_eq!(magic, MAGIC, "o200k vocabulary asset is corrupt");
let (count, mut rest) = rest.split_at(4);
let count = u32::from_le_bytes(count.try_into().expect("4 bytes")) as usize;
let mut map = HashMap::with_capacity(count);
for rank in 0..count {
let (len, tail) = rest.split_first().expect("vocabulary truncated");
let (token, tail) = tail.split_at(*len as usize);
map.insert(token, rank as Rank);
rest = tail;
}
assert!(rest.is_empty(), "vocabulary has trailing bytes");
map
})
}
fn pattern() -> &'static Regex {
static PATTERN: OnceLock<Regex> = OnceLock::new();
PATTERN.get_or_init(|| Regex::new(PIECE_PATTERN).expect("the piece pattern is valid"))
}
fn split_pieces(text: &str) -> impl Iterator<Item = &str> {
let mut cursor = 0usize;
std::iter::from_fn(move || {
if cursor >= text.len() {
return None;
}
let m = pattern().find_at(text, cursor)?;
let piece = m.as_str();
let is_inline_whitespace = !piece.is_empty()
&& piece.chars().all(char::is_whitespace)
&& !piece.contains(['\r', '\n']);
if is_inline_whitespace && m.end() < text.len() && piece.chars().count() > 1 {
let last = piece.char_indices().next_back().expect("non-empty").0;
cursor = m.start() + last;
return Some(&piece[..last]);
}
cursor = m.end();
Some(piece)
})
}
fn count_piece(piece: &[u8], ranks: &HashMap<&'static [u8], Rank>) -> usize {
if piece.len() <= 1 {
return piece.len();
}
const NONE: usize = usize::MAX;
let len = piece.len();
let mut next: Vec<usize> = (1..=len + 1).collect();
let mut prev: Vec<usize> = std::iter::once(NONE).chain(0..len).collect();
let rank_of = |i: usize, next: &[usize]| -> Option<Rank> {
let mid = *next.get(i)?;
let end = *next.get(mid)?;
if end > len {
return None;
}
ranks.get(&piece[i..end]).copied()
};
let mut heap: BinaryHeap<Reverse<(Rank, usize)>> = (0..len)
.filter_map(|i| rank_of(i, &next).map(|r| Reverse((r, i))))
.collect();
let mut parts = len;
while let Some(Reverse((rank, i))) = heap.pop() {
if next[i] > len || rank_of(i, &next) != Some(rank) {
continue;
}
let mid = next[i];
let end = next[mid];
next[i] = end;
if end <= len {
prev[end] = i;
}
next[mid] = NONE; parts -= 1;
if let Some(r) = rank_of(i, &next) {
heap.push(Reverse((r, i)));
}
let before = prev[i];
if before != NONE {
if let Some(r) = rank_of(before, &next) {
heap.push(Reverse((r, before)));
}
}
}
parts
}
pub fn count_tokens(text: &str) -> usize {
let ranks = ranks();
split_pieces(text)
.map(|piece| {
let bytes = piece.as_bytes();
if ranks.contains_key(bytes) {
1
} else {
count_piece(bytes, ranks)
}
})
.sum()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_vocabulary_asset_loads_and_ranks_are_dense() {
let r = ranks();
assert_eq!(r.len(), 199_998, "o200k_base has 199998 entries");
assert_eq!(r.get(b"!".as_slice()), Some(&0));
}
#[test]
fn a_single_byte_is_one_token() {
assert_eq!(count_tokens("a"), 1);
assert_eq!(count_tokens(""), 0);
}
#[test]
fn an_inline_whitespace_run_hands_its_last_space_to_the_next_word() {
let pieces: Vec<&str> = split_pieces(" leading").collect();
assert_eq!(pieces, vec![" ", " leading"]);
let pieces: Vec<&str> = split_pieces("a b c").collect();
assert_eq!(pieces, vec!["a", " b", " ", " c"]);
}
#[test]
fn a_whitespace_run_that_ends_the_text_is_taken_whole() {
let pieces: Vec<&str> = split_pieces("hi ").collect();
assert_eq!(pieces, vec!["hi", " "]);
}
#[test]
fn newlines_are_taken_whole_with_their_leading_whitespace() {
let pieces: Vec<&str> = split_pieces("a\n\nb").collect();
assert_eq!(pieces, vec!["a", "\n\n", "b"]);
}
#[test]
fn known_strings_count_to_their_published_token_counts() {
assert_eq!(count_tokens("hello world"), 2);
assert_eq!(count_tokens("The quick brown fox"), 4);
assert_eq!(count_tokens("1234567890"), 4);
}
#[test]
fn multibyte_text_counts_without_panicking_on_char_boundaries() {
assert!(count_tokens("这是一段中文文本") > 0);
assert!(count_tokens("emoji 🚀🔥") > 0);
assert!(count_tokens("café über naïve") > 0);
}
#[test]
fn a_pathological_single_piece_does_not_blow_up() {
let time = |n: usize| {
let piece = "!".repeat(n);
let start = std::time::Instant::now();
assert!(count_tokens(&piece) > 0);
start.elapsed()
};
time(1_000);
let small = time(10_000);
let large = time(40_000);
assert!(
large < small * 6,
"4x the input took {large:?} vs {small:?}: the merge went quadratic"
);
}
}
#[cfg(test)]
mod reference_equivalence {
use super::count_tokens;
#[test]
fn it_matches_the_reference_on_the_shapes_a_prompt_is_made_of() {
let vectors: &[(&str, usize)] = &[
("", 0),
("a", 1),
("hello world", 2),
("The quick brown fox jumps over the lazy dog.", 10),
(" leading spaces", 3),
("trailing spaces ", 4),
("x y", 3),
("a b c d", 6),
("\n\n\n", 1),
("\t\t", 1),
("a\n\nb", 3),
("mix \n \n end", 3),
("word \n", 2),
(" ", 1),
(" ", 1),
("a\t\tb", 3),
("end. ", 3),
("def f():\n return 1\n", 8),
(
"fn main(){let x:Vec<u32>=(0..10).filter(|n|n%2==0).collect();}",
28,
),
(
"{\"a\": 1, \"bb\": [1, 2, 3], \"ccc\": {\"d\": true, \"e\": null}}",
31,
),
("https://openrouter.ai/api/v1/models?limit=50&offset=0", 17),
("3f8a1c2e-9b4d-4f21-8e7a-1c0d5b6e2a94", 34),
("SELECT * FROM t WHERE x > 1 GROUP BY 1;", 14),
("CamelCaseIdentifier snake_case_name SCREAMING_CASE", 10),
("it's don't we're I've I'm they'll he'd", 7),
("IT'S DON'T WE'RE", 6),
("1234567890", 4),
("0.000003", 4),
("-42", 2),
("1e-9", 4),
("aa", 1),
("aaa", 1),
("aaaa", 1),
("!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!", 3),
("这是一段中文文本,用于测试分词器的行为差异。", 16),
("こんにちは、これは日本語です。", 7),
("Здравствуйте, это русский текст.", 6),
("café über naïve", 5),
("emoji 🚀🔥 and ✨", 7),
("🇫🇷🇯🇵", 8),
("\u{200b}zero width", 3),
];
for (text, expected) in vectors {
assert_eq!(
count_tokens(text),
*expected,
"token count diverged on {text:?}"
);
}
}
#[test]
fn it_matches_the_reference_on_a_real_block_of_source() {
const SOURCE: &str = r#"fn count_piece(piece: &[u8], ranks: &HashMap<&'static [u8], Rank>) -> usize {
if piece.len() <= 1 {
return piece.len();
}
let mut parts: Vec<usize> = (0..=piece.len()).collect();
let mut pair_rank: Vec<Option<Rank>> =
(0..parts.len() - 1).map(|i| rank_at(piece, &parts, ranks, i)).collect();
loop {
let Some((i, _)) = pair_rank
.iter()
.enumerate()
.filter_map(|(i, r)| r.map(|r| (i, r)))
.min_by_key(|&(i, r)| (r, i))
else {
break;
};
parts.remove(i + 1);
pair_rank.remove(i + 1);
pair_rank[i] = rank_at(piece, &parts, ranks, i);
if i > 0 {
pair_rank[i - 1] = rank_at(piece, &parts, ranks, i - 1);
}
}
parts.len() - 1
}"#;
assert_eq!(
SOURCE.len(),
800,
"the recorded count is for exactly this text"
);
assert_eq!(count_tokens(SOURCE), 246);
}
#[test]
fn it_matches_the_reference_across_a_pseudorandom_corpus() {
const REFERENCE_TOTAL: usize = 8266;
let mut state = 0x2545_F491_4F6C_DD1Du64;
let mut next = move || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
let alphabet: Vec<char> = " \t\n\rabcXYZ0129,.;:'\"/\\_-()[]{}日本語中文🚀é∀×"
.chars()
.collect();
let total: usize = (0..500)
.map(|_| {
let len = (next() % 40) as usize;
let text: String = (0..len)
.map(|_| alphabet[(next() as usize) % alphabet.len()])
.collect();
count_tokens(&text)
})
.sum();
assert_eq!(
total, REFERENCE_TOTAL,
"the encoder diverged somewhere in the corpus"
);
}
}