#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[cfg(target_arch = "aarch64")]
use super::cl100k_family::batch_masks;
#[cfg(target_arch = "x86_64")]
use super::cl100k_family::batch_masks_x86;
use super::mask::{MaskScheme, MaskState};
use super::{
decode_cp, is_ascii_ws, is_digit, is_letter, letter_end_at, scan_letters_from, scan_newlines,
scan_numbers_max3, scan_other_from,
};
use crate::pretokenize::unicode::{self, CharClass};
pub(crate) struct Olmo3Scheme;
impl MaskScheme for Olmo3Scheme {
#[inline(always)]
fn advance(bytes: &[u8], pos: usize) -> usize {
advance_pos(bytes, pos)
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
fn batch_masks(bytes: &[u8], scan: usize) -> (u64, u64) {
let ct = unicode::ClassTable::get();
batch_masks(bytes, scan, true, move |cp| ct.class_of(cp))
}
#[cfg(target_arch = "x86_64")]
#[inline(always)]
unsafe fn batch_masks_x86<const AVX512: bool>(bytes: &[u8], scan: usize) -> (u64, u64) {
let ct = unicode::ClassTable::get();
unsafe { batch_masks_x86::<AVX512>(bytes, scan, true, move |cp| ct.class_of(cp)) }
}
}
pub struct FastOlmo3Pretokenizer<'a> {
bytes: &'a [u8],
state: MaskState,
}
super::impl_mask_pretokenizer!(FastOlmo3Pretokenizer, Olmo3Scheme);
#[inline(always)]
fn ws_token_end(bytes: &[u8], start: usize) -> usize {
super::whitespace_token_end::<true>(bytes, start, |codepoint| {
unicode::class_of(codepoint) == CharClass::Whitespace
})
}
#[inline(always)]
fn advance_pos(bytes: &[u8], pos: usize) -> usize {
let b0 = unsafe { *bytes.get_unchecked(pos) };
if is_letter(b0) {
return scan_letters_from(bytes, pos + 1);
}
if b0 == b' ' {
let Some(&b1) = bytes.get(pos + 1) else {
return pos + 1; };
if is_letter(b1) {
return scan_letters_from(bytes, pos + 2); }
if b1 < 0x80 {
if is_digit(b1) {
return pos + 1; }
if is_ascii_ws(b1) {
return ws_token_end(bytes, pos);
}
let p = scan_other_from(bytes, pos + 2);
return scan_newlines(bytes, p);
}
let (cp, l) = unsafe { decode_cp(bytes, pos + 1) };
let p1 = pos + 1 + l;
match unicode::class_of(cp) {
CharClass::Letter => return scan_letters_from(bytes, p1),
CharClass::Whitespace => return ws_token_end(bytes, pos),
CharClass::Number => return pos + 1,
CharClass::Other => {
let p = scan_other_from(bytes, p1);
return scan_newlines(bytes, p);
}
}
}
if b0 >= 0x80 {
let (cp, l) = unsafe { decode_cp(bytes, pos) };
let p0 = pos + l;
let class = unicode::class_of(cp);
if class == CharClass::Letter {
return scan_letters_from(bytes, p0);
}
if class == CharClass::Number {
return scan_numbers_max3(bytes, p0, 1);
}
if let Some(p) = letter_end_at(bytes, p0) {
return scan_letters_from(bytes, p);
}
if class == CharClass::Whitespace {
return ws_token_end(bytes, pos);
}
let p = scan_other_from(bytes, p0);
return scan_newlines(bytes, p);
}
if is_digit(b0) {
return scan_numbers_max3(bytes, pos + 1, 1);
}
if b0 == b'\'' {
if let Some(end) = super::contraction_end(bytes, pos) {
return end;
}
if let Some(p) = letter_end_at(bytes, pos + 1) {
return scan_letters_from(bytes, p);
}
let p = scan_other_from(bytes, pos + 1);
return scan_newlines(bytes, p);
}
if b0 == b'\r' || b0 == b'\n' {
return ws_token_end(bytes, pos);
}
if is_ascii_ws(b0) {
if let Some(p) = letter_end_at(bytes, pos + 1) {
return scan_letters_from(bytes, p);
}
return ws_token_end(bytes, pos);
}
if let Some(p) = letter_end_at(bytes, pos + 1) {
return scan_letters_from(bytes, p);
}
let p = scan_other_from(bytes, pos + 1);
scan_newlines(bytes, p)
}