use crate::pretokenize::unicode::{self, CharClass};
#[cfg(target_arch = "x86_64")]
#[inline]
pub(crate) fn avx512_scanner_available() -> bool {
std::arch::is_x86_feature_detected!("avx512f")
&& std::arch::is_x86_feature_detected!("avx512bw")
&& std::arch::is_x86_feature_detected!("avx512vl")
&& std::arch::is_x86_feature_detected!("bmi1")
&& std::arch::is_x86_feature_detected!("bmi2")
&& std::arch::is_x86_feature_detected!("lzcnt")
&& std::arch::is_x86_feature_detected!("popcnt")
}
#[cfg(target_arch = "x86_64")]
#[inline]
pub(crate) fn avx512_fill_available() -> bool {
avx512_scanner_available() && std::arch::is_x86_feature_detected!("avx512vbmi2")
}
#[cfg(target_arch = "x86_64")]
#[inline]
pub(crate) fn avx2_scanner_available() -> bool {
std::arch::is_x86_feature_detected!("avx2")
&& std::arch::is_x86_feature_detected!("bmi1")
&& std::arch::is_x86_feature_detected!("bmi2")
&& std::arch::is_x86_feature_detected!("lzcnt")
&& std::arch::is_x86_feature_detected!("popcnt")
}
#[cfg(target_arch = "x86_64")]
#[inline]
pub(crate) fn simd_scanner_available() -> bool {
avx512_scanner_available() || avx2_scanner_available()
}
#[cfg(not(target_arch = "x86_64"))]
#[inline]
pub(crate) fn simd_scanner_available() -> bool {
cfg!(target_arch = "aarch64")
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
pub(crate) unsafe fn movemask64(
v0: std::arch::aarch64::uint8x16_t,
v1: std::arch::aarch64::uint8x16_t,
v2: std::arch::aarch64::uint8x16_t,
v3: std::arch::aarch64::uint8x16_t,
) -> u64 {
use std::arch::aarch64::*;
unsafe {
const W: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
let w = vld1q_u8(W.as_ptr());
let mut a0 = vandq_u8(v0, w);
let a1 = vandq_u8(v1, w);
let mut a2 = vandq_u8(v2, w);
let a3 = vandq_u8(v3, w);
core::arch::asm!(
"addp {a0:v}.16b, {a0:v}.16b, {a1:v}.16b",
"addp {a2:v}.16b, {a2:v}.16b, {a3:v}.16b",
"addp {a0:v}.16b, {a0:v}.16b, {a2:v}.16b",
"addp {a0:v}.16b, {a0:v}.16b, {a0:v}.16b",
a0 = inout(vreg) a0,
a1 = in(vreg) a1,
a2 = inout(vreg) a2,
a3 = in(vreg) a3,
options(pure, nomem, nostack, preserves_flags),
);
vgetq_lane_u64::<0>(vreinterpretq_u64_u8(a0))
}
}
#[derive(Clone, Copy, Default)]
pub(crate) struct AsciiMasks {
pub l: u64,
pub d: u64,
pub s: u64,
pub wt: u64,
pub n: u64,
pub hi: u64,
pub ap: u64,
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
pub(crate) fn ascii_masks(bytes: &[u8], scan: usize) -> AsciiMasks {
use std::arch::aarch64::*;
unsafe {
let p = bytes.as_ptr().add(scan);
let mut l = [vdupq_n_u8(0); 4];
let mut d = [vdupq_n_u8(0); 4];
let mut s = [vdupq_n_u8(0); 4];
let mut wt = [vdupq_n_u8(0); 4];
let mut n = [vdupq_n_u8(0); 4];
let mut hi = [vdupq_n_u8(0); 4];
let mut ap = [vdupq_n_u8(0); 4];
for i in 0..4 {
let v = vld1q_u8(p.add(16 * i));
let lowered = vorrq_u8(v, vdupq_n_u8(0x20));
l[i] = vcleq_u8(vsubq_u8(lowered, vdupq_n_u8(b'a')), vdupq_n_u8(25));
d[i] = vcleq_u8(vsubq_u8(v, vdupq_n_u8(b'0')), vdupq_n_u8(9));
s[i] = vceqq_u8(v, vdupq_n_u8(b' '));
n[i] = vorrq_u8(
vceqq_u8(v, vdupq_n_u8(b'\r')),
vceqq_u8(v, vdupq_n_u8(b'\n')),
);
wt[i] = vbicq_u8(vcleq_u8(vsubq_u8(v, vdupq_n_u8(9)), vdupq_n_u8(4)), n[i]);
hi[i] = vcltzq_s8(vreinterpretq_s8_u8(v));
ap[i] = vceqq_u8(v, vdupq_n_u8(b'\''));
}
AsciiMasks {
l: movemask64(l[0], l[1], l[2], l[3]),
d: movemask64(d[0], d[1], d[2], d[3]),
s: movemask64(s[0], s[1], s[2], s[3]),
wt: movemask64(wt[0], wt[1], wt[2], wt[3]),
n: movemask64(n[0], n[1], n[2], n[3]),
hi: movemask64(hi[0], hi[1], hi[2], hi[3]),
ap: movemask64(ap[0], ap[1], ap[2], ap[3]),
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw,avx512vl,bmi1,bmi2,lzcnt,popcnt")]
#[inline]
pub(crate) fn ascii_masks_avx512(bytes: &[u8], scan: usize) -> AsciiMasks {
use std::arch::x86_64::*;
unsafe {
let v = _mm512_loadu_si512(bytes.as_ptr().add(scan) as *const _);
let lowered = _mm512_or_si512(v, _mm512_set1_epi8(0x20));
let l = _mm512_cmple_epu8_mask(
_mm512_sub_epi8(lowered, _mm512_set1_epi8(b'a' as i8)),
_mm512_set1_epi8(25),
);
let d = _mm512_cmple_epu8_mask(
_mm512_sub_epi8(v, _mm512_set1_epi8(b'0' as i8)),
_mm512_set1_epi8(9),
);
let s = _mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(b' ' as i8));
let n = _mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(b'\r' as i8))
| _mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(b'\n' as i8));
let wt =
_mm512_cmple_epu8_mask(_mm512_sub_epi8(v, _mm512_set1_epi8(9)), _mm512_set1_epi8(4))
& !n;
let hi = _mm512_movepi8_mask(v) as u64;
let ap = _mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(b'\'' as i8));
AsciiMasks {
l,
d,
s,
wt,
n,
hi,
ap,
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,bmi1,bmi2,lzcnt,popcnt")]
#[inline(never)]
pub(crate) fn ascii_masks_avx2(bytes: &[u8], scan: usize) -> AsciiMasks {
use std::arch::x86_64::*;
unsafe {
let le =
|v: __m256i, lim: __m256i| -> __m256i { _mm256_cmpeq_epi8(_mm256_min_epu8(v, lim), v) };
let mm = |m0: __m256i, m1: __m256i| -> u64 {
(_mm256_movemask_epi8(m0) as u32 as u64)
| ((_mm256_movemask_epi8(m1) as u32 as u64) << 32)
};
let p = bytes.as_ptr().add(scan);
let v0 = _mm256_loadu_si256(p as *const _);
let v1 = _mm256_loadu_si256(p.add(32) as *const _);
let x20 = _mm256_set1_epi8(0x20);
let ca = _mm256_set1_epi8(b'a' as i8);
let c25 = _mm256_set1_epi8(25);
let l = mm(
le(_mm256_sub_epi8(_mm256_or_si256(v0, x20), ca), c25),
le(_mm256_sub_epi8(_mm256_or_si256(v1, x20), ca), c25),
);
let c0 = _mm256_set1_epi8(b'0' as i8);
let c9 = _mm256_set1_epi8(9);
let d = mm(
le(_mm256_sub_epi8(v0, c0), c9),
le(_mm256_sub_epi8(v1, c0), c9),
);
let sp = _mm256_set1_epi8(b' ' as i8);
let s = mm(_mm256_cmpeq_epi8(v0, sp), _mm256_cmpeq_epi8(v1, sp));
let cr = _mm256_set1_epi8(b'\r' as i8);
let lf = _mm256_set1_epi8(b'\n' as i8);
let n = mm(
_mm256_or_si256(_mm256_cmpeq_epi8(v0, cr), _mm256_cmpeq_epi8(v0, lf)),
_mm256_or_si256(_mm256_cmpeq_epi8(v1, cr), _mm256_cmpeq_epi8(v1, lf)),
);
let c4 = _mm256_set1_epi8(4);
let wt = mm(
le(_mm256_sub_epi8(v0, c9), c4),
le(_mm256_sub_epi8(v1, c9), c4),
) & !n;
let hi = mm(v0, v1); let apc = _mm256_set1_epi8(b'\'' as i8);
let ap = mm(_mm256_cmpeq_epi8(v0, apc), _mm256_cmpeq_epi8(v1, apc));
AsciiMasks {
l,
d,
s,
wt,
n,
hi,
ap,
}
}
}
#[inline(always)]
pub(crate) unsafe fn nn_at_full(bytes: &[u8], idx: usize) -> bool {
use super::{decode_cp_inbounds, is_ascii_ws};
let b = bytes[idx];
if b < 0x80 {
return !is_ascii_ws(b);
}
let (cp, _) = unsafe { decode_cp_inbounds(bytes, idx) };
unicode::class_of(cp) != CharClass::Whitespace
}
#[inline(always)]
pub(crate) unsafe fn char_through(
bytes: &[u8],
pos: usize,
class: impl Fn(u32) -> CharClass,
) -> (CharClass, usize, usize) {
use super::{decode_cp_inbounds, is_ascii_ws, is_digit, is_letter};
let b = bytes[pos - 1];
if b < 0x80 {
let cls = if is_letter(b) {
CharClass::Letter
} else if is_digit(b) {
CharClass::Number
} else if is_ascii_ws(b) {
CharClass::Whitespace
} else {
CharClass::Other
};
return (cls, pos - 1, pos);
}
let mut j = pos - 1;
while j > 0 && bytes[j] & 0xC0 == 0x80 {
j -= 1;
}
let (cp, l) = unsafe { decode_cp_inbounds(bytes, j) };
(class(cp), j, j + l)
}
#[derive(Clone, Copy, Default)]
pub(crate) struct UniClasses {
pub l: u64,
pub n: u64,
pub o: u64,
pub ws: u64,
pub w2: u64,
pub w3: u64,
pub lead2: u64,
pub lead3: u64,
pub lead4: u64,
pub cont: u64,
pub resid: u64,
}
#[inline(always)]
pub(crate) unsafe fn classify_uni_chars<const NUMBERS: bool, const LEADS: bool>(
bytes: &[u8],
scan: usize,
mut m: u64,
class: impl Fn(u32) -> CharClass,
) -> UniClasses {
use super::decode_cp_inbounds;
let mut u = UniClasses::default();
while m != 0 {
let i = m.trailing_zeros() as usize;
m &= m - 1;
let b = bytes[scan + i];
if b < 0xE0 {
if b < 0xC2 {
u.resid |= 1 << i; continue;
}
let lead = 1u64 << i;
let chm = 3u64 << i; let b1 = unsafe { *bytes.get_unchecked(scan + i + 1) };
let cp = ((b as u32 & 0x1F) << 6) | (b1 as u32 & 0x3F);
match class(cp) {
CharClass::Letter => u.l |= chm,
CharClass::Number => {
u.n |= chm;
if !NUMBERS {
u.resid |= chm;
}
}
CharClass::Other => u.o |= chm,
CharClass::Whitespace => {
u.ws |= chm;
if i + 2 > 64 {
u.resid |= chm;
} else {
u.w2 |= lead;
}
}
}
if LEADS {
u.lead2 |= lead;
}
u.cont |= chm & !lead;
m &= !chm;
continue;
}
let l = if b < 0xF0 { 3 } else { 4 };
let chm = ((1u64 << l) - 1) << i; let lead = 1u64 << i;
let (cp, _) = unsafe { decode_cp_inbounds(bytes, scan + i) };
match class(cp) {
CharClass::Letter => u.l |= chm,
CharClass::Number => {
u.n |= chm;
if !NUMBERS {
u.resid |= chm;
}
}
CharClass::Other => u.o |= chm,
CharClass::Whitespace => {
u.ws |= chm;
if i + l > 64 || l == 4 {
u.resid |= chm;
} else {
u.w3 |= lead;
}
}
}
if LEADS {
if l == 3 {
u.lead3 |= lead;
} else {
u.lead4 |= lead;
}
}
u.cont |= chm & !lead;
m &= !chm;
}
u
}
#[inline(always)]
pub(crate) fn digit_run_splits3(d: u64) -> u64 {
let mut b = d & !(d << 1); let mut c = d & (d >> 1) & (d >> 2) & (d >> 3);
let mut sh = 3u32;
while sh < 64 {
b |= (b & c) << sh;
c &= c >> sh;
sh <<= 1;
}
b
}
pub(crate) trait MaskScheme {
fn advance(bytes: &[u8], pos: usize) -> usize;
#[cfg(target_arch = "aarch64")]
fn batch_masks(bytes: &[u8], scan: usize) -> (u64, u64);
#[cfg(target_arch = "x86_64")]
unsafe fn batch_masks_x86<const AVX512: bool>(bytes: &[u8], scan: usize) -> (u64, u64);
#[cfg(target_arch = "x86_64")]
#[inline(always)]
fn batch_masks(bytes: &[u8], scan: usize) -> (u64, u64)
where
Self: Sized,
{
debug_assert!(simd_scanner_available());
if avx512_scanner_available() {
unsafe { batch_masks_dyn_avx512::<Self>(bytes, scan) }
} else {
unsafe { batch_masks_dyn_avx2::<Self>(bytes, scan) }
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw,avx512vl,bmi1,bmi2,lzcnt,popcnt")]
#[inline]
unsafe fn batch_masks_dyn_avx512<S: MaskScheme>(bytes: &[u8], scan: usize) -> (u64, u64) {
unsafe { S::batch_masks_x86::<true>(bytes, scan) }
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,bmi1,bmi2,lzcnt,popcnt")]
#[inline]
unsafe fn batch_masks_dyn_avx2<S: MaskScheme>(bytes: &[u8], scan: usize) -> (u64, u64) {
unsafe { S::batch_masks_x86::<false>(bytes, scan) }
}
pub(crate) const X86_TIER_DYN: u8 = 0;
pub(crate) const X86_TIER_AVX2: u8 = 1;
pub(crate) const X86_TIER_AVX512: u8 = 2;
pub(crate) const X86_TIER_AVX512_VBMI2: u8 = 3;
pub(crate) struct MaskState {
pub pos: usize,
scan: usize,
mask_base: usize,
rem: u64,
batch_usable: u64,
batch_bad: u64,
scalar_until: usize,
pre_base: usize,
pre_usable: u64,
pre_bad: u64,
}
impl MaskState {
#[inline]
pub(crate) fn new(pos: usize) -> Self {
let scalar_until = if simd_scanner_available() {
pos
} else {
usize::MAX
};
Self {
pos,
scan: pos,
mask_base: pos,
rem: 0,
batch_usable: 0,
batch_bad: 0,
scalar_until,
pre_base: usize::MAX,
pre_usable: 0,
pre_bad: 0,
}
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
fn load_segment(&mut self, from_bit: u32) {
let live = u64::MAX << from_bit;
let seg_bad = self.batch_bad & live;
if seg_bad == 0 {
self.rem = self.batch_usable & live;
self.batch_bad = 0;
} else {
let nb = seg_bad.trailing_zeros();
self.rem = self.batch_usable & live & ((1u64 << nb) - 1);
let rest = self.batch_usable & (u64::MAX << nb);
self.scalar_until = if rest != 0 {
self.mask_base + rest.trailing_zeros() as usize
} else {
self.mask_base + 64
};
}
let at_start = self.pos == self.mask_base + from_bit as usize;
self.rem &= !(u64::from(at_start) << from_bit);
}
#[inline(always)]
pub(crate) fn next_span<S: MaskScheme>(&mut self, bytes: &[u8]) -> Option<(usize, usize)> {
let len = bytes.len();
loop {
if self.rem != 0 {
let tz = self.rem.trailing_zeros() as usize;
let end = self.mask_base + tz;
self.rem &= self.rem - 1;
let start = self.pos;
self.pos = end;
return Some((start, end));
}
if self.pos < self.scalar_until {
if self.pos >= len {
return None;
}
let start = self.pos;
let end = S::advance(bytes, start);
self.pos = end;
return Some((start, end));
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
{
if self.batch_bad != 0 && self.pos < self.mask_base + 64 {
self.load_segment((self.pos - self.mask_base) as u32);
continue;
}
self.batch_bad = 0;
while self.scan + 64 <= self.pos {
self.scan += 64;
}
if self.scan + 64 > len {
self.scalar_until = usize::MAX;
continue;
}
let (usable, bad) = if self.pre_base == self.scan {
(self.pre_usable, self.pre_bad)
} else {
S::batch_masks(bytes, self.scan)
};
self.mask_base = self.scan;
self.scan += 64;
self.batch_usable = usable;
self.batch_bad = bad;
if self.scan + 64 <= len {
let (u2, b2) = S::batch_masks(bytes, self.scan);
self.pre_base = self.scan;
self.pre_usable = u2;
self.pre_bad = b2;
} else {
self.pre_base = usize::MAX;
}
if self.pos > self.mask_base {
self.load_segment((self.pos - self.mask_base) as u32);
} else {
self.load_segment(0);
}
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
{
self.scalar_until = usize::MAX;
}
}
}
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
static BIT_POS: [[u16; 8]; 256] = {
let mut t = [[0u16; 8]; 256];
let mut b = 1usize;
while b < 256 {
let mut j = 0;
let mut w = 0;
while j < 8 {
if b >> j & 1 == 1 {
t[b][w] = j as u16;
w += 1;
}
j += 1;
}
b += 1;
}
t
};
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
unsafe fn flatten_bits(m: u64, rel: u16, out: *mut u16) -> usize {
let mut x = m;
x -= (x >> 1) & 0x5555_5555_5555_5555;
x = (x & 0x3333_3333_3333_3333) + ((x >> 2) & 0x3333_3333_3333_3333);
x = (x + (x >> 4)) & 0x0F0F_0F0F_0F0F_0F0F;
let incl = x.wrapping_mul(0x0101_0101_0101_0101);
let excl = incl << 8;
#[cfg(target_arch = "aarch64")]
unsafe {
use std::arch::aarch64::*;
for j in 0..8 {
let b = (m >> (8 * j)) as u8 as usize;
let w = (excl >> (8 * j)) as u8 as usize;
let v = vld1q_u16(BIT_POS[b].as_ptr());
let v = vaddq_u16(v, vdupq_n_u16(rel.wrapping_add(8 * j as u16)));
vst1q_u16(out.add(w), v);
}
}
#[cfg(not(target_arch = "aarch64"))]
unsafe {
for j in 0..8 {
let b = (m >> (8 * j)) as u8 as usize;
let w = (excl >> (8 * j)) as u8 as usize;
let e = &BIT_POS[b];
let base = rel.wrapping_add(8 * j as u16);
for (t, &offset) in e.iter().enumerate() {
out.add(w + t).write(offset.wrapping_add(base));
}
}
}
(incl >> 56) as usize
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw,avx512vbmi2")]
#[inline]
unsafe fn flatten_bits_avx512(m: u64, rel: u16, out: *mut u16) -> usize {
use std::arch::x86_64::*;
const IOTA: [u8; 64] = {
let mut a = [0u8; 64];
let mut i = 0;
while i < 64 {
a[i] = i as u8;
i += 1;
}
a
};
unsafe {
let iota = _mm512_loadu_si512(IOTA.as_ptr() as *const _);
let comp = _mm512_maskz_compress_epi8(m, iota);
let relv = _mm512_set1_epi16(rel as i16);
let lo = _mm512_add_epi16(_mm512_cvtepu8_epi16(_mm512_castsi512_si256(comp)), relv);
let hi = _mm512_add_epi16(
_mm512_cvtepu8_epi16(_mm512_extracti64x4_epi64::<1>(comp)),
relv,
);
_mm512_storeu_si512(out as *mut _, lo);
_mm512_storeu_si512(out.add(32) as *mut _, hi);
}
m.count_ones() as usize
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
unsafe fn flatten_bits_dispatch<const X86_TIER: u8>(m: u64, rel: u16, out: *mut u16) -> usize {
#[cfg(target_arch = "x86_64")]
if X86_TIER == X86_TIER_AVX512_VBMI2 {
return unsafe { flatten_bits_avx512(m, rel, out) };
}
unsafe { flatten_bits(m, rel, out) }
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
static PACK_MASK_TABLE: [[u64; 2]; 16] = {
let mut t = [[0u64; 2]; 16];
let mut n = 1;
while n <= 15 {
let (lo, hi) = crate::pretokenize::pack_mask_halves(n);
t[n] = [lo, hi];
n += 1;
}
t
};
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
const BOUND_BUF: usize = crate::pretokenize::PRETOKEN_CHUNK + 208;
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
const REL_LIMIT: isize = u16::MAX as isize - 127;
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
impl MaskState {
#[inline(always)]
pub(crate) fn fill_spans_two_phase<'a, S: MaskScheme>(
&mut self,
bytes: &'a [u8],
batch: &mut crate::pretokenize::SpanBatch<'a>,
prefetch: &impl Fn(u64),
) -> usize {
#[cfg(target_arch = "x86_64")]
if crate::pretokenize::crc_hash_selected() {
if avx512_fill_available() {
return unsafe {
self.fill_spans_two_phase_avx512_vbmi2_crc::<S>(bytes, batch, prefetch)
};
}
if avx512_scanner_available() {
return unsafe {
self.fill_spans_two_phase_avx512_crc::<S>(bytes, batch, prefetch)
};
}
if avx2_scanner_available() {
return unsafe { self.fill_spans_two_phase_avx2_crc::<S>(bytes, batch, prefetch) };
}
return unsafe { self.fill_spans_two_phase_crc::<S>(bytes, batch, prefetch) };
}
self.fill_spans_two_phase_impl::<S, false, X86_TIER_DYN>(bytes, batch, prefetch)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(
enable = "avx512f,avx512bw,avx512vl,avx512vbmi2,bmi1,bmi2,lzcnt,popcnt,sse4.2"
)]
unsafe fn fill_spans_two_phase_avx512_vbmi2_crc<'a, S: MaskScheme>(
&mut self,
bytes: &'a [u8],
batch: &mut crate::pretokenize::SpanBatch<'a>,
prefetch: &impl Fn(u64),
) -> usize {
self.fill_spans_two_phase_impl::<S, true, X86_TIER_AVX512_VBMI2>(bytes, batch, prefetch)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw,avx512vl,bmi1,bmi2,lzcnt,popcnt,sse4.2")]
unsafe fn fill_spans_two_phase_avx512_crc<'a, S: MaskScheme>(
&mut self,
bytes: &'a [u8],
batch: &mut crate::pretokenize::SpanBatch<'a>,
prefetch: &impl Fn(u64),
) -> usize {
self.fill_spans_two_phase_impl::<S, true, X86_TIER_AVX512>(bytes, batch, prefetch)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,bmi1,bmi2,lzcnt,popcnt,sse4.2")]
unsafe fn fill_spans_two_phase_avx2_crc<'a, S: MaskScheme>(
&mut self,
bytes: &'a [u8],
batch: &mut crate::pretokenize::SpanBatch<'a>,
prefetch: &impl Fn(u64),
) -> usize {
self.fill_spans_two_phase_impl::<S, true, X86_TIER_AVX2>(bytes, batch, prefetch)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse4.2")]
unsafe fn fill_spans_two_phase_crc<'a, S: MaskScheme>(
&mut self,
bytes: &'a [u8],
batch: &mut crate::pretokenize::SpanBatch<'a>,
prefetch: &impl Fn(u64),
) -> usize {
self.fill_spans_two_phase_impl::<S, true, X86_TIER_DYN>(bytes, batch, prefetch)
}
#[inline(always)]
fn fill_spans_two_phase_impl<'a, S: MaskScheme, const X86_CRC: bool, const X86_TIER: u8>(
&mut self,
bytes: &'a [u8],
batch: &mut crate::pretokenize::SpanBatch<'a>,
prefetch: &impl Fn(u64),
) -> usize {
use crate::pretokenize::{PRETOKEN_CHUNK, fill_span_hash, pack_pretoken_key};
debug_assert!(simd_scanner_available());
let len = bytes.len();
let mut pending = self.pos;
let mut scan = self.scan;
if scan > pending {
scan -= 64 * (scan - pending).div_ceil(64);
}
let mut n = 0usize;
let pack_masks: *const [u64; 2] = std::hint::black_box(PACK_MASK_TABLE.as_ptr());
'refill: while n < PRETOKEN_CHUNK && pending < len {
if pending >= scan + 64 {
scan += 64 * ((pending - scan) / 64);
}
let fill_base = pending;
let needed = PRETOKEN_CHUNK - n;
let mut buf = [std::mem::MaybeUninit::<u16>::uninit(); BOUND_BUF];
let bufp = buf.as_mut_ptr() as *mut u16;
let mut nb = 0usize;
let mut resume = pending;
let mut exhausted = false;
let mut overflow_end: Option<usize> = None;
'harvest: while nb < needed {
if scan.wrapping_sub(fill_base) as isize > REL_LIMIT {
break; }
if scan + 64 > len {
let mut p = if nb > 0 {
fill_base + unsafe { *bufp.add(nb - 1) } as usize
} else {
fill_base
};
while p < len && nb < needed {
p = S::advance(bytes, p);
unsafe { bufp.add(nb).write((p - fill_base) as u16) };
nb += 1;
}
exhausted = p >= len;
break;
}
let base = scan;
#[cfg(target_arch = "x86_64")]
let (usable, bad) = match X86_TIER {
X86_TIER_AVX512 | X86_TIER_AVX512_VBMI2 => unsafe {
S::batch_masks_x86::<true>(bytes, base)
},
X86_TIER_AVX2 => unsafe { S::batch_masks_x86::<false>(bytes, base) },
_ => S::batch_masks(bytes, base),
};
#[cfg(not(target_arch = "x86_64"))]
let (usable, bad) = S::batch_masks(bytes, base);
let (mut ulive, mut blive) = if resume >= base {
debug_assert!(resume - base < 64);
let k = resume - base;
((u64::MAX << k) << 1, u64::MAX << k)
} else {
(u64::MAX, u64::MAX)
};
let rel = base.wrapping_sub(fill_base) as u16;
if bad & blive == 0 {
debug_assert!(
nb + if X86_TIER == X86_TIER_AVX512_VBMI2 {
128
} else {
72
} <= BOUND_BUF
);
nb += unsafe {
flatten_bits_dispatch::<X86_TIER>(usable & ulive, rel, bufp.add(nb))
};
scan = base + 64;
continue;
}
loop {
let seg_bad = bad & blive;
if seg_bad == 0 {
debug_assert!(
nb + if X86_TIER == X86_TIER_AVX512_VBMI2 {
128
} else {
72
} <= BOUND_BUF
);
nb += unsafe {
flatten_bits_dispatch::<X86_TIER>(usable & ulive, rel, bufp.add(nb))
};
scan = base + 64;
break;
}
let fb = seg_bad.trailing_zeros();
let prefix = usable & ulive & !(u64::MAX << fb);
debug_assert!(
nb + if X86_TIER == X86_TIER_AVX512_VBMI2 {
128
} else {
72
} <= BOUND_BUF
);
nb += unsafe { flatten_bits_dispatch::<X86_TIER>(prefix, rel, bufp.add(nb)) };
let mut p = if nb > 0 {
fill_base + unsafe { *bufp.add(nb - 1) } as usize
} else {
fill_base
};
let rest = usable & (u64::MAX << fb);
let until = if rest != 0 {
base + rest.trailing_zeros() as usize
} else {
base + 64
};
while p < until {
p = S::advance(bytes, p);
let relp = p - fill_base;
if relp > u16::MAX as usize {
overflow_end = Some(p);
break 'harvest;
}
debug_assert!(nb < BOUND_BUF);
unsafe { bufp.add(nb).write(relp as u16) };
nb += 1;
}
if p >= base + 64 {
scan = base + 64 * ((p - base) / 64);
resume = p;
break;
}
blive = u64::MAX << (p - base);
ulive = blive << 1;
}
}
if nb == 0 {
debug_assert!(!exhausted);
let end = overflow_end.unwrap_or_else(|| S::advance(bytes, fill_base));
let span = &bytes[fill_base..end];
let (key, h) = match pack_pretoken_key(span) {
Some(key) => (key, fill_span_hash::<X86_CRC>(key)),
None => (0, 0),
};
prefetch(h);
let meta = if key != 0 { h } else { span.len() as u64 };
batch.entries[n] = crate::pretokenize::BatchEntry {
key,
ptr: span.as_ptr(),
meta,
};
n += 1;
pending = end;
continue 'refill;
}
let emit_n = nb.min(needed);
let last_end = unsafe { *bufp.add(emit_n - 1) } as usize;
let entries = &mut batch.entries[n..n + emit_n];
let base_ptr = unsafe { bytes.as_ptr().add(fill_base) };
let mut prev = 0usize;
if fill_base + last_end + 16 <= len {
for (i, e) in entries.iter_mut().enumerate() {
let end = unsafe { *bufp.add(i) } as usize;
let tok_len = end - prev;
let p = unsafe { base_ptr.add(prev) };
prev = end;
let raw = unsafe { (p as *const u128).read_unaligned() };
let m = tok_len.min(15);
let [mask_lo, mask_hi] = unsafe { *pack_masks.add(m) };
let keep = ((tok_len <= 15) as u64).wrapping_neg();
let klo = (raw as u64) & mask_lo & keep;
let khi = (((raw >> 64) as u64 & mask_hi) | ((m as u64) << 56)) & keep;
let key = (klo as u128) | ((khi as u128) << 64);
let hv = fill_span_hash::<X86_CRC>(key);
prefetch(hv);
let meta = (hv & keep) | (tok_len as u64 & !keep);
e.key = key;
e.ptr = p;
e.meta = meta;
}
} else {
for (i, e) in entries.iter_mut().enumerate() {
let end = unsafe { *bufp.add(i) } as usize;
let tok_len = end - prev;
let p = unsafe { base_ptr.add(prev) };
prev = end;
let span = unsafe { std::slice::from_raw_parts(p, tok_len) };
let (key, hv) = match pack_pretoken_key(span) {
Some(key) => (key, fill_span_hash::<X86_CRC>(key)),
None => (0, 0),
};
prefetch(hv);
let meta = if key != 0 { hv } else { tok_len as u64 };
e.key = key;
e.ptr = p;
e.meta = meta;
}
}
n += emit_n;
pending = fill_base + prev;
if exhausted {
debug_assert_eq!(pending, len);
break;
}
}
if scan > pending {
scan -= 64 * (scan - pending).div_ceil(64);
}
self.pos = pending;
self.scan = scan;
self.mask_base = scan;
self.rem = 0;
self.batch_usable = 0;
self.batch_bad = 0;
self.scalar_until = pending;
self.pre_base = usize::MAX;
n
}
}