use pyo3::exceptions::{PyIOError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::PyDict;
#[cfg(feature = "rayon")]
use rayon::prelude::*;
use rustc_hash::{FxHashMap, FxHashSet};
use crate::core::SentencePieceTokenizer;
use crate::core::hf_json::{
from_json_bytes as core_from_json_bytes, from_json_path as core_from_json_path,
};
use crate::core::spm::SpmTokenizer;
use crate::core::wordpiece::WordPieceTokenizer;
use crate::core::{byte_level_decode_bytes, Tokenize, Tokenizer};
use crate::core::{AnyTokenizer, Backend, SpecialMode, SpecialPolicy};
#[pyclass(name = "Tokenizer")]
pub struct PyTokenizer {
inner: Tokenizer,
policy: SpecialPolicy,
}
#[pymethods]
impl PyTokenizer {
#[new]
#[pyo3(signature = (vocab_path, pattern, special_tokens=None))]
fn new(
vocab_path: &str,
pattern: &str,
special_tokens: Option<&Bound<'_, PyDict>>,
) -> PyResult<Self> {
let special = parse_special_tokens(special_tokens)?;
let inner = Tokenizer::from_file(vocab_path, pattern, special)
.map_err(|e| PyIOError::new_err(e.to_string()))?;
Ok(Self {
inner,
policy: SpecialPolicy::default(),
})
}
#[staticmethod]
fn from_pretrained(py: Python<'_>, name: &str) -> PyResult<Py<PyAny>> {
let loaded = crate::core::pretrained::from_pretrained(name)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
any_tokenizer_to_py(py, loaded)
}
#[staticmethod]
#[pyo3(signature = (vocab_data, pattern, special_tokens=None))]
fn from_bytes(
vocab_data: &[u8],
pattern: &str,
special_tokens: Option<&Bound<'_, PyDict>>,
) -> PyResult<Self> {
let special = parse_special_tokens(special_tokens)?;
let inner = Tokenizer::from_bytes(vocab_data, pattern, special)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
Ok(Self {
inner,
policy: SpecialPolicy::default(),
})
}
#[pyo3(signature = (use_pcre2=true))]
fn pcre2(&self, use_pcre2: bool) -> PyResult<Self> {
let new_inner = self.inner.clone();
let result = new_inner
.pcre2(use_pcre2)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
Ok(Self {
inner: result,
policy: self.policy.clone(),
})
}
#[pyo3(signature = (use_jit=true))]
fn jit(&self, use_jit: bool) -> PyResult<Self> {
let new_inner = self.inner.clone();
let result = new_inner
.jit(use_jit)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
Ok(Self {
inner: result,
policy: self.policy.clone(),
})
}
fn encode(&self, text: &str) -> Vec<u32> {
self.policy.apply_single(self.inner.encode(text))
}
fn encode_raw(&self, text: &str) -> Vec<u32> {
self.inner.encode(text)
}
fn encode_rayon(&self, text: &str) -> Vec<u32> {
self.policy.apply_single(self.inner.encode_rayon(text))
}
fn encode_with_special(&self, text: &str) -> Vec<u32> {
self.policy
.apply_single(self.inner.encode_with_special(text))
}
fn encode_ordinary(&self, text: &str) -> Vec<u32> {
self.policy.apply_single(self.inner.encode_ordinary(text))
}
fn encode_allowed_special(
&self,
text: &str,
allowed_special: Vec<String>,
) -> PyResult<Vec<u32>> {
let allowed: FxHashSet<String> = allowed_special.into_iter().collect();
let ids = self
.inner
.encode_with(text, &SpecialMode::Allow(&allowed))
.map_err(|e| PyValueError::new_err(e.to_string()))?;
Ok(self.policy.apply_single(ids))
}
fn decode(&self, tokens: Vec<u32>) -> PyResult<String> {
self.inner
.decode(&tokens)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_bytes(&self, tokens: Vec<u32>) -> Vec<u8> {
self.inner.decode_bytes(&tokens)
}
fn decode_lossy(&self, tokens: Vec<u32>) -> String {
self.inner.decode_lossy(&tokens)
}
fn encode_batch(&self, texts: Vec<String>) -> Vec<Vec<u32>> {
self.inner
.encode_batch(&texts)
.into_iter()
.map(|ids| self.policy.apply_single(ids))
.collect()
}
fn encode_batch_with_special(&self, texts: Vec<String>) -> Vec<Vec<u32>> {
self.inner
.encode_batch_with_special(&texts)
.into_iter()
.map(|ids| self.policy.apply_single(ids))
.collect()
}
fn decode_batch(&self, token_lists: Vec<Vec<u32>>) -> PyResult<Vec<String>> {
self.inner
.decode_batch(&token_lists)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_batch_lossy(&self, token_lists: Vec<Vec<u32>>) -> Vec<String> {
self.inner.decode_batch_lossy(&token_lists)
}
#[getter]
fn vocab_size(&self) -> usize {
self.inner.vocab_size()
}
fn streaming_decoder(&self) -> PyStreamingDecoder {
PyStreamingDecoder::new(
self.inner.decoder().clone(),
self.inner.special_tokens_decoder().clone(),
)
}
fn byte_level_streaming_decoder(&self) -> PyByteLevelStreamingDecoder {
PyByteLevelStreamingDecoder::new(
self.inner.decoder().clone(),
self.inner.special_tokens_decoder().clone(),
)
}
fn clear_cache(&self) {
self.inner.clear_cache();
}
#[getter]
fn cache_len(&self) -> usize {
self.inner.cache_len()
}
fn __repr__(&self) -> String {
format!("Tokenizer(vocab_size={})", self.inner.vocab_size())
}
}
#[pyclass(name = "SentencePieceTokenizer")]
pub struct PySentencePieceTokenizer {
inner: SentencePieceTokenizer,
policy: SpecialPolicy,
bos_token_id: Option<u32>,
}
#[pymethods]
impl PySentencePieceTokenizer {
#[new]
#[pyo3(signature = (tokens, scores, eos_token_id, bos_token_id=None))]
fn new(
tokens: Vec<String>,
scores: Vec<f64>,
eos_token_id: u32,
bos_token_id: Option<u32>,
) -> PyResult<Self> {
let inner = SentencePieceTokenizer::new(tokens, scores, bos_token_id, eos_token_id)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
Ok(Self {
inner,
policy: SpecialPolicy::default(),
bos_token_id,
})
}
fn encode(&self, text: &str) -> Vec<u32> {
self.policy
.apply_single(self.with_bos(self.inner.encode(text)))
}
fn encode_raw(&self, text: &str) -> Vec<u32> {
self.inner.encode(text)
}
fn encode_with_special(&self, text: &str) -> Vec<u32> {
self.encode(text)
}
fn encode_ordinary(&self, text: &str) -> Vec<u32> {
self.policy
.apply_single(self.with_bos(self.inner.encode_ordinary(text)))
}
fn encode_allowed_special(
&self,
text: &str,
allowed_special: Vec<String>,
) -> PyResult<Vec<u32>> {
let allowed: FxHashSet<String> = allowed_special.into_iter().collect();
let ids = self
.inner
.encode_with(text, &SpecialMode::Allow(&allowed))
.map_err(|e| PyValueError::new_err(e.to_string()))?;
Ok(self.policy.apply_single(self.with_bos(ids)))
}
fn encode_batch(&self, texts: Vec<String>) -> Vec<Vec<u32>> {
#[cfg(feature = "rayon")]
{
texts.par_iter().map(|text| self.encode(text)).collect()
}
#[cfg(not(feature = "rayon"))]
{
texts.iter().map(|text| self.encode(text)).collect()
}
}
fn decode(&self, ids: Vec<u32>) -> PyResult<String> {
self.inner
.decode(&ids)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_lossy(&self, ids: Vec<u32>) -> String {
self.inner.decode_lossy(&ids)
}
#[getter]
fn vocab_size(&self) -> usize {
self.inner.vocab_size()
}
fn is_eos(&self, token_id: u32) -> bool {
self.inner.is_eos(token_id)
}
#[getter]
fn eos_token_id(&self) -> u32 {
self.inner.eos_token_id()
}
#[getter]
fn bos_token_id(&self) -> Option<u32> {
self.inner.bos_token_id()
}
fn __repr__(&self) -> String {
format!(
"SentencePieceTokenizer(vocab_size={})",
self.inner.vocab_size()
)
}
}
impl PySentencePieceTokenizer {
fn with_bos(&self, mut ids: Vec<u32>) -> Vec<u32> {
if let Some(bos) = self.bos_token_id {
ids.insert(0, bos);
}
ids
}
}
#[pyclass(name = "SpmTokenizer")]
pub struct PySpmTokenizer {
inner: SpmTokenizer,
policy: SpecialPolicy,
}
#[pymethods]
impl PySpmTokenizer {
#[new]
#[pyo3(signature = (tokens, scores, bos_token_id=None, eos_token_id=None))]
fn new(
tokens: Vec<String>,
scores: Vec<f32>,
bos_token_id: Option<u32>,
eos_token_id: Option<u32>,
) -> PyResult<Self> {
let inner = SpmTokenizer::new(tokens, scores, bos_token_id, eos_token_id)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
Ok(Self {
inner,
policy: SpecialPolicy::default(),
})
}
fn encode(&self, text: &str) -> Vec<u32> {
self.policy
.apply_single(Tokenize::encode(&self.inner, text))
}
fn encode_raw(&self, text: &str) -> Vec<u32> {
Tokenize::encode(&self.inner, text)
}
fn encode_with_special(&self, text: &str) -> Vec<u32> {
self.encode(text)
}
fn encode_batch(&self, texts: Vec<String>) -> Vec<Vec<u32>> {
#[cfg(feature = "rayon")]
{
texts.par_iter().map(|text| self.encode(text)).collect()
}
#[cfg(not(feature = "rayon"))]
{
texts.iter().map(|text| self.encode(text)).collect()
}
}
fn encode_ordinary(&self, text: &str) -> Vec<u32> {
self.policy.apply_single(self.inner.encode_ordinary(text))
}
fn encode_allowed_special(
&self,
text: &str,
allowed_special: Vec<String>,
) -> PyResult<Vec<u32>> {
let allowed: FxHashSet<String> = allowed_special.into_iter().collect();
let ids = self
.inner
.encode_with(text, &SpecialMode::Allow(&allowed))
.map_err(|e| PyValueError::new_err(e.to_string()))?;
Ok(self.policy.apply_single(ids))
}
fn decode(&self, ids: Vec<u32>) -> PyResult<String> {
Tokenize::decode(&self.inner, &ids).map_err(|e| PyValueError::new_err(e.to_string()))
}
#[getter]
fn vocab_size(&self) -> usize {
Tokenize::vocab_size(&self.inner)
}
#[getter]
fn eos_token_id(&self) -> Option<u32> {
self.inner.eos_token_id()
}
#[getter]
fn bos_token_id(&self) -> Option<u32> {
self.inner.bos_token_id()
}
fn __repr__(&self) -> String {
format!(
"SpmTokenizer(vocab_size={})",
Tokenize::vocab_size(&self.inner)
)
}
}
#[pyclass(name = "WordPieceTokenizer")]
pub struct PyWordPieceTokenizer {
inner: WordPieceTokenizer,
policy: SpecialPolicy,
}
#[pymethods]
impl PyWordPieceTokenizer {
#[new]
#[pyo3(signature = (vocab, unk_token_id, max_word_len=100, do_lower_case=false, strip_accents=None))]
fn new(
vocab: Vec<String>,
unk_token_id: u32,
max_word_len: usize,
do_lower_case: bool,
strip_accents: Option<bool>,
) -> Self {
let inner = WordPieceTokenizer::new(vocab, unk_token_id, max_word_len, do_lower_case);
Self {
inner: match strip_accents {
Some(strip) => inner.with_strip_accents(strip),
None => inner,
},
policy: SpecialPolicy::default(),
}
}
fn encode(&self, text: &str) -> Vec<u32> {
self.policy
.apply_single(Tokenize::encode(&self.inner, text))
}
fn encode_raw(&self, text: &str) -> Vec<u32> {
Tokenize::encode(&self.inner, text)
}
fn encode_with_special(&self, text: &str) -> Vec<u32> {
self.encode(text)
}
fn encode_ordinary(&self, text: &str) -> Vec<u32> {
self.policy.apply_single(self.inner.encode_ordinary(text))
}
fn encode_allowed_special(
&self,
text: &str,
allowed_special: Vec<String>,
) -> PyResult<Vec<u32>> {
let allowed: FxHashSet<String> = allowed_special.into_iter().collect();
let ids = self
.inner
.encode_with(text, &SpecialMode::Allow(&allowed))
.map_err(|e| PyValueError::new_err(e.to_string()))?;
Ok(self.policy.apply_single(ids))
}
fn encode_batch(&self, texts: Vec<String>) -> Vec<Vec<u32>> {
#[cfg(feature = "rayon")]
{
texts.par_iter().map(|text| self.encode(text)).collect()
}
#[cfg(not(feature = "rayon"))]
{
texts.iter().map(|text| self.encode(text)).collect()
}
}
fn decode(&self, ids: Vec<u32>) -> PyResult<String> {
Tokenize::decode(&self.inner, &ids).map_err(|e| PyValueError::new_err(e.to_string()))
}
fn vocab_size(&self) -> usize {
Tokenize::vocab_size(&self.inner)
}
#[getter]
fn unk_token_id(&self) -> u32 {
self.inner.unk_token_id()
}
#[getter]
fn cls_token_id(&self) -> Option<u32> {
self.inner.cls_token_id()
}
#[getter]
fn sep_token_id(&self) -> Option<u32> {
self.inner.sep_token_id()
}
#[getter]
fn pad_token_id(&self) -> Option<u32> {
self.inner.pad_token_id()
}
}
#[pyclass(name = "AnyTokenizer")]
pub struct PyAnyTokenizer {
inner: AnyTokenizer,
}
#[pymethods]
impl PyAnyTokenizer {
fn encode(&self, text: &str) -> Vec<u32> {
self.inner.encode(text)
}
fn encode_raw(&self, text: &str) -> Vec<u32> {
self.inner.encode_raw(text)
}
fn encode_with_special(&self, text: &str) -> PyResult<Vec<u32>> {
self.inner
.encode_with(text, &SpecialMode::All)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn encode_ordinary(&self, text: &str) -> PyResult<Vec<u32>> {
self.inner
.encode_with(text, &SpecialMode::Ordinary)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn encode_allowed_special(
&self,
text: &str,
allowed_special: Vec<String>,
) -> PyResult<Vec<u32>> {
let allowed: FxHashSet<String> = allowed_special.into_iter().collect();
self.inner
.encode_with(text, &SpecialMode::Allow(&allowed))
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn encode_batch(&self, texts: Vec<String>) -> Vec<Vec<u32>> {
let refs: Vec<&str> = texts.iter().map(String::as_str).collect();
self.inner.encode_batch(&refs)
}
fn encode_batch_with_special(&self, texts: Vec<String>) -> PyResult<Vec<Vec<u32>>> {
let refs: Vec<&str> = texts.iter().map(String::as_str).collect();
self.inner
.encode_batch_with(&refs, &SpecialMode::All)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn encode_rayon(&self, text: &str) -> Vec<u32> {
self.inner.encode_rayon(text)
}
fn decode(&self, ids: Vec<u32>) -> PyResult<String> {
Tokenize::decode(&self.inner, &ids).map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_batch(&self, token_lists: Vec<Vec<u32>>) -> PyResult<Vec<String>> {
self.inner
.decode_batch(&token_lists)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_bytes(&self, tokens: Vec<u32>) -> PyResult<Vec<u8>> {
Ok(self.bpe_raw()?.decode_bytes(&tokens))
}
fn decode_lossy(&self, tokens: Vec<u32>) -> PyResult<String> {
Ok(self.bpe_raw()?.decode_lossy(&tokens))
}
fn decode_batch_lossy(&self, token_lists: Vec<Vec<u32>>) -> PyResult<Vec<String>> {
Ok(self.bpe_raw()?.decode_batch_lossy(&token_lists))
}
#[pyo3(signature = (use_pcre2=true))]
fn pcre2<'py>(mut slf: PyRefMut<'py, Self>, use_pcre2: bool) -> PyResult<PyRefMut<'py, Self>> {
slf.inner
.set_pcre2(use_pcre2)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
Ok(slf)
}
#[pyo3(signature = (use_jit=true))]
fn jit<'py>(mut slf: PyRefMut<'py, Self>, use_jit: bool) -> PyResult<PyRefMut<'py, Self>> {
slf.inner
.set_jit(use_jit)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
Ok(slf)
}
fn streaming_decoder(&self) -> PyResult<PyStreamingDecoder> {
let bpe = self.bpe_raw()?;
Ok(PyStreamingDecoder::new(
bpe.decoder().clone(),
bpe.special_tokens_decoder().clone(),
))
}
fn byte_level_streaming_decoder(&self) -> PyResult<PyByteLevelStreamingDecoder> {
let bpe = self.bpe_raw()?;
Ok(PyByteLevelStreamingDecoder::new(
bpe.decoder().clone(),
bpe.special_tokens_decoder().clone(),
))
}
fn clear_cache(&self) -> PyResult<()> {
self.bpe()?.clear_cache();
Ok(())
}
#[getter]
fn cache_len(&self) -> PyResult<usize> {
Ok(self.bpe()?.cache_len())
}
#[getter]
fn vocab_size(&self) -> usize {
Tokenize::vocab_size(&self.inner)
}
fn is_eos(&self, token_id: u32) -> bool {
self.inner.is_eos(token_id)
}
#[getter]
fn eos_token_id(&self) -> Option<u32> {
self.inner.eos_token_id()
}
fn special_token_id(&self, name: &str) -> Option<u32> {
self.inner.special_token_id(name)
}
#[getter]
fn family(&self) -> &'static str {
self.inner.family()
}
fn __repr__(&self) -> String {
format!(
"AnyTokenizer(family={}, vocab_size={})",
self.inner.family(),
Tokenize::vocab_size(&self.inner)
)
}
}
impl PyAnyTokenizer {
fn bpe(&self) -> PyResult<&Tokenizer> {
match self.inner.backend() {
Backend::Bpe(bpe) => Ok(bpe),
_ => Err(PyValueError::new_err(format!(
"this method needs the byte-level BPE backend; this tokenizer is {}",
self.inner.family()
))),
}
}
fn bpe_raw(&self) -> PyResult<&Tokenizer> {
if self.inner.declares_decoder() {
return Err(PyValueError::new_err(
"this tokenizer declares a `decoder` pipeline, which raw byte-level \
decoding would bypass; use `decode` / `decode_batch` instead",
));
}
self.bpe()
}
}
fn any_tokenizer_to_py(py: Python<'_>, any: AnyTokenizer) -> PyResult<Py<PyAny>> {
Ok(Py::new(py, PyAnyTokenizer { inner: any })?.into_any())
}
#[pyfunction]
pub fn from_json(py: Python<'_>, path: &str) -> PyResult<Py<PyAny>> {
let any = core_from_json_path(path).map_err(|e| PyValueError::new_err(e.to_string()))?;
any_tokenizer_to_py(py, any)
}
#[pyfunction]
pub fn from_json_bytes(py: Python<'_>, data: &[u8]) -> PyResult<Py<PyAny>> {
let any = core_from_json_bytes(data).map_err(|e| PyValueError::new_err(e.to_string()))?;
any_tokenizer_to_py(py, any)
}
#[pyfunction]
pub fn base_vocab_size(name: &str) -> PyResult<u32> {
crate::core::pretrained::base_vocab_size_by_name(name)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn parse_special_tokens(
special_tokens: Option<&Bound<'_, PyDict>>,
) -> PyResult<FxHashMap<String, u32>> {
let mut result = FxHashMap::default();
if let Some(dict) = special_tokens {
for (key, value) in dict.iter() {
let k: String = key.extract()?;
let v: u32 = value.extract()?;
result.insert(k, v);
}
}
Ok(result)
}
#[pyclass(name = "StreamingDecoder")]
pub struct PyStreamingDecoder {
decoder: FxHashMap<u32, Vec<u8>>,
special_decoder: FxHashMap<u32, String>,
buffer: Vec<u8>,
}
#[pymethods]
impl PyStreamingDecoder {
fn add_token(&mut self, token_id: u32) -> Option<String> {
let bytes = match self.decoder.get(&token_id) {
Some(b) => b.as_slice(),
None => self.special_decoder.get(&token_id)?.as_bytes(),
};
self.buffer.extend_from_slice(bytes);
self.extract_complete_utf8()
}
fn add_tokens(&mut self, token_ids: Vec<u32>) -> Option<String> {
for token_id in token_ids {
let bytes = if let Some(b) = self.decoder.get(&token_id) {
b.as_slice()
} else if let Some(s) = self.special_decoder.get(&token_id) {
s.as_bytes()
} else {
continue;
};
self.buffer.extend_from_slice(bytes);
}
self.extract_complete_utf8()
}
fn flush(&mut self) -> String {
if self.buffer.is_empty() {
return String::new();
}
let result = String::from_utf8_lossy(&self.buffer).into_owned();
self.buffer.clear();
result
}
fn reset(&mut self) {
self.buffer.clear();
}
#[getter]
fn has_pending(&self) -> bool {
!self.buffer.is_empty()
}
#[getter]
fn pending_bytes(&self) -> usize {
self.buffer.len()
}
fn __repr__(&self) -> String {
format!("StreamingDecoder(pending_bytes={})", self.buffer.len())
}
}
impl PyStreamingDecoder {
fn new(decoder: FxHashMap<u32, Vec<u8>>, special_decoder: FxHashMap<u32, String>) -> Self {
Self {
decoder,
special_decoder,
buffer: Vec::with_capacity(16),
}
}
fn extract_complete_utf8(&mut self) -> Option<String> {
if self.buffer.is_empty() {
return None;
}
let valid_len = self.find_valid_utf8_len();
if valid_len == 0 {
return None;
}
let valid_bytes: Vec<u8> = self.buffer.drain(..valid_len).collect();
let result = unsafe { String::from_utf8_unchecked(valid_bytes) };
Some(result)
}
fn find_valid_utf8_len(&self) -> usize {
let bytes = &self.buffer;
let len = bytes.len();
if len == 0 {
return 0;
}
if std::str::from_utf8(bytes).is_ok() {
return len;
}
for incomplete_len in 1..=3.min(len) {
let check_len = len - incomplete_len;
if check_len == 0 {
continue;
}
if std::str::from_utf8(&bytes[..check_len]).is_ok()
&& Self::could_be_incomplete_sequence(&bytes[check_len..])
{
return check_len;
}
}
for i in (0..len).rev() {
if std::str::from_utf8(&bytes[..=i]).is_ok() {
return i + 1;
}
}
0
}
fn could_be_incomplete_sequence(bytes: &[u8]) -> bool {
if bytes.is_empty() {
return false;
}
let first = bytes[0];
match first {
0xC0..=0xDF => bytes.len() < 2,
0xE0..=0xEF => bytes.len() < 3,
0xF0..=0xF7 => bytes.len() < 4,
_ => false,
}
}
}
#[pyclass(name = "ByteLevelStreamingDecoder")]
pub struct PyByteLevelStreamingDecoder {
decoder: FxHashMap<u32, Vec<u8>>,
special_decoder: FxHashMap<u32, String>,
buffer: Vec<u8>,
}
#[pymethods]
impl PyByteLevelStreamingDecoder {
fn add_token(&mut self, token_id: u32) -> Option<String> {
match self.decoder.get(&token_id) {
Some(encoded_bytes) => {
if let Some(raw_bytes) = byte_level_decode_bytes(encoded_bytes) {
self.buffer.extend_from_slice(&raw_bytes);
} else {
self.buffer.extend_from_slice(encoded_bytes);
}
}
None => self
.buffer
.extend_from_slice(self.special_decoder.get(&token_id)?.as_bytes()),
}
self.extract_complete_utf8()
}
fn add_tokens(&mut self, token_ids: Vec<u32>) -> Option<String> {
for token_id in token_ids {
if let Some(encoded_bytes) = self.decoder.get(&token_id) {
if let Some(raw_bytes) = byte_level_decode_bytes(encoded_bytes) {
self.buffer.extend_from_slice(&raw_bytes);
} else {
self.buffer.extend_from_slice(encoded_bytes);
}
} else if let Some(special) = self.special_decoder.get(&token_id) {
self.buffer.extend_from_slice(special.as_bytes());
}
}
self.extract_complete_utf8()
}
fn flush(&mut self) -> String {
if self.buffer.is_empty() {
return String::new();
}
let result = String::from_utf8_lossy(&self.buffer).into_owned();
self.buffer.clear();
result
}
fn reset(&mut self) {
self.buffer.clear();
}
#[getter]
fn has_pending(&self) -> bool {
!self.buffer.is_empty()
}
#[getter]
fn pending_bytes(&self) -> usize {
self.buffer.len()
}
fn __repr__(&self) -> String {
format!(
"ByteLevelStreamingDecoder(pending_bytes={})",
self.buffer.len()
)
}
}
impl PyByteLevelStreamingDecoder {
fn new(decoder: FxHashMap<u32, Vec<u8>>, special_decoder: FxHashMap<u32, String>) -> Self {
Self {
decoder,
special_decoder,
buffer: Vec::with_capacity(16),
}
}
fn extract_complete_utf8(&mut self) -> Option<String> {
if self.buffer.is_empty() {
return None;
}
let valid_len = self.find_valid_utf8_len();
if valid_len == 0 {
return None;
}
let valid_bytes: Vec<u8> = self.buffer.drain(..valid_len).collect();
let result = unsafe { String::from_utf8_unchecked(valid_bytes) };
Some(result)
}
fn find_valid_utf8_len(&self) -> usize {
let bytes = &self.buffer;
let len = bytes.len();
if len == 0 {
return 0;
}
if std::str::from_utf8(bytes).is_ok() {
return len;
}
for incomplete_len in 1..=3.min(len) {
let check_len = len - incomplete_len;
if check_len == 0 {
continue;
}
if std::str::from_utf8(&bytes[..check_len]).is_ok()
&& Self::could_be_incomplete_sequence(&bytes[check_len..])
{
return check_len;
}
}
for i in (0..len).rev() {
if std::str::from_utf8(&bytes[..=i]).is_ok() {
return i + 1;
}
}
0
}
fn could_be_incomplete_sequence(bytes: &[u8]) -> bool {
if bytes.is_empty() {
return false;
}
let first = bytes[0];
match first {
0xC0..=0xDF => bytes.len() < 2,
0xE0..=0xEF => bytes.len() < 3,
0xF0..=0xF7 => bytes.len() < 4,
_ => false,
}
}
}
include!("agent_tokens_generated.rs");