toklen 0.1.0

A single-threaded, lightweight, and fast token counter.
Documentation
use fancy_regex::Regex;
use serde::Deserialize;

use crate::pre_tokenized::{PreTokenizedString, Split as PtSplit};

use crate::pre_tokenizer::Error;

/// GPT-2 byte-to-unicode mapping table.
pub(crate) const BYTE_TO_CHAR: [char; 256] = build_byte_to_char();

const fn build_byte_to_char() -> [char; 256] {
    let mut table = ['\0'; 256];
    let mut next: u32 = 256;
    let mut i: u16 = 0;
    while i < 256 {
        let b = i as u8;
        let nice = (b >= b'!' && b <= b'~') || (b >= 0xA1 && b <= 0xAC) || b >= 0xAE;
        let cp = if nice {
            i as u32
        } else {
            let cp = next;
            next += 1;
            cp
        };
        table[i as usize] = char::from_u32(cp).expect("valid codepoint");
        i += 1;
    }
    table
}

/// Pre-computed UTF-8 encoding of each byte-to-char mapping.
const BYTE_TO_UTF8: [[u8; 2]; 256] = build_byte_to_utf8();

/// Length of each entry in [`BYTE_TO_UTF8`]: 1 for ASCII, 2 otherwise.
const BYTE_TO_UTF8_LEN: [u8; 256] = build_byte_to_utf8_len();

const fn build_byte_to_utf8() -> [[u8; 2]; 256] {
    let mut table = [[0u8; 2]; 256];
    let mut i: u16 = 0;
    while i < 256 {
        let cp = BYTE_TO_CHAR[i as usize] as u32;
        if cp < 0x80 {
            table[i as usize] = [cp as u8, 0];
        } else {
            table[i as usize] = [(0xC0 | (cp >> 6)) as u8, (0x80 | (cp & 0x3F)) as u8];
        }
        i += 1;
    }
    table
}

const fn build_byte_to_utf8_len() -> [u8; 256] {
    let mut table = [0u8; 256];
    let mut i: u16 = 0;
    while i < 256 {
        let cp = BYTE_TO_CHAR[i as usize] as u32;
        table[i as usize] = if cp < 0x80 { 1 } else { 2 };
        i += 1;
    }
    table
}

/// Encode an entire byte slice into GPT-2 byte-level characters.
///
/// # Safety
/// `out` must have at least `src.len() * 2` bytes of remaining capacity.
unsafe fn encode_bytes_bulk(src: &[u8], out: &mut Vec<u8>) {
    let mut pos = out.len();
    let base = out.as_mut_ptr();
    for &b in src {
        let utf8 = BYTE_TO_UTF8[b as usize];
        let len = BYTE_TO_UTF8_LEN[b as usize] as usize;
        unsafe {
            std::ptr::copy_nonoverlapping(utf8.as_ptr(), base.add(pos), len);
        }
        pos += len;
    }
    unsafe { out.set_len(pos) };
}

/// Append the byte-level encoding of `s` to `out`.
pub(crate) fn encode_bytes_into(s: &str, out: &mut String) {
    unsafe {
        let v = out.as_mut_vec();
        v.reserve(s.len() << 1);
        encode_bytes_bulk(s.as_bytes(), v);
    }
}

/// Append the byte-level encoding of `s` to `out` without checking capacity.
///
/// # Safety
/// `out` must have at least `s.len() * 2` bytes of spare capacity.
unsafe fn encode_bytes_into_unchecked(s: &str, out: &mut String) {
    unsafe { encode_bytes_bulk(s.as_bytes(), out.as_mut_vec()) };
}

/// GPT-2 pretokenization regex.
const GPT2_PATTERN: &str = concat!(
    r"'(?i:[sdmt])",
    r"|'(?i:ll|ve|re)",
    r"| ?\p{L}+",
    r"| ?\p{N}+",
    r"| ?[^\s\p{L}\p{N}]+",
    r"|\s+(?!\S)",
    r"|\s+",
);

#[inline(always)]
const fn default_true() -> bool {
    true
}

