use super::{
decode_cp, is_ascii_ws, is_digit, is_letter, scan_newlines, scan_numbers_max3,
swar_scan_letters,
};
use crate::pretokenize::Pretoken;
use crate::pretokenize::unicode::{DsCharClass, ds_class_of, is_deepseek_cjk};
pub struct FastDeepSeekV3Pretokenizer<'a> {
bytes: &'a [u8],
pos: usize,
}
impl<'a> FastDeepSeekV3Pretokenizer<'a> {
#[inline]
pub fn new(bytes: &'a [u8]) -> Self {
Self { bytes, pos: 0 }
}
#[inline]
pub fn with_pos(bytes: &'a [u8], pos: usize) -> Self {
Self { bytes, pos }
}
#[inline]
pub fn pos(&self) -> usize {
self.pos
}
}
impl<'a> Iterator for FastDeepSeekV3Pretokenizer<'a> {
type Item = Pretoken<'a>;
#[inline]
fn next(&mut self) -> Option<Pretoken<'a>> {
if self.pos >= self.bytes.len() {
return None;
}
let start = self.pos;
self.pos = advance_pos(self.bytes, start);
Some(Pretoken(&self.bytes[start..self.pos]))
}
}
unsafe impl<'a> crate::pretokenize::PretokenSpans<'a> for FastDeepSeekV3Pretokenizer<'a> {
#[inline(never)]
fn fill_spans_keyed(
&mut self,
batch: &mut crate::pretokenize::SpanBatch<'a>,
prefetch: &impl Fn(u64),
) -> usize {
let (bytes, len) = (self.bytes, self.bytes.len());
let mut pos = self.pos;
let n = crate::pretokenize::fill_spans_keyed_with_buf(
bytes,
|| {
if pos >= len {
return None;
}
let start = pos;
pos = advance_pos(bytes, start);
Some((start, pos))
},
batch,
prefetch,
);
self.pos = pos;
n
}
}
#[inline(always)]
fn lm_end_at(bytes: &[u8], pos: usize, cjk_region: bool) -> Option<usize> {
let &b = bytes.get(pos)?;
if b < 0x80 {
if !cjk_region && is_letter(b) {
return Some(pos + 1);
}
return None;
}
let (cp, l) = unsafe { decode_cp(bytes, pos) };
if is_deepseek_cjk(cp) != cjk_region {
return None;
}
match ds_class_of(cp) {
DsCharClass::Letter | DsCharClass::Mark => Some(pos + l),
_ => None,
}
}
#[inline(always)]
fn scan_lm_from(bytes: &[u8], pos: usize, cjk_region: bool) -> usize {
let len = bytes.len();
let mut p = pos;
loop {
if !cjk_region {
p = swar_scan_letters(bytes, p);
}
if p >= len || unsafe { *bytes.get_unchecked(p) } < 0x80 {
return p; }
let (cp, l) = unsafe { decode_cp(bytes, p) };
if is_deepseek_cjk(cp) != cjk_region {
return p;
}
match ds_class_of(cp) {
DsCharClass::Letter | DsCharClass::Mark => p += l,
_ => return p,
}
}
}
#[inline(always)]
fn scan_ps_from(bytes: &[u8], pos: usize, cjk_region: bool) -> usize {
let len = bytes.len();
let mut p = pos;
loop {
if !cjk_region {
while p < len {
let b = unsafe { *bytes.get_unchecked(p) };
if b >= 0x80 {
break;
}
if !b.is_ascii_punctuation() {
return p;
}
p += 1;
}
}
if p >= len || unsafe { *bytes.get_unchecked(p) } < 0x80 {
return p;
}
let (cp, l) = unsafe { decode_cp(bytes, p) };
if is_deepseek_cjk(cp) != cjk_region || ds_class_of(cp) != DsCharClass::PunctSym {
return p;
}
p += l;
}
}
#[inline(always)]
fn ws_token_end(bytes: &[u8], start: usize) -> usize {
let len = bytes.len();
let mut p = start;
let mut last_nl_end = 0usize; let mut last_char_start = start;
let mut at_boundary = false;
while p < len {
let b = unsafe { *bytes.get_unchecked(p) };
if b == b'\r' || b == b'\n' {
last_char_start = p;
p += 1;
last_nl_end = p;
} else if is_ascii_ws(b) {
last_char_start = p;
p += 1;
} else if b < 0x80 {
at_boundary = is_digit(b);
break;
} else {
let (cp, l) = unsafe { decode_cp(bytes, p) };
if ds_class_of(cp) == DsCharClass::Whitespace {
last_char_start = p;
p += l;
} else {
at_boundary = ds_class_of(cp) == DsCharClass::Number || is_deepseek_cjk(cp);
break;
}
}
}
if last_nl_end != 0 {
return last_nl_end; }
if p >= len || at_boundary {
return p; }
if last_char_start > start {
return last_char_start; }
p }
#[inline(always)]
fn scan_gap_from(bytes: &[u8], pos: usize, first_len: usize, cjk_region: bool) -> usize {
let len = bytes.len();
let mut p = pos + first_len;
while p < len {
let b = unsafe { *bytes.get_unchecked(p) };
let (cp, l) = if b < 0x80 {
(b as u32, 1)
} else {
unsafe { decode_cp(bytes, p) }
};
if is_deepseek_cjk(cp) != cjk_region || ds_class_of(cp) != DsCharClass::Other {
return p;
}
if lm_end_at(bytes, p + l, cjk_region).is_some() {
return p;
}
p += l;
}
p
}
#[inline(always)]
fn advance_main(bytes: &[u8], pos: usize) -> usize {
let b0 = unsafe { *bytes.get_unchecked(pos) };
if is_letter(b0) {
return scan_lm_from(bytes, pos + 1, false);
}
if b0 == b' ' {
let Some(&b1) = bytes.get(pos + 1) else {
return pos + 1; };
if is_letter(b1) {
return scan_lm_from(bytes, pos + 2, false); }
if b1 < 0x80 {
if is_digit(b1) {
return pos + 1; }
if b1.is_ascii_punctuation() {
let p = scan_ps_from(bytes, pos + 2, false);
return scan_newlines(bytes, p);
}
if is_ascii_ws(b1) {
return ws_token_end(bytes, pos);
}
return pos + 1; }
let (cp, l) = unsafe { decode_cp(bytes, pos + 1) };
if is_deepseek_cjk(cp) {
return pos + 1; }
let p1 = pos + 1 + l;
match ds_class_of(cp) {
DsCharClass::Letter | DsCharClass::Mark => scan_lm_from(bytes, p1, false),
DsCharClass::Number => pos + 1,
DsCharClass::Whitespace => ws_token_end(bytes, pos),
DsCharClass::PunctSym => {
let p = scan_ps_from(bytes, p1, false);
scan_newlines(bytes, p)
}
DsCharClass::Other => pos + 1, }
} else if b0 < 0x80 {
if b0 == b'\r' || b0 == b'\n' {
return ws_token_end(bytes, pos); }
if is_ascii_ws(b0) {
if let Some(e) = lm_end_at(bytes, pos + 1, false) {
return scan_lm_from(bytes, e, false);
}
return ws_token_end(bytes, pos);
}
if b0.is_ascii_punctuation() {
if let Some(&b1) = bytes.get(pos + 1)
&& is_letter(b1)
{
return swar_scan_letters(bytes, pos + 1);
}
let p = scan_ps_from(bytes, pos + 1, false);
return scan_newlines(bytes, p);
}
if let Some(e) = lm_end_at(bytes, pos + 1, false) {
return scan_lm_from(bytes, e, false);
}
scan_gap_from(bytes, pos, 1, false)
} else {
let (cp, l) = unsafe { decode_cp(bytes, pos) };
let p0 = pos + l;
match ds_class_of(cp) {
DsCharClass::Letter | DsCharClass::Mark => scan_lm_from(bytes, p0, false),
DsCharClass::Whitespace => {
if let Some(e) = lm_end_at(bytes, p0, false) {
return scan_lm_from(bytes, e, false);
}
ws_token_end(bytes, pos)
}
DsCharClass::PunctSym => {
let p = scan_ps_from(bytes, p0, false);
scan_newlines(bytes, p)
}
DsCharClass::Number | DsCharClass::Other => {
if let Some(e) = lm_end_at(bytes, p0, false) {
return scan_lm_from(bytes, e, false);
}
scan_gap_from(bytes, pos, l, false)
}
}
}
}
#[inline(always)]
fn advance_cjk(bytes: &[u8], pos: usize) -> usize {
let (cp, l) = unsafe { decode_cp(bytes, pos) };
let p0 = pos + l;
match ds_class_of(cp) {
DsCharClass::Letter | DsCharClass::Mark => scan_lm_from(bytes, p0, true),
DsCharClass::PunctSym => scan_ps_from(bytes, p0, true),
_ => {
if let Some(e) = lm_end_at(bytes, p0, true) {
return scan_lm_from(bytes, e, true);
}
scan_gap_from(bytes, pos, l, true)
}
}
}
#[inline(always)]
fn advance_pos(bytes: &[u8], pos: usize) -> usize {
let b0 = unsafe { *bytes.get_unchecked(pos) };
if b0 < 0x80 {
if is_digit(b0) {
return scan_numbers_max3(bytes, pos + 1, 1); }
return advance_main(bytes, pos);
}
let (cp, l) = unsafe { decode_cp(bytes, pos) };
if is_deepseek_cjk(cp) {
return advance_cjk(bytes, pos);
}
if ds_class_of(cp) == DsCharClass::Number {
return scan_numbers_max3(bytes, pos + l, 1);
}
advance_main(bytes, pos)
}