use regex::Regex;
pub trait PreTokenizer: Send + Sync {
fn next_match(&self, text: &str, pos: usize) -> Option<(usize, usize)>;
}
#[derive(Clone, Copy, PartialEq, Eq)]
pub(crate) enum FastPath {
None,
Cl100k,
O200k,
Qwen2,
Deepseek,
Tekken,
MiniMax,
}
#[derive(Clone, Copy, PartialEq, Eq)]
pub(crate) enum WhitespaceRules {
Generic,
NewlineFirst,
NewlineFirstSplitOnNumCjk,
}
#[inline]
fn is_deepseek_split_boundary(c: char) -> bool {
c.is_numeric() || matches!(c, '一'..='龥' | '\u{3040}'..='\u{309F}' | '\u{30A0}'..='\u{30FF}')
}
pub struct RegexPreTokenizer {
regex: Regex,
fast: FastPath,
ws: WhitespaceRules,
}
impl RegexPreTokenizer {
pub(crate) fn new(pattern: &str, fast: FastPath, ws: WhitespaceRules) -> Self {
Self {
regex: Regex::new(pattern).expect("invalid regex pattern"),
fast,
ws,
}
}
}
impl PreTokenizer for RegexPreTokenizer {
#[inline]
fn next_match(&self, text: &str, pos: usize) -> Option<(usize, usize)> {
let bytes = text.as_bytes();
let fast = match self.fast {
FastPath::Cl100k => cl100k_ascii_next::<3>(bytes, pos),
FastPath::Qwen2 => cl100k_ascii_next::<1>(bytes, pos),
FastPath::O200k => o200k_like_ascii_next::<true, 3, false>(bytes, pos),
FastPath::Tekken => o200k_like_ascii_next::<false, 1, true>(bytes, pos),
FastPath::MiniMax => o200k_like_ascii_next::<true, 3, true>(bytes, pos),
FastPath::Deepseek => deepseek_ascii_next(bytes, pos),
FastPath::None => None,
};
if let Some(r) = fast {
return Some(r);
}
let mat = self.regex.find_at(text, pos)?;
let start = mat.start();
let end = adjust_whitespace_end(bytes, start, mat.end(), self.ws);
Some((start, end))
}
}
#[inline]
fn take_line_tail<const SLASH: bool>(b: &[u8], mut k: usize) -> usize {
while k < b.len() && (b[k] == b'\r' || b[k] == b'\n' || (SLASH && b[k] == b'/')) {
k += 1;
}
k
}
#[inline(always)]
fn ascii_num_punct<const MAX_DIGITS: usize, const SLASH_TAIL: bool>(
b: &[u8],
i: usize,
) -> Option<(usize, usize)> {
let n = b.len();
let c0 = b[i];
if c0.is_ascii_digit() {
let mut j = i;
let mut k = 0;
while j < n && k < MAX_DIGITS && b[j] < 0x80 && b[j].is_ascii_digit() {
j += 1;
k += 1;
}
if k < MAX_DIGITS && j < n && b[j] >= 0x80 {
return None;
}
return Some((i, j));
}
let mut j = i;
if c0 == b' ' {
match b.get(i + 1) {
Some(&c1)
if c1 < 0x80
&& !is_ascii_ws(c1)
&& !c1.is_ascii_alphabetic()
&& !c1.is_ascii_digit() =>
{
j = i + 1;
}
_ => return None,
}
}
let cj = b[j];
if cj < 0x80 && !is_ascii_ws(cj) && !cj.is_ascii_alphabetic() && !cj.is_ascii_digit() {
let mut k = j;
while k < n
&& b[k] < 0x80
&& !is_ascii_ws(b[k])
&& !b[k].is_ascii_alphabetic()
&& !b[k].is_ascii_digit()
{
k += 1;
}
if k < n && b[k] >= 0x80 {
return None;
}
k = take_line_tail::<SLASH_TAIL>(b, k);
return Some((i, k));
}
None
}
#[inline(always)]
fn cl100k_ascii_next<const MAX_DIGITS: usize>(b: &[u8], i: usize) -> Option<(usize, usize)> {
let n = b.len();
if i >= n {
return None;
}
let c0 = b[i];
if c0 >= 0x80 {
return None;
}
if c0 == b'\''
&& let Some(len) = match_contraction(b, i)
{
return Some((i, i + len));
}
if c0 != b'\r'
&& c0 != b'\n'
&& !c0.is_ascii_alphabetic()
&& !c0.is_ascii_digit()
&& let Some(&c1) = b.get(i + 1)
&& c1 < 0x80
&& c1.is_ascii_alphabetic()
{
let mut j = i + 1;
while j < n && b[j] < 0x80 && b[j].is_ascii_alphabetic() {
j += 1;
}
if j < n && b[j] >= 0x80 {
return None;
}
return Some((i, j));
}
if c0.is_ascii_alphabetic() {
let mut j = i;
while j < n && b[j] < 0x80 && b[j].is_ascii_alphabetic() {
j += 1;
}
if j < n && b[j] >= 0x80 {
return None;
}
return Some((i, j));
}
ascii_num_punct::<MAX_DIGITS, false>(b, i)
}
#[inline(always)]
fn o200k_like_ascii_next<
const CONTRACTIONS: bool,
const MAX_DIGITS: usize,
const SLASH_TAIL: bool,
>(
b: &[u8],
i: usize,
) -> Option<(usize, usize)> {
let n = b.len();
if i >= n {
return None;
}
let c0 = b[i];
if c0 >= 0x80 {
return None;
}
let p = if c0.is_ascii_alphabetic() {
i
} else if c0 != b'\r' && c0 != b'\n' && !c0.is_ascii_digit() {
match b.get(i + 1) {
Some(&c1) if c1 < 0x80 && c1.is_ascii_alphabetic() => i + 1,
_ => return ascii_num_punct::<MAX_DIGITS, SLASH_TAIL>(b, i),
}
} else {
return ascii_num_punct::<MAX_DIGITS, SLASH_TAIL>(b, i);
};
let mut q = p;
while q < n && b[q] < 0x80 && b[q].is_ascii_uppercase() {
q += 1;
}
if q < n && b[q] >= 0x80 {
return None;
}
let letters_end = if q > p {
if q < n && b[q].is_ascii_lowercase() {
let mut r = q;
while r < n && b[r] < 0x80 && b[r].is_ascii_lowercase() {
r += 1;
}
if r < n && b[r] >= 0x80 {
return None;
}
r
} else {
q
}
} else {
let mut r = p;
while r < n && b[r] < 0x80 && b[r].is_ascii_lowercase() {
r += 1;
}
if r < n && b[r] >= 0x80 {
return None;
}
r
};
let mut end = letters_end;
if CONTRACTIONS
&& end < n
&& b[end] == b'\''
&& let Some(len) = match_contraction(b, end)
{
end += len;
}
Some((i, end))
}
#[inline]
fn deepseek_ascii_next(b: &[u8], i: usize) -> Option<(usize, usize)> {
let n = b.len();
if i >= n {
return None;
}
let c0 = b[i];
if c0 >= 0x80 {
return None; }
if c0.is_ascii_digit() {
let mut j = i;
let mut k = 0;
while j < n && k < 3 && b[j].is_ascii_digit() {
j += 1;
k += 1;
}
if k < 3 && j < n && b[j] >= 0x80 {
return None; }
return Some((i, j));
}
if c0.is_ascii_punctuation() {
if let Some(&c1) = b.get(i + 1)
&& c1 < 0x80
&& c1.is_ascii_alphabetic()
{
let mut j = i + 1;
while j < n && b[j].is_ascii_alphabetic() {
j += 1;
}
return Some((i, j));
}
let mut k = i;
while k < n && b[k] < 0x80 && b[k].is_ascii_punctuation() {
k += 1;
}
if k < n && b[k] >= 0x80 {
return None; }
k = take_line_tail::<false>(b, k);
return Some((i, k));
}
if c0.is_ascii_alphabetic() {
let mut j = i;
while j < n && b[j].is_ascii_alphabetic() {
j += 1;
}
if j < n && b[j] >= 0x80 {
return None; }
return Some((i, j));
}
if c0 == b' ' {
match b.get(i + 1) {
Some(&c1) if c1 >= 0x80 => return None, Some(&c1) if c1.is_ascii_alphabetic() => {
let mut j = i + 1;
while j < n && b[j].is_ascii_alphabetic() {
j += 1;
}
if j < n && b[j] >= 0x80 {
return None;
}
return Some((i, j));
}
Some(&c1) if c1.is_ascii_punctuation() => {
let mut k = i + 1;
while k < n && b[k] < 0x80 && b[k].is_ascii_punctuation() {
k += 1;
}
if k < n && b[k] >= 0x80 {
return None;
}
k = take_line_tail::<false>(b, k);
return Some((i, k));
}
_ => return None,
}
}
None
}
#[inline]
fn match_contraction(b: &[u8], i: usize) -> Option<usize> {
let c1 = b.get(i + 1).copied()?.to_ascii_lowercase();
match c1 {
b's' | b't' | b'm' | b'd' => Some(2),
b'r' if b.get(i + 2).map(|c| c.to_ascii_lowercase()) == Some(b'e') => Some(3),
b'v' if b.get(i + 2).map(|c| c.to_ascii_lowercase()) == Some(b'e') => Some(3),
b'l' if b.get(i + 2).map(|c| c.to_ascii_lowercase()) == Some(b'l') => Some(3),
_ => None,
}
}
#[inline]
fn adjust_whitespace_end(bytes: &[u8], start: usize, end: usize, ws: WhitespaceRules) -> usize {
if end - start <= 1 || end >= bytes.len() {
return end;
}
if ws != WhitespaceRules::Generic && matches!(bytes[end - 1], b'\r' | b'\n') {
return end;
}
let first = bytes[start];
if first > 0x20 && first < 0x7F {
return end;
}
if ws == WhitespaceRules::NewlineFirstSplitOnNumCjk
&& let Some(next) = bytes[end..].iter().next()
&& (next.is_ascii_digit() || *next >= 0x80)
&& let Some(c) = std::str::from_utf8(&bytes[end..])
.ok()
.and_then(|s| s.chars().next())
&& is_deepseek_split_boundary(c)
{
return end;
}
let piece = &bytes[start..end];
if piece.iter().all(|&b| is_ascii_ws(b)) {
let next = bytes[end];
if is_ascii_ws(next) {
return end;
}
return end - 1;
}
let matched = std::str::from_utf8(&bytes[start..end]).unwrap();
if !matched.chars().all(|c| c.is_whitespace()) {
return end;
}
let tail = std::str::from_utf8(&bytes[end..]).unwrap();
let next_char = match tail.chars().next() {
Some(c) => c,
None => return end,
};
if next_char.is_whitespace() {
return end;
}
let last_len = matched.chars().next_back().unwrap().len_utf8();
if end - last_len <= start {
return end;
}
end - last_len
}
#[inline(always)]
const fn is_ascii_ws(b: u8) -> bool {
matches!(b, b' ' | b'\t' | b'\n' | b'\r' | 0x0B | 0x0C)
}
#[cfg(test)]
mod tests {
use super::*;
fn collect_matches(pt: &dyn PreTokenizer, text: &str) -> Vec<(usize, usize)> {
let mut result = vec![];
let mut pos = 0;
while let Some((start, end)) = pt.next_match(text, pos) {
result.push((start, end));
pos = end;
}
result
}
use crate::encoding::{
CL100K_PATTERN, DEEPSEEK_V3_PATTERN, MISTRAL_V3_PATTERN, O200K_PATTERN, P50K_PATTERN,
QWEN2_PATTERN,
};
#[derive(Clone, Copy)]
struct Spec {
pattern: &'static str,
fast: FastPath,
ws: WhitespaceRules,
}
const CL100K: Spec = Spec {
pattern: CL100K_PATTERN,
fast: FastPath::Cl100k,
ws: WhitespaceRules::NewlineFirst,
};
const O200K: Spec = Spec {
pattern: O200K_PATTERN,
fast: FastPath::O200k,
ws: WhitespaceRules::NewlineFirst,
};
const QWEN2: Spec = Spec {
pattern: QWEN2_PATTERN,
fast: FastPath::Qwen2,
ws: WhitespaceRules::NewlineFirst,
};
const DEEPSEEK: Spec = Spec {
pattern: DEEPSEEK_V3_PATTERN,
fast: FastPath::Deepseek,
ws: WhitespaceRules::NewlineFirst,
};
const MISTRAL: Spec = Spec {
pattern: MISTRAL_V3_PATTERN,
fast: FastPath::Tekken,
ws: WhitespaceRules::NewlineFirst,
};
const P50K: Spec = Spec {
pattern: P50K_PATTERN,
fast: FastPath::None,
ws: WhitespaceRules::Generic,
};
impl Spec {
fn tokenizer(self) -> RegexPreTokenizer {
RegexPreTokenizer::new(self.pattern, self.fast, self.ws)
}
}
fn reference_matches(spec: Spec, text: &str) -> Vec<(usize, usize)> {
let regex = Regex::new(spec.pattern).unwrap();
let bytes = text.as_bytes();
let mut result = vec![];
let mut pos = 0;
while pos < text.len() {
let mat = match regex.find_at(text, pos) {
Some(m) => m,
None => break,
};
let start = mat.start();
let end = adjust_whitespace_end(bytes, start, mat.end(), spec.ws);
result.push((start, end));
pos = end;
}
result
}
fn assert_fast_matches_reference(spec: Spec, text: &str) {
let pt = spec.tokenizer();
assert_eq!(
reference_matches(spec, text),
collect_matches(&pt, text),
"fast/regex mismatch for {text:?}"
);
}
#[test]
fn test_cl100k_english() {
assert_fast_matches_reference(CL100K, "Hello, world!");
}
#[test]
fn test_cl100k_cjk() {
assert_fast_matches_reference(CL100K, "你好世界");
}
#[test]
fn test_cl100k_contractions() {
assert_fast_matches_reference(CL100K, "I'm don't they're we've she'll it'd");
}
#[test]
fn test_o200k_english() {
assert_fast_matches_reference(O200K, "Hello, world! CamelCase mixedScript123");
}
#[test]
fn test_p50k_english() {
assert_fast_matches_reference(P50K, "Hello world, I'm testing!");
}
#[test]
fn test_empty_input() {
let pt = CL100K.tokenizer();
assert_eq!(collect_matches(&pt, ""), vec![]);
}
#[test]
fn test_only_whitespace() {
assert_fast_matches_reference(CL100K, " \n \t ");
}
#[test]
fn test_emoji() {
assert_fast_matches_reference(CL100K, "🎉🚀💡");
}
#[test]
fn test_mixed_script() {
assert_fast_matches_reference(CL100K, "Hello 你好 World 🌍");
}
use WhitespaceRules::{Generic, NewlineFirst};
#[test]
fn test_adjust_whitespace_single_byte() {
assert_eq!(adjust_whitespace_end(b"a b", 0, 1, Generic), 1);
}
#[test]
fn test_adjust_whitespace_at_end_of_input() {
assert_eq!(adjust_whitespace_end(b" ", 0, 2, Generic), 2);
}
#[test]
fn test_adjust_whitespace_non_ws_piece() {
assert_eq!(adjust_whitespace_end(b"hello world", 0, 5, Generic), 5);
}
#[test]
fn test_adjust_whitespace_trim_before_nonws() {
let bytes = b" x";
assert_eq!(adjust_whitespace_end(bytes, 0, 2, Generic), 1);
}
#[test]
fn test_adjust_whitespace_no_trim_before_ws() {
let bytes = b" ";
assert_eq!(adjust_whitespace_end(bytes, 0, 2, Generic), 2);
}
#[test]
fn test_adjust_whitespace_unicode_slow_path() {
let input = "\u{3000}\u{3000}x";
let bytes = input.as_bytes();
assert_eq!(adjust_whitespace_end(bytes, 0, 6, Generic), 3);
}
#[test]
fn test_adjust_whitespace_unicode_followed_by_unicode_ws() {
let input = "\u{3000}\u{3000}\u{3000}";
let bytes = input.as_bytes();
assert_eq!(adjust_whitespace_end(bytes, 0, 6, Generic), 6);
}
#[test]
fn test_adjust_whitespace_single_multibyte_ws_before_nonws() {
let input = "\u{3000}x";
let bytes = input.as_bytes();
assert_eq!(adjust_whitespace_end(bytes, 0, 3, Generic), 3);
}
#[test]
fn test_adjust_whitespace_newline_branch_keeps_double_newline() {
let bytes = b"\n\nx";
assert_eq!(adjust_whitespace_end(bytes, 0, 2, NewlineFirst), 2);
assert_eq!(adjust_whitespace_end(bytes, 0, 2, Generic), 1);
}
#[test]
fn test_adjust_whitespace_newline_branch_keeps_crlf() {
let bytes = b"\r\n@";
assert_eq!(adjust_whitespace_end(bytes, 0, 2, NewlineFirst), 2);
assert_eq!(adjust_whitespace_end(bytes, 0, 2, Generic), 1);
}
#[test]
fn test_adjust_whitespace_newline_branch_still_trims_spaces() {
let bytes = b" x";
assert_eq!(adjust_whitespace_end(bytes, 0, 2, NewlineFirst), 1);
}
#[test]
fn test_adjust_whitespace_newline_branch_trims_trailing_spaces_after_newline() {
let bytes = b"\n x";
assert_eq!(adjust_whitespace_end(bytes, 1, 3, NewlineFirst), 2);
}
#[test]
fn test_all_patterns_match_reference() {
let texts = vec![
"Hello, world!",
"你好世界",
"fn main() { }",
" hello ",
"line1\nline2\n",
"café résumé",
"100% of $1,000",
"a@b.com",
" \t\n ",
"",
"a",
"hello world! 你好 🚀 test 123",
"word\n\nnext",
"\r\n@rem",
"a\n\n\nb",
"a \n\n b",
];
for spec in [CL100K, O200K, QWEN2, DEEPSEEK, MISTRAL, P50K] {
for text in &texts {
assert_fast_matches_reference(spec, text);
}
}
}
proptest::proptest! {
#![proptest_config(proptest::prelude::ProptestConfig::with_cases(20000))]
#[test]
fn prop_cl100k_fast_matches_regex(text in ".*") {
let pt = CL100K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(CL100K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_cl100k_fast_matches_regex_ascii(text in "[ -~ \t\r\n]*") {
let pt = CL100K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(CL100K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_cl100k_fast_matches_regex_newlines(text in "[\r\n \tabc.!]*") {
let pt = CL100K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(CL100K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_o200k_fast_matches_regex(text in ".*") {
let pt = O200K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(O200K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_o200k_fast_matches_regex_ascii(text in "[ -~ \t\r\n]*") {
let pt = O200K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(O200K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_o200k_fast_matches_regex_newlines(text in "[\r\n \tabc.!]*") {
let pt = O200K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(O200K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_qwen2_fast_matches_regex(text in ".*") {
let pt = QWEN2.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(QWEN2, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_qwen2_fast_matches_regex_ascii(text in "[ -~ \t\r\n]*") {
let pt = QWEN2.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(QWEN2, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_deepseek_fast_matches_regex(text in ".*") {
let pt = DEEPSEEK.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(DEEPSEEK, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_deepseek_fast_matches_regex_ascii(text in "[ -~ \t\r\n]*") {
let pt = DEEPSEEK.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(DEEPSEEK, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_mistral_fast_matches_regex(text in ".*") {
let pt = MISTRAL.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(MISTRAL, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_mistral_fast_matches_regex_ascii(text in "[ -~ \t\r\n]*") {
let pt = MISTRAL.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(MISTRAL, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_mistral_fast_matches_regex_slashes(text in "[/\r\n .!abcAB0]*") {
let pt = MISTRAL.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(MISTRAL, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_p50k_fast_matches_regex(text in "[ -~ \t\r\n]*") {
let pt = P50K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(P50K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
}
}