use fancy_regex::Regex;
use serde::Deserialize;
use crate::pre_tokenized::{PreTokenizedString, Split as PtSplit};
use crate::pre_tokenizer::Error;
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
}
const BYTE_TO_UTF8: [[u8; 2]; 256] = build_byte_to_utf8();
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
}
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) };
}
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);
}
}
unsafe fn encode_bytes_into_unchecked(s: &str, out: &mut String) {
unsafe { encode_bytes_bulk(s.as_bytes(), out.as_mut_vec()) };
}
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
}
#[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,
}
#[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 {
pub fn is_bulk_only(&self) -> bool {
self.regex.is_none() && !self.add_prefix_space
}
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(())
}
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(())
}
}