use super::mask::{self, MaskScheme, MaskState};
use super::{
decode_cp, is_ascii_ws, is_digit, is_letter, scan_digits_from, scan_letters_from,
scan_other_from,
};
use crate::pretokenize::unicode::{self, CharClass};
#[cfg(target_arch = "aarch64")]
#[inline(always)]
fn batch_masks(bytes: &[u8], scan: usize) -> (u64, u64) {
use std::arch::aarch64::*;
let len = bytes.len();
if scan + 70 > len {
return (0, u64::MAX);
}
unsafe {
let p = bytes.as_ptr().add(scan);
let zero = vdupq_n_u8(0);
let mut lv = [zero; 4];
let mut dv = [zero; 4];
let mut sv = [zero; 4];
let mut wsv = [zero; 4];
let mut hiv = [zero; 4];
let mut apv = [zero; 4];
for i in 0..4 {
let v = vld1q_u8(p.add(16 * i));
let lowered = vorrq_u8(v, vdupq_n_u8(0x20));
lv[i] = vcleq_u8(vsubq_u8(lowered, vdupq_n_u8(b'a')), vdupq_n_u8(25));
dv[i] = vcleq_u8(vsubq_u8(v, vdupq_n_u8(b'0')), vdupq_n_u8(9));
sv[i] = vceqq_u8(v, vdupq_n_u8(b' '));
wsv[i] = vorrq_u8(sv[i], vcleq_u8(vsubq_u8(v, vdupq_n_u8(9)), vdupq_n_u8(4)));
hiv[i] = vcltzq_s8(vreinterpretq_s8_u8(v));
apv[i] = vceqq_u8(v, vdupq_n_u8(b'\''));
}
let lb = mask::movemask64(lv[0], lv[1], lv[2], lv[3]);
let db = mask::movemask64(dv[0], dv[1], dv[2], dv[3]);
let s64 = mask::movemask64(sv[0], sv[1], sv[2], sv[3]);
let wsa = mask::movemask64(wsv[0], wsv[1], wsv[2], wsv[3]);
let ap_any = vorrq_u8(vorrq_u8(apv[0], apv[1]), vorrq_u8(apv[2], apv[3]));
let ap64 = if vmaxvq_u8(ap_any) != 0 {
mask::movemask64(apv[0], apv[1], apv[2], apv[3])
} else {
0
};
let hi_any = vorrq_u8(vorrq_u8(hiv[0], hiv[1]), vorrq_u8(hiv[2], hiv[3]));
if vmaxvq_u8(hi_any) != 0 {
let hi64 = mask::movemask64(hiv[0], hiv[1], hiv[2], hiv[3]);
return extended_masks(bytes, scan, lb, db, s64, wsa, hi64, ap64);
}
ascii_batch_algebra(bytes, scan, lb, db, s64, wsa, ap64)
}
}
#[cfg(target_arch = "x86_64")]
#[inline(always)]
unsafe fn batch_masks_x86<const AVX512: bool>(bytes: &[u8], scan: usize) -> (u64, u64) {
let len = bytes.len();
if scan + 70 > len {
return (0, u64::MAX);
}
let am = if AVX512 {
unsafe { mask::ascii_masks_avx512(bytes, scan) }
} else {
unsafe { mask::ascii_masks_avx2(bytes, scan) }
};
let wsa = am.s | am.wt | am.n;
if am.hi != 0 {
return unsafe { extended_masks(bytes, scan, am.l, am.d, am.s, wsa, am.hi, am.ap) };
}
ascii_batch_algebra(bytes, scan, am.l, am.d, am.s, wsa, am.ap)
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
fn ascii_batch_algebra(
bytes: &[u8],
scan: usize,
lb: u64,
db: u64,
s64: u64,
wsa: u64,
ap64: u64,
) -> (u64, u64) {
let ob = !(lb | db | wsa);
let (pl, pd, ps, pws, po) = if scan == 0 {
(0, 0, 0, 0, 0)
} else {
carries_at(bytes, scan)
};
let cont_same = (lb & ((lb << 1) | pl)) | (db & ((db << 1) | pd)) | (ob & ((ob << 1) | po));
let after_sp = (s64 << 1) | ps;
let nb = !wsa & !cont_same & !after_sp;
let mut split_ok = wsa & (!wsa >> 1); let nb64 = bytes[scan + 64]; if nb64 < 0x80 {
split_ok |= (u64::from(!is_ascii_ws(nb64)) << 63) & wsa;
} else if wsa >> 63 != 0
&& unsafe { mask::nn_at_full(bytes, scan + 64) }
{
split_ok |= 1 << 63;
}
let pwsb = (wsa << 1) | pws;
let wsboundary = wsa & (!pwsb | split_ok);
let mut boundary = nb | wsboundary;
let mut bad = 0u64;
if ap64 != 0 {
let mut cand = ap64 & boundary;
while cand != 0 {
let i = cand.trailing_zeros() as usize;
cand &= cand - 1;
if i >= 61 {
bad |= u64::MAX << i;
break;
}
let k = match bytes[scan + i + 1] {
b's' | b'd' | b'm' | b't' => 2,
b'l' if bytes[scan + i + 2] == b'l' => 3,
b'v' if bytes[scan + i + 2] == b'e' => 3,
b'r' if bytes[scan + i + 2] == b'e' => 3,
_ => 0,
};
if k != 0 {
boundary &= !(1u64 << (i + 1));
boundary |= 1u64 << (i + k);
}
}
}
(boundary & !bad, bad)
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[cfg_attr(
target_arch = "x86_64",
target_feature(enable = "bmi1,bmi2,lzcnt,popcnt")
)]
#[inline(never)]
#[allow(clippy::too_many_arguments)]
fn extended_masks(
bytes: &[u8],
scan: usize,
l64: u64,
d64: u64,
s64: u64,
ws64: u64,
hi64: u64,
ap64: u64,
) -> (u64, u64) {
let wsa = ws64;
let ct = unicode::ClassTable::get();
let class = move |cp| ct.class_of(cp);
let mut claim = mask::UniClasses::default();
let (pl, pd, ps, pws, po) = if scan == 0 {
(0, 0, 0, 0, 0)
} else if bytes[scan - 1] < 0x80 {
carries_at(bytes, scan)
} else {
let (cls, _lead, end) = unsafe { mask::char_through(bytes, scan, class) };
let chm = if end > scan {
(1u64 << (end - scan)) - 1
} else {
0
};
claim.cont = chm;
match cls {
CharClass::Letter => {
claim.l = chm;
(1, 0, 0, 0, 0)
}
CharClass::Number => {
claim.n = chm;
(0, 1, 0, 0, 0)
}
CharClass::Other => {
claim.o = chm;
(0, 0, 0, 0, 1)
}
CharClass::Whitespace => {
claim.ws = chm;
claim.resid = chm;
(0, 0, u64::from(bytes[scan - 1] == b' '), 1, 0)
}
}
};
let uni =
unsafe { mask::classify_uni_chars::<true, false>(bytes, scan, hi64 & !claim.cont, class) };
let lb = l64 | claim.l | uni.l;
let db = d64 | claim.n | uni.n;
let wsb = wsa | claim.ws | uni.ws;
let ob = !(l64 | d64 | wsa | hi64) | claim.o | uni.o;
let contm = claim.cont | uni.cont;
let resid = claim.resid | uni.resid;
let cont_same = (lb & ((lb << 1) | pl)) | (db & ((db << 1) | pd)) | (ob & ((ob << 1) | po));
let after_sp = (s64 << 1) | ps;
let nb = !wsb & !cont_same & !after_sp & !contm;
let nn = !wsb;
let mut split_ok = (wsa & (nn >> 1)) | (uni.w2 & (nn >> 2)) | (uni.w3 & (nn >> 3));
let ws_leads = wsa | uni.w2 | uni.w3;
let edge_mb = (uni.w2 & (1 << 62)) | (uni.w3 & (1 << 61));
let nb64 = bytes[scan + 64]; if nb64 < 0x80 && edge_mb == 0 {
split_ok = (split_ok & !(1 << 63)) | ((u64::from(!is_ascii_ws(nb64)) << 63) & wsa);
} else {
let edge = edge_mb | ((1 << 63) & wsa);
if edge != 0 {
if unsafe { mask::nn_at_full(bytes, scan + 64) } {
split_ok |= edge;
} else {
split_ok &= !edge;
}
}
}
let pwsb = (wsb << 1) | pws;
let wsboundary = ws_leads & (!pwsb | split_ok);
let mut boundary = nb | wsboundary;
let mut bad = resid | resid << 1 | resid >> 1;
let mut cand = ap64 & boundary & !bad;
while cand != 0 {
let i = cand.trailing_zeros() as usize;
cand &= cand - 1;
if i >= 61 {
bad |= u64::MAX << i;
break;
}
let k = match bytes[scan + i + 1] {
b's' | b'd' | b'm' | b't' => 2,
b'l' if bytes[scan + i + 2] == b'l' => 3,
b'v' if bytes[scan + i + 2] == b'e' => 3,
b'r' if bytes[scan + i + 2] == b'e' => 3,
_ => 0,
};
if k != 0 {
boundary &= !(1u64 << (i + 1));
boundary |= 1u64 << (i + k);
}
}
(boundary & !bad, bad)
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
fn carries_at(bytes: &[u8], scan: usize) -> (u64, u64, u64, u64, u64) {
let b = bytes[scan - 1];
if b < 0x80 {
let (l, d, w) = (is_letter(b), is_digit(b), is_ascii_ws(b));
let bit = |c: bool| u64::from(c);
return (bit(l), bit(d), bit(b == b' '), bit(w), bit(!l && !d && !w));
}
match unsafe { mask::char_through(bytes, scan, unicode::class_of) }.0 {
CharClass::Letter => (1, 0, 0, 0, 0),
CharClass::Number => (0, 1, 0, 0, 0),
CharClass::Whitespace => (0, 0, 0, 1, 0),
CharClass::Other => (0, 0, 0, 0, 1),
}
}
pub(crate) struct R50kScheme;
impl MaskScheme for R50kScheme {
#[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) {
batch_masks(bytes, scan)
}
#[cfg(target_arch = "x86_64")]
#[inline(always)]
unsafe fn batch_masks_x86<const AVX512: bool>(bytes: &[u8], scan: usize) -> (u64, u64) {
unsafe { batch_masks_x86::<AVX512>(bytes, scan) }
}
}
pub struct FastR50kPretokenizer<'a> {
bytes: &'a [u8],
state: MaskState,
}
super::impl_mask_pretokenizer!(FastR50kPretokenizer, R50kScheme);
#[inline(always)]
fn advance_pos(bytes: &[u8], start: usize) -> usize {
let len = bytes.len();
let b0 = unsafe { *bytes.get_unchecked(start) };
if is_letter(b0) {
return scan_letters_from(bytes, start + 1);
}
if b0 == b' ' {
if start + 1 < len {
let b1 = unsafe { *bytes.get_unchecked(start + 1) };
if is_letter(b1) {
return scan_letters_from(bytes, start + 2);
}
if is_digit(b1) {
return scan_digits_from(bytes, start + 2);
}
if b1 >= 0x80 {
let (cp, l) = unsafe { decode_cp(bytes, start + 1) };
let p = start + 1 + l;
return match unicode::class_of(cp) {
CharClass::Letter => scan_letters_from(bytes, p),
CharClass::Number => scan_digits_from(bytes, p),
CharClass::Whitespace => advance_ws(bytes, p, start),
CharClass::Other => scan_other_from(bytes, p),
};
}
if is_ascii_ws(b1) {
return advance_ws(bytes, start + 1, start);
}
return scan_other_from(bytes, start + 2);
}
return start + 1;
}
if b0 >= 0x80 {
let (cp, l) = unsafe { decode_cp(bytes, start) };
let p = start + l;
return match unicode::class_of(cp) {
CharClass::Letter => scan_letters_from(bytes, p),
CharClass::Number => scan_digits_from(bytes, p),
CharClass::Whitespace => advance_ws(bytes, p, start),
CharClass::Other => scan_other_from(bytes, p),
};
}
if is_digit(b0) {
return scan_digits_from(bytes, start + 1);
}
if b0 == b'\'' {
match bytes.get(start + 1) {
Some(b's' | b'd' | b'm' | b't') => return start + 2,
Some(b'l') if bytes.get(start + 2) == Some(&b'l') => return start + 3,
Some(b'v') if bytes.get(start + 2) == Some(&b'e') => return start + 3,
Some(b'r') if bytes.get(start + 2) == Some(&b'e') => return start + 3,
_ => return scan_other_from(bytes, start + 1),
}
}
if b0.wrapping_sub(9) < 5 {
return advance_ws(bytes, start + 1, start);
}
scan_other_from(bytes, start + 1)
}
#[inline(always)]
fn advance_ws(bytes: &[u8], scan_pos: usize, token_start: usize) -> usize {
let len = bytes.len();
let mut p = scan_pos;
while p < len {
let b = unsafe { *bytes.get_unchecked(p) };
if is_ascii_ws(b) {
p += 1;
} else if b >= 0x80 {
let (cp, l) = unsafe { decode_cp(bytes, p) };
if unicode::class_of(cp) == CharClass::Whitespace {
p += l;
} else {
break;
}
} else {
break;
}
}
if p < len {
let ws_bytes = p - token_start;
if ws_bytes >= 2 {
let mut last = p - 1;
while last > token_start && unsafe { *bytes.get_unchecked(last) } & 0xC0 == 0x80 {
last -= 1;
}
if last > token_start {
return last;
}
}
}
p
}