use super::unicode_tables as ucd;
pub fn in_ranges(cp: u32, ranges: &[(u32, u32)]) -> bool {
match ranges.binary_search_by(|&(lo, _)| lo.cmp(&cp)) {
Ok(_) => true, Err(0) => false,
Err(idx) => {
let (_, hi) = ranges[idx - 1];
cp <= hi
}
}
}
#[inline]
fn is_l(c: char) -> bool {
in_ranges(c as u32, ucd::LETTER)
}
#[inline]
fn is_m(c: char) -> bool {
in_ranges(c as u32, ucd::MARK)
}
#[inline]
fn is_n(c: char) -> bool {
in_ranges(c as u32, ucd::NUMBER)
}
#[inline]
fn is_p(c: char) -> bool {
in_ranges(c as u32, ucd::PUNCTUATION)
}
#[inline]
fn is_s(c: char) -> bool {
in_ranges(c as u32, ucd::SYMBOL)
}
#[inline]
fn is_ws(c: char) -> bool {
c.is_whitespace()
}
#[inline]
fn is_ascii_punct_lead(c: char) -> bool {
matches!(
c,
'!' | '"'
| '#'
| '$'
| '%'
| '&'
| '\''
| '('
| ')'
| '*'
| '+'
| ','
| '-'
| '.'
| '/'
| ':'
| ';'
| '<'
| '='
| '>'
| '?'
| '@'
| '['
| '\\'
| ']'
| '^'
| '_'
| '`'
| '{'
| '|'
| '}'
| '~'
)
}
#[inline]
fn is_cjk_kana(c: char) -> bool {
let cp = c as u32;
(0x4E00..=0x9FA5).contains(&cp)
|| (0x3040..=0x309F).contains(&cp)
|| (0x30A0..=0x30FF).contains(&cp)
}
pub fn pretokenize(text: &str) -> Vec<String> {
let mut pieces = vec![text.to_string()];
pieces = split_stage(&pieces, split_digit_groups);
pieces = split_stage(&pieces, split_cjk_kana);
pieces = split_stage(&pieces, split_gpt_word);
pieces.iter().map(|p| byte_level_map(p)).collect()
}
pub fn pretokenize_smollm2(text: &str) -> Vec<String> {
let mut pieces = vec![text.to_string()];
pieces = split_stage(&pieces, split_digits_individual);
pieces = split_stage(&pieces, split_gpt2_word);
pieces.iter().map(|p| byte_level_map(p)).collect()
}
pub fn pretokenize_gpt2(text: &str) -> Vec<String> {
split_gpt2_word(text)
.iter()
.map(|p| byte_level_map(p))
.collect()
}
fn split_stage(pieces: &[String], f: fn(&str) -> Vec<String>) -> Vec<String> {
let mut out = Vec::with_capacity(pieces.len());
for p in pieces {
out.extend(f(p));
}
out
}
fn split_digit_groups(s: &str) -> Vec<String> {
let chars: Vec<char> = s.chars().collect();
let mut out = Vec::new();
let mut buf = String::new();
let mut i = 0;
while i < chars.len() {
if is_n(chars[i]) {
if !buf.is_empty() {
out.push(std::mem::take(&mut buf));
}
let mut grp = String::new();
let mut taken = 0;
while i < chars.len() && taken < 3 && is_n(chars[i]) {
grp.push(chars[i]);
i += 1;
taken += 1;
}
out.push(grp);
} else {
buf.push(chars[i]);
i += 1;
}
}
if !buf.is_empty() {
out.push(buf);
}
out
}
fn split_cjk_kana(s: &str) -> Vec<String> {
isolate_runs(s, is_cjk_kana)
}
fn split_digits_individual(s: &str) -> Vec<String> {
let mut out = Vec::new();
let mut buf = String::new();
for c in s.chars() {
if is_n(c) {
if !buf.is_empty() {
out.push(std::mem::take(&mut buf));
}
out.push(c.to_string());
} else {
buf.push(c);
}
}
if !buf.is_empty() {
out.push(buf);
}
out
}
fn isolate_runs(s: &str, pred: fn(char) -> bool) -> Vec<String> {
let mut out = Vec::new();
let mut buf = String::new();
let mut in_run = false;
for c in s.chars() {
let hit = pred(c);
if hit != in_run {
if !buf.is_empty() {
out.push(std::mem::take(&mut buf));
}
in_run = hit;
}
buf.push(c);
}
if !buf.is_empty() {
out.push(buf);
}
out
}
fn split_gpt_word(s: &str) -> Vec<String> {
let chars: Vec<char> = s.chars().collect();
let mut out = Vec::new();
let mut gap = String::new();
let mut i = 0;
while i < chars.len() {
if let Some(len) = match_word(&chars, i) {
if !gap.is_empty() {
out.push(std::mem::take(&mut gap));
}
let tok: String = chars[i..i + len].iter().collect();
out.push(tok);
i += len;
} else {
gap.push(chars[i]);
i += 1;
}
}
if !gap.is_empty() {
out.push(gap);
}
out
}
fn match_word(chars: &[char], i: usize) -> Option<usize> {
let n = chars.len();
let at = |k: usize| -> Option<char> { chars.get(k).copied() };
if let Some(c0) = at(i)
&& is_ascii_punct_lead(c0)
{
let mut j = i + 1;
while j < n && chars[j].is_ascii_alphabetic() {
j += 1;
}
if j > i + 1 {
return Some(j - i);
}
}
{
let mut j = i;
if let Some(c) = at(j)
&& c != '\r'
&& c != '\n'
&& !is_l(c)
&& !is_p(c)
&& !is_s(c)
{
if let Some(c1) = at(j + 1)
&& (is_l(c1) || is_m(c1))
{
j += 1;
}
}
let start_lm = j;
while j < n && (is_l(chars[j]) || is_m(chars[j])) {
j += 1;
}
if j > start_lm {
return Some(j - i);
}
}
{
let mut j = i;
let lead_space = at(j) == Some(' ');
if lead_space {
j += 1;
}
let start_ps = j;
while j < n && (is_p(chars[j]) || is_s(chars[j])) {
j += 1;
}
if j > start_ps {
while j < n && (chars[j] == '\r' || chars[j] == '\n') {
j += 1;
}
return Some(j - i);
}
}
{
let mut last_crlf_end = None;
let mut k = i;
while k < n && is_ws(chars[k]) {
if chars[k] == '\r' || chars[k] == '\n' {
last_crlf_end = Some(k + 1);
}
k += 1;
}
if let Some(end) = last_crlf_end {
return Some(end - i);
}
}
{
let mut j = i;
while j < n && is_ws(chars[j]) {
j += 1;
}
let w = j - i;
if w >= 1 {
if j == n {
return Some(w); } else if w >= 2 {
return Some(w - 1); }
}
}
{
let mut j = i;
while j < n && is_ws(chars[j]) {
j += 1;
}
if j > i {
return Some(j - i);
}
}
None
}
fn split_gpt2_word(s: &str) -> Vec<String> {
let chars: Vec<char> = s.chars().collect();
let mut out = Vec::new();
let mut gap = String::new();
let mut i = 0;
while i < chars.len() {
if let Some(len) = match_gpt2_word(&chars, i) {
if !gap.is_empty() {
out.push(std::mem::take(&mut gap));
}
let tok: String = chars[i..i + len].iter().collect();
out.push(tok);
i += len;
} else {
gap.push(chars[i]);
i += 1;
}
}
if !gap.is_empty() {
out.push(gap);
}
out
}
fn match_gpt2_word(chars: &[char], i: usize) -> Option<usize> {
let n = chars.len();
if chars[i] == '\'' {
for suf in ["s", "t", "re", "ve", "m", "ll", "d"] {
let sl = suf.len(); if i + 1 + sl <= n && chars[i + 1..i + 1 + sl].iter().collect::<String>() == *suf {
return Some(1 + sl);
}
}
}
{
let mut j = i;
if chars[j] == ' '
&& let Some(&c1) = chars.get(j + 1)
&& is_l(c1)
{
j += 1;
}
let start = j;
while j < n && is_l(chars[j]) {
j += 1;
}
if j > start {
return Some(j - i);
}
}
{
let mut j = i;
if chars[j] == ' '
&& let Some(&c1) = chars.get(j + 1)
&& is_n(c1)
{
j += 1;
}
let start = j;
while j < n && is_n(chars[j]) {
j += 1;
}
if j > start {
return Some(j - i);
}
}
{
let other = |c: char| !is_ws(c) && !is_l(c) && !is_n(c);
let mut j = i;
if chars[j] == ' '
&& let Some(&c1) = chars.get(j + 1)
&& other(c1)
{
j += 1;
}
let start = j;
while j < n && other(chars[j]) {
j += 1;
}
if j > start {
return Some(j - i);
}
}
{
let mut j = i;
while j < n && is_ws(chars[j]) {
j += 1;
}
let w = j - i;
if w >= 1 {
if j == n {
return Some(w);
} else if w >= 2 {
return Some(w - 1);
}
}
}
{
let mut j = i;
while j < n && is_ws(chars[j]) {
j += 1;
}
if j > i {
return Some(j - i);
}
}
None
}
pub fn byte_level_map(s: &str) -> String {
s.bytes().map(byte_to_char).collect()
}
#[inline]
fn byte_to_char(b: u8) -> char {
let printable =
(0x21..=0x7E).contains(&b) || (0xA1..=0xAC).contains(&b) || (0xAE..=0xFF).contains(&b);
if printable {
char::from_u32(b as u32).expect("printable byte is a valid codepoint")
} else {
let mut offset = 0u32;
for x in 0u8..b {
let x_printable = (0x21..=0x7E).contains(&x)
|| (0xA1..=0xAC).contains(&x)
|| (0xAE..=0xFF).contains(&x);
if !x_printable {
offset += 1;
}
}
char::from_u32(0x100 + offset).expect("byte-level remap stays below 0x144")
}
}
#[inline]
pub fn char_to_byte(c: char) -> Option<u8> {
let cp = c as u32;
if cp < 0x100 {
let b = cp as u8;
let printable =
(0x21..=0x7E).contains(&b) || (0xA1..=0xAC).contains(&b) || (0xAE..=0xFF).contains(&b);
if printable {
return Some(b);
}
return None;
}
if (0x100..0x144).contains(&cp) {
let target = cp - 0x100;
let mut offset = 0u32;
for x in 0u8..=0xFF {
let x_printable = (0x21..=0x7E).contains(&x)
|| (0xA1..=0xAC).contains(&x)
|| (0xAE..=0xFF).contains(&x);
if !x_printable {
if offset == target {
return Some(x);
}
offset += 1;
}
}
}
None
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ranges_membership() {
assert!(is_l('A'));
assert!(is_l('é'));
assert!(is_l('一')); assert!(!is_l('5'));
assert!(is_n('5'));
assert!(is_n('²')); assert!(!is_n('A'));
assert!(is_p('!'));
assert!(is_p('.'));
assert!(!is_p('+')); assert!(is_s('+'));
assert!(is_s('$')); assert!(is_m('\u{0301}')); assert!(!is_m('a'));
}
#[test]
fn byte_level_roundtrip_is_total() {
let mut seen = std::collections::HashSet::new();
for b in 0u8..=0xFF {
let c = byte_to_char(b);
assert!(seen.insert(c), "byte-level map not injective at {b}");
assert_eq!(char_to_byte(c), Some(b), "roundtrip failed for byte {b}");
}
assert_eq!(byte_to_char(b' '), 'Ġ');
assert_eq!(byte_to_char(b'\n'), 'Ċ');
assert_eq!(byte_to_char(b'\t'), 'ĉ');
}
#[test]
fn byte_level_maps_utf8() {
let s = byte_level_map("é");
assert_eq!(s, "é");
assert_eq!(byte_level_map(" a"), "Ġa");
}
#[test]
fn digit_grouping_groups_of_three() {
assert_eq!(split_digit_groups("1234567"), vec!["123", "456", "7"]);
assert_eq!(split_digit_groups("ab12cd"), vec!["ab", "12", "cd"]);
assert_eq!(split_digit_groups("12"), vec!["12"]);
assert_eq!(split_digit_groups("abc"), vec!["abc"]);
}
#[test]
fn cjk_isolation() {
let out = split_cjk_kana("a日本b");
assert_eq!(out, vec!["a", "日本", "b"]);
}
#[test]
fn gpt_word_basic_split() {
let out = split_gpt_word("Hello world");
assert_eq!(out, vec!["Hello", " world"]);
}
#[test]
fn gpt_word_punct_run() {
let out = split_gpt_word("a...");
assert_eq!(out, vec!["a", "..."]);
let out2 = split_gpt_word(" !!!");
assert_eq!(out2, vec![" !!!"]);
}
#[test]
fn gpt_word_trailing_whitespace() {
let out = split_gpt_word("a ");
assert_eq!(out, vec!["a", " "]);
}
#[test]
fn full_pretokenize_byte_mapped() {
let p = pretokenize("Hi");
assert_eq!(p, vec!["Hi"]);
let p2 = pretokenize(" a");
assert_eq!(p2, vec!["Ġa"]);
}
#[test]
fn smollm2_digits_are_individual() {
assert_eq!(split_digits_individual("1234"), vec!["1", "2", "3", "4"]);
assert_eq!(
split_digits_individual("ab12cd"),
vec!["ab", "1", "2", "cd"]
);
assert_eq!(split_digits_individual("²x"), vec!["²", "x"]); assert_eq!(split_digits_individual("abc"), vec!["abc"]);
}
#[test]
fn gpt2_contractions_match_first() {
assert_eq!(split_gpt2_word("it's"), vec!["it", "'s"]);
assert_eq!(split_gpt2_word("we'll've"), vec!["we", "'ll", "'ve"]);
assert_eq!(split_gpt2_word("IT'S"), vec!["IT", "'", "S"]);
assert_eq!(split_gpt2_word("dogs'"), vec!["dogs", "'"]);
}
#[test]
fn gpt2_letters_have_no_mark_class() {
assert_eq!(split_gpt2_word("e\u{0301}x"), vec!["e", "\u{0301}", "x"]);
assert_eq!(split_gpt2_word("e\u{0301}"), vec!["e", "\u{0301}"]);
}
#[test]
fn gpt2_number_runs_after_digit_stage() {
assert_eq!(split_gpt2_word(" 5"), vec![" 5"]);
assert_eq!(pretokenize_smollm2("a 12"), vec!["a", "Ġ", "1", "2"]);
}
#[test]
fn gpt2_whitespace_backtracking() {
assert_eq!(split_gpt2_word("a b"), vec!["a", " ", " b"]);
assert_eq!(split_gpt2_word("a "), vec!["a", " "]);
assert_eq!(split_gpt2_word("a\n\n"), vec!["a", "\n\n"]);
assert_eq!(split_gpt2_word("a\n\nb"), vec!["a", "\n", "\n", "b"]);
}
#[test]
fn gpt2_punct_and_symbols() {
assert_eq!(split_gpt2_word(" !!!"), vec![" !!!"]);
assert_eq!(split_gpt2_word("a..."), vec!["a", "..."]);
assert_eq!(split_gpt2_word("+=<>"), vec!["+=<>"]);
}
#[test]
fn smollm2_full_scheme_byte_mapped() {
assert_eq!(
pretokenize_smollm2("Hi 42!"),
vec!["Hi", "Ġ", "4", "2", "!"]
);
assert_eq!(pretokenize_smollm2("it's"), vec!["it", "'s"]);
assert_eq!(pretokenize_smollm2("é"), vec!["é"]);
}
}