/// Raw deserialization helper for [`ByteLevel`].
#[derive(Deserialize)]
struct ByteLevelRaw {
    #[serde(default = "default_true")]
    add_prefix_space: bool,
    #[serde(default = "default_true")]
    trim_offsets: bool,
    #[serde(default = "default_true")]
    use_regex: bool,
}

/// A compiled ByteLevel pre-tokenizer.
#[derive(Clone, Debug, Deserialize)]
#[serde(try_from = "ByteLevelRaw")]
pub struct ByteLevel {
    regex: Option<Regex>,
    add_prefix_space: bool,
    #[allow(dead_code)]
    trim_offsets: bool,
}

impl TryFrom<ByteLevelRaw> for ByteLevel {
    type Error = Error;

    fn try_from(raw: ByteLevelRaw) -> Result<Self, Error> {
        let regex = if raw.use_regex {
            Some(Regex::new(GPT2_PATTERN)?)
        } else {
            None
        };
        Ok(Self {
            regex,
            add_prefix_space: raw.add_prefix_space,
            trim_offsets: raw.trim_offsets,
        })
    }
}

impl ByteLevel {
    /// Returns `true` when this instance does no regex splitting or prefix-space
    /// insertion — only bulk byte encoding.
    pub fn is_bulk_only(&self) -> bool {
        self.regex.is_none() && !self.add_prefix_space
    }

    /// Pre-tokenize in place using byte-level encoding.
    pub fn pre_tokenize(&self, pts: &mut PreTokenizedString) -> Result<(), Error> {
        if self.regex.is_none() && !self.add_prefix_space {
            return self.pre_tokenize_bulk(pts);
        }

        let old_buf = pts.buffer();
        let mut new_buf = String::with_capacity(old_buf.len().saturating_mul(2));
        let mut new_splits = Vec::with_capacity(pts.splits().len() << 2);

        for split in pts.splits() {
            let text = pts.split_text(split);

            if split.token_id.is_some() {
                let start = new_buf.len();
                encode_bytes_into(text, &mut new_buf);
                let end = new_buf.len();
                new_splits.push(PtSplit {
                    range: start..end,
                    token_id: split.token_id,
                });
                continue;
            }

            if text.is_empty() {
                continue;
            }

            let prefixed;
            let text = if self.add_prefix_space && !text.starts_with(' ') {
                prefixed = format!(" {text}");
                prefixed.as_str()
            } else {
                text
            };

            match &self.regex {
                Some(re) => {
                    for m in re.find_iter(text) {
                        let m = m?;
                        if m.start() < m.end() {
                            let start = new_buf.len();
                            encode_bytes_into(&text[m.start()..m.end()], &mut new_buf);
                            let end = new_buf.len();
                            new_splits.push(PtSplit {
                                range: start..end,
                                token_id: None,
                            });
                        }
                    }
                }
                None => {
                    let start = new_buf.len();
                    encode_bytes_into(text, &mut new_buf);
                    let end = new_buf.len();
                    if start < end {
                        new_splits.push(PtSplit {
                            range: start..end,
                            token_id: None,
                        });
                    }
                }
            }
        }

        pts.set_buffer(new_buf, new_splits);
        Ok(())
    }

    /// Encode the buffer and remap split ranges without per-split overhead.
    fn pre_tokenize_bulk(&self, pts: &mut PreTokenizedString) -> Result<(), Error> {
        let old_buf = pts.buffer();
        let mut new_buf = String::with_capacity(old_buf.len() << 1);
        let mut new_splits = Vec::with_capacity(pts.splits().len());

        for split in pts.splits() {
            let text = pts.split_text(split);
            if text.is_empty() && split.token_id.is_none() {
                continue;
            }
            let start = new_buf.len();
            unsafe { encode_bytes_into_unchecked(text, &mut new_buf) };
            let end = new_buf.len();
            new_splits.push(PtSplit {
                range: start..end,
                token_id: split.token_id,
            });
        }

        pts.set_buffer(new_buf, new_splits);
        Ok(())
    }
}