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::{AnyTokenizer, Backend, SpecialDecode, SpecialMode, SpecialPolicy};
use crate::core::{StreamingDecoder, Tokenize, Tokenizer};
#[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_with_special(&self, tokens: Vec<u32>) -> PyResult<String> {
self.inner
.decode_with(&tokens, SpecialDecode::Render)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_bytes(&self, tokens: Vec<u32>) -> PyResult<Vec<u8>> {
self.inner
.decode_bytes(&tokens)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_lossy(&self, tokens: Vec<u32>) -> String {
self.inner.decode_lossy(&tokens)
}
fn decode_token_bytes(&self, id: u32) -> PyResult<Vec<u8>> {
self.inner
.decode_token_bytes(id)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_token(&self, id: u32) -> PyResult<String> {
self.inner
.decode_token(id)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
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.streaming_decoder())
}
fn streaming_decoder_with_special(&self) -> PyStreamingDecoder {
PyStreamingDecoder::new(self.inner.streaming_decoder_with(SpecialDecode::Render))
}
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_with_special(&self, ids: Vec<u32>) -> PyResult<String> {
self.inner
.decode_with(&ids, SpecialDecode::Render)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_lossy(&self, ids: Vec<u32>) -> String {
self.inner.decode_lossy(&ids)
}
fn decode_token_bytes(&self, id: u32) -> PyResult<Vec<u8>> {
self.inner
.decode_token_bytes(id)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_token(&self, id: u32) -> PyResult<String> {
self.inner
.decode_token(id)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
#[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 streaming_decoder(&self) -> PyStreamingDecoder {
PyStreamingDecoder::new(self.inner.streaming_decoder())
}
fn streaming_decoder_with_special(&self) -> PyStreamingDecoder {
PyStreamingDecoder::new(self.inner.streaming_decoder_with(SpecialDecode::Render))
}
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()))
}
fn decode_with_special(&self, ids: Vec<u32>) -> PyResult<String> {
Tokenize::decode_with(&self.inner, &ids, SpecialDecode::Render)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_token_bytes(&self, id: u32) -> PyResult<Vec<u8>> {
Tokenize::decode_token_bytes(&self.inner, id)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_token(&self, id: u32) -> PyResult<String> {
Tokenize::decode_token(&self.inner, id).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 streaming_decoder(&self) -> PyStreamingDecoder {
PyStreamingDecoder::new(self.inner.streaming_decoder())
}
fn streaming_decoder_with_special(&self) -> PyStreamingDecoder {
PyStreamingDecoder::new(self.inner.streaming_decoder_with(SpecialDecode::Render))
}
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 decode_with_special(&self, ids: Vec<u32>) -> PyResult<String> {
Tokenize::decode_with(&self.inner, &ids, SpecialDecode::Render)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_token_bytes(&self, id: u32) -> PyResult<Vec<u8>> {
Tokenize::decode_token_bytes(&self.inner, id)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_token(&self, id: u32) -> PyResult<String> {
Tokenize::decode_token(&self.inner, id).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()
}
fn streaming_decoder(&self) -> PyStreamingDecoder {
PyStreamingDecoder::new(self.inner.streaming_decoder())
}
fn streaming_decoder_with_special(&self) -> PyStreamingDecoder {
PyStreamingDecoder::new(self.inner.streaming_decoder_with(SpecialDecode::Render))
}
}
#[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_with_special(&self, ids: Vec<u32>) -> PyResult<String> {
Tokenize::decode_with(&self.inner, &ids, SpecialDecode::Render)
.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>> {
self.bpe_raw()?
.decode_bytes(&tokens)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
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))
}
fn decode_token_bytes(&self, id: u32) -> PyResult<Vec<u8>> {
self.inner
.decode_token_bytes(id)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn decode_token(&self, id: u32) -> PyResult<String> {
self.inner
.decode_token(id)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
#[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> {
self.inner
.streaming_decoder()
.map(PyStreamingDecoder::new)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
fn streaming_decoder_with_special(&self) -> PyResult<PyStreamingDecoder> {
self.inner
.streaming_decoder_with(SpecialDecode::Render)
.map(PyStreamingDecoder::new)
.map_err(|e| PyValueError::new_err(e.to_string()))
}
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 {
inner: StreamingDecoder,
}
#[pymethods]
impl PyStreamingDecoder {
fn add_token(&mut self, token_id: u32) -> Option<String> {
self.inner.add_token_lossy(token_id)
}
fn add_tokens(&mut self, token_ids: Vec<u32>) -> Option<String> {
self.inner.add_tokens_lossy(&token_ids)
}
fn flush(&mut self) -> String {
self.inner.flush()
}
fn reset(&mut self) {
self.inner.reset();
}
#[getter]
fn has_pending(&self) -> bool {
self.inner.has_pending()
}
#[getter]
fn pending_bytes(&self) -> usize {
self.inner.pending_bytes()
}
fn __repr__(&self) -> String {
format!(
"StreamingDecoder(pending_bytes={})",
self.inner.pending_bytes()
)
}
}
impl PyStreamingDecoder {
fn new(inner: StreamingDecoder) -> Self {
Self { inner }
}
}
include!("agent_tokens_generated.rs");