From 894eb7dc1f1aba507a898814bf79c05a23881ea2 Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Wed, 22 Jul 2026 17:24:30 +0200 Subject: [PATCH 01/17] implement ultra-basic, erroneous decode --- tokenizers/tk-encode/src/models/bpe/model.rs | 4 ++ .../tk-encode/src/models/unigram/model.rs | 4 ++ .../tk-encode/src/models/wordlevel/mod.rs | 4 ++ .../tk-encode/src/models/wordpiece/mod.rs | 53 +++++++++++++++++++ .../tk-encode/src/tokenizer/pipeline.rs | 23 +++++++- 5 files changed, 86 insertions(+), 2 deletions(-) diff --git a/tokenizers/tk-encode/src/models/bpe/model.rs b/tokenizers/tk-encode/src/models/bpe/model.rs index fa7bcec23..94fe27f18 100644 --- a/tokenizers/tk-encode/src/models/bpe/model.rs +++ b/tokenizers/tk-encode/src/models/bpe/model.rs @@ -855,6 +855,10 @@ impl pipeline::Model for PipelineBPE { skip: Vec::new(), } } + + fn id_to_token_bytes(&self, id: &PipelineToken) -> Option<&[u8]> { + self.vocab.id_to_token_bytes(id.id) + } } pub struct BpeScratch { diff --git a/tokenizers/tk-encode/src/models/unigram/model.rs b/tokenizers/tk-encode/src/models/unigram/model.rs index 1cc3f4b98..a4b3c1410 100644 --- a/tokenizers/tk-encode/src/models/unigram/model.rs +++ b/tokenizers/tk-encode/src/models/unigram/model.rs @@ -549,6 +549,10 @@ impl pipeline::Model for Unigram { } Ok(()) } + + fn id_to_token_bytes(&self, id: &PipelineToken) -> Option<&[u8]> { + self.token_to_ids.id_to_token_bytes(id.id) + } } #[cfg(test)] diff --git a/tokenizers/tk-encode/src/models/wordlevel/mod.rs b/tokenizers/tk-encode/src/models/wordlevel/mod.rs index e26f16387..081760e67 100644 --- a/tokenizers/tk-encode/src/models/wordlevel/mod.rs +++ b/tokenizers/tk-encode/src/models/wordlevel/mod.rs @@ -228,6 +228,10 @@ impl pipeline::Model for WordLevel { } Ok(()) } + + fn id_to_token_bytes(&self, id: &PipelineToken) -> Option<&[u8]> { + self.vocab_r.get(&id.id).map(|s| s.as_bytes()) + } } #[cfg(test)] diff --git a/tokenizers/tk-encode/src/models/wordpiece/mod.rs b/tokenizers/tk-encode/src/models/wordpiece/mod.rs index a1286fe95..d9c86cde6 100644 --- a/tokenizers/tk-encode/src/models/wordpiece/mod.rs +++ b/tokenizers/tk-encode/src/models/wordpiece/mod.rs @@ -322,6 +322,11 @@ impl pipeline::ModelScratch for WordPieceScratch {} pub struct PipelineWordPiece { vocab_trie: yada::DoubleArray>, + // Token bytes for decoding. The trie can't hand its keys back, so we keep + // a second copy; ids are dense line indices, so a flat arena indexed by id + // is cheaper than a hashmap. Token `id` is `vocab_r[offs[id]..offs[id + 1]]`. + vocab_r: Box<[u8]>, + vocab_r_offsets: Box<[u32]>, unk_token: Option, continuing_subword_prefix: String, max_input_chars_per_word: usize, @@ -344,11 +349,31 @@ impl TryFrom for PipelineWordPiece { keyset.sort_unstable_by(|(a, _), (b, _)| a.cmp(b)); let vocab_trie = DoubleArray::new(DoubleArrayBuilder::build(&keyset)?)?; + let vocab_size = keyset + .iter() + .map(|(_, id)| *id as usize + 1) + .max() + .unwrap_or(0); + let mut vocab_r_offsets = vec![0u32; vocab_size + 1]; + for (token, id) in &keyset { + vocab_r_offsets[*id as usize + 1] = token.len() as u32; + } + for i in 1..vocab_r_offsets.len() { + vocab_r_offsets[i] += vocab_r_offsets[i - 1]; + } + let mut vocab_r = vec![0u8; vocab_r_offsets[vocab_size] as usize]; + for (token, id) in &keyset { + let start = vocab_r_offsets[*id as usize] as usize; + vocab_r[start..start + token.len()].copy_from_slice(token.as_bytes()); + } + Ok(Self { continuing_subword_prefix, max_input_chars_per_word, unk_token, vocab_trie, + vocab_r: vocab_r.into_boxed_slice(), + vocab_r_offsets: vocab_r_offsets.into_boxed_slice(), }) } } @@ -409,6 +434,13 @@ impl pipeline::Model for PipelineWordPiece { } Ok(()) } + + fn id_to_token_bytes(&self, id: &PipelineToken) -> Option<&[u8]> { + let i = id.id as usize; + let start = *self.vocab_r_offsets.get(i)? as usize; + let end = *self.vocab_r_offsets.get(i + 1)? as usize; + Some(&self.vocab_r[start..end]) + } } #[cfg(test)] @@ -419,4 +451,25 @@ mod tests { fn test_error_display() { assert!(format!("{}", Error::MissingUnkToken).contains("Missing [UNK] token")); } + + #[test] + fn id_to_token_bytes_round_trips() { + use crate::pipeline::{Model as _, PipelineToken}; + + let vocab: Vocab = [ + ("[UNK]".to_string(), 0), + ("hello".to_string(), 1), + ("##world".to_string(), 2), + ] + .into_iter() + .collect(); + let wp = WordPiece::builder().vocab(vocab).build().unwrap(); + let model = PipelineWordPiece::try_from(wp).unwrap(); + + for (token, id) in [("[UNK]", 0), ("hello", 1), ("##world", 2)] { + let bytes = model.id_to_token_bytes(&PipelineToken { id }).unwrap(); + assert_eq!(bytes, token.as_bytes()); + } + assert_eq!(model.id_to_token_bytes(&PipelineToken { id: 3 }), None); + } } diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index efb12eb67..ef8c4a748 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -542,8 +542,16 @@ impl PipelineTokenizer { /// test and the comparative benchmark report decode as *pending* instead of /// silently validating garbage. Implementing this flips the ignored /// `pipeline_decode_oracle` test on and lights up the decode charts. - pub fn decode(&self, _ids: &[u32], _skip_special_tokens: bool) -> Result { - Err("PipelineTokenizer::decode is not implemented yet".into()) + pub fn decode(&self, ids: &[PipelineToken], _skip_special_tokens: bool) -> Result { + let mut output = Vec::with_capacity(ids.len()); + for id in ids { + let slice = self + .model + .id_to_token_bytes(id) + .ok_or::(format!("Invalid token id: {}", id.id).into())?; + output.extend_from_slice(slice); + } + Ok(String::from_utf8(output)?) } /// Decode several id sequences at once, one `String` per input. Mirrors the @@ -923,6 +931,8 @@ pub trait Model { ) -> Result<()>; fn init_scratch(&self) -> Self::Scratch; + + fn id_to_token_bytes(&self, id: &PipelineToken) -> Option<&[u8]>; } #[allow( @@ -970,6 +980,15 @@ impl Model for PipelineModel { Self::Unigram(unigram) => Self::Scratch::Unigram(unigram.init_scratch()), } } + + fn id_to_token_bytes(&self, id: &PipelineToken) -> Option<&[u8]> { + match self { + Self::BPE(model) => model.id_to_token_bytes(id), + Self::WordLevel(model) => model.id_to_token_bytes(id), + Self::WordPiece(model) => model.id_to_token_bytes(id), + Self::Unigram(model) => model.id_to_token_bytes(id), + } + } } pub enum PipelineModelScratch { From c3848b25715137195ec00546ebbf291d77f0b9da Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Wed, 22 Jul 2026 17:37:19 +0200 Subject: [PATCH 02/17] fix: api --- tokenizers/tk-encode/src/models/bpe/model.rs | 4 ++-- .../tk-encode/src/models/unigram/model.rs | 4 ++-- .../tk-encode/src/models/wordlevel/mod.rs | 4 ++-- .../tk-encode/src/models/wordpiece/mod.rs | 10 +++++----- tokenizers/tk-encode/src/tokenizer/pipeline.rs | 18 ++++++++---------- 5 files changed, 19 insertions(+), 21 deletions(-) diff --git a/tokenizers/tk-encode/src/models/bpe/model.rs b/tokenizers/tk-encode/src/models/bpe/model.rs index 94fe27f18..4cc9905c6 100644 --- a/tokenizers/tk-encode/src/models/bpe/model.rs +++ b/tokenizers/tk-encode/src/models/bpe/model.rs @@ -856,8 +856,8 @@ impl pipeline::Model for PipelineBPE { } } - fn id_to_token_bytes(&self, id: &PipelineToken) -> Option<&[u8]> { - self.vocab.id_to_token_bytes(id.id) + fn id_to_token_bytes(&self, id: u32) -> Option<&[u8]> { + self.vocab.id_to_token_bytes(id) } } diff --git a/tokenizers/tk-encode/src/models/unigram/model.rs b/tokenizers/tk-encode/src/models/unigram/model.rs index a4b3c1410..b0d9daa51 100644 --- a/tokenizers/tk-encode/src/models/unigram/model.rs +++ b/tokenizers/tk-encode/src/models/unigram/model.rs @@ -550,8 +550,8 @@ impl pipeline::Model for Unigram { Ok(()) } - fn id_to_token_bytes(&self, id: &PipelineToken) -> Option<&[u8]> { - self.token_to_ids.id_to_token_bytes(id.id) + fn id_to_token_bytes(&self, id: u32) -> Option<&[u8]> { + self.token_to_ids.id_to_token_bytes(id) } } diff --git a/tokenizers/tk-encode/src/models/wordlevel/mod.rs b/tokenizers/tk-encode/src/models/wordlevel/mod.rs index 081760e67..03e73b3b5 100644 --- a/tokenizers/tk-encode/src/models/wordlevel/mod.rs +++ b/tokenizers/tk-encode/src/models/wordlevel/mod.rs @@ -229,8 +229,8 @@ impl pipeline::Model for WordLevel { Ok(()) } - fn id_to_token_bytes(&self, id: &PipelineToken) -> Option<&[u8]> { - self.vocab_r.get(&id.id).map(|s| s.as_bytes()) + fn id_to_token_bytes(&self, id: u32) -> Option<&[u8]> { + self.vocab_r.get(&id).map(|s| s.as_bytes()) } } diff --git a/tokenizers/tk-encode/src/models/wordpiece/mod.rs b/tokenizers/tk-encode/src/models/wordpiece/mod.rs index d9c86cde6..c8aa7be66 100644 --- a/tokenizers/tk-encode/src/models/wordpiece/mod.rs +++ b/tokenizers/tk-encode/src/models/wordpiece/mod.rs @@ -435,8 +435,8 @@ impl pipeline::Model for PipelineWordPiece { Ok(()) } - fn id_to_token_bytes(&self, id: &PipelineToken) -> Option<&[u8]> { - let i = id.id as usize; + fn id_to_token_bytes(&self, id: u32) -> Option<&[u8]> { + let i = id as usize; let start = *self.vocab_r_offsets.get(i)? as usize; let end = *self.vocab_r_offsets.get(i + 1)? as usize; Some(&self.vocab_r[start..end]) @@ -454,7 +454,7 @@ mod tests { #[test] fn id_to_token_bytes_round_trips() { - use crate::pipeline::{Model as _, PipelineToken}; + use crate::pipeline::Model as _; let vocab: Vocab = [ ("[UNK]".to_string(), 0), @@ -467,9 +467,9 @@ mod tests { let model = PipelineWordPiece::try_from(wp).unwrap(); for (token, id) in [("[UNK]", 0), ("hello", 1), ("##world", 2)] { - let bytes = model.id_to_token_bytes(&PipelineToken { id }).unwrap(); + let bytes = model.id_to_token_bytes(id).unwrap(); assert_eq!(bytes, token.as_bytes()); } - assert_eq!(model.id_to_token_bytes(&PipelineToken { id: 3 }), None); + assert_eq!(model.id_to_token_bytes(3), None); } } diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index ef8c4a748..295dd0b15 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -537,18 +537,16 @@ impl PipelineTokenizer { /// Decode token ids back to a `String`. /// - /// Not implemented yet — the pipeline decode path is being built. It fails - /// loud (rather than returning a plausible-but-wrong string) so the oracle - /// test and the comparative benchmark report decode as *pending* instead of - /// silently validating garbage. Implementing this flips the ignored - /// `pipeline_decode_oracle` test on and lights up the decode charts. - pub fn decode(&self, ids: &[PipelineToken], _skip_special_tokens: bool) -> Result { + /// Incomplete: it concatenates raw token bytes only. No decoder, no added- + /// vocab lookup, no `skip_special_tokens` — so the `pipeline_decode_oracle` + /// test fails on purpose until those land. + pub fn decode(&self, ids: &[u32], _skip_special_tokens: bool) -> Result { let mut output = Vec::with_capacity(ids.len()); - for id in ids { + for &id in ids { let slice = self .model .id_to_token_bytes(id) - .ok_or::(format!("Invalid token id: {}", id.id).into())?; + .ok_or::(format!("Invalid token id: {id}").into())?; output.extend_from_slice(slice); } Ok(String::from_utf8(output)?) @@ -932,7 +930,7 @@ pub trait Model { fn init_scratch(&self) -> Self::Scratch; - fn id_to_token_bytes(&self, id: &PipelineToken) -> Option<&[u8]>; + fn id_to_token_bytes(&self, id: u32) -> Option<&[u8]>; } #[allow( @@ -981,7 +979,7 @@ impl Model for PipelineModel { } } - fn id_to_token_bytes(&self, id: &PipelineToken) -> Option<&[u8]> { + fn id_to_token_bytes(&self, id: u32) -> Option<&[u8]> { match self { Self::BPE(model) => model.id_to_token_bytes(id), Self::WordLevel(model) => model.id_to_token_bytes(id), From fd7ce0e7ee8e7d24d65fefb7ae62ac81157e1049 Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Wed, 22 Jul 2026 18:22:42 +0200 Subject: [PATCH 03/17] impl: special tokens --- tokenizers/tk-encode/src/tokenizer/pipeline.rs | 8 +++++++- .../tk-encode/src/vocab/bucket_added_vocabulary.rs | 13 ++++++++++--- tokenizers/tk-encode/src/vocab/buckets.rs | 3 +++ 3 files changed, 20 insertions(+), 4 deletions(-) diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index 295dd0b15..474db894c 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -540,9 +540,15 @@ impl PipelineTokenizer { /// Incomplete: it concatenates raw token bytes only. No decoder, no added- /// vocab lookup, no `skip_special_tokens` — so the `pipeline_decode_oracle` /// test fails on purpose until those land. - pub fn decode(&self, ids: &[u32], _skip_special_tokens: bool) -> Result { + pub fn decode(&self, ids: &[u32], skip_special_tokens: bool) -> Result { let mut output = Vec::with_capacity(ids.len()); for &id in ids { + if let Some(special_token) = self.added_vocabulary.simple_id_to_token_bytes(id) { + if !skip_special_tokens { + output.extend_from_slice(special_token); + } + continue; + } let slice = self .model .id_to_token_bytes(id) diff --git a/tokenizers/tk-encode/src/vocab/bucket_added_vocabulary.rs b/tokenizers/tk-encode/src/vocab/bucket_added_vocabulary.rs index 5bc0be6dd..6ff84fbb7 100644 --- a/tokenizers/tk-encode/src/vocab/bucket_added_vocabulary.rs +++ b/tokenizers/tk-encode/src/vocab/bucket_added_vocabulary.rs @@ -249,10 +249,17 @@ impl AddedVocabulary { /// this returns the cached normalized form so that the configured `Decoder` can /// invert the transformation correctly. For all other tokens, the original /// `content` is returned. - pub fn simple_id_to_token(&self, _id: u32) -> Option { + pub fn simple_id_to_token(&self, id: u32) -> Option { self.vocab - .id_to_token(_id) - .or_else(|| self.normalized_vocab.id_to_token(_id)) + .id_to_token(id) + .or_else(|| self.normalized_vocab.id_to_token(id)) + } + + /// todo: docs + pub fn simple_id_to_token_bytes(&self, id: u32) -> Option<&[u8]> { + self.vocab + .id_to_token_bytes(id) + .or_else(|| self.normalized_vocab.id_to_token_bytes(id)) } // diff --git a/tokenizers/tk-encode/src/vocab/buckets.rs b/tokenizers/tk-encode/src/vocab/buckets.rs index 9a78a8cec..4db9e0202 100644 --- a/tokenizers/tk-encode/src/vocab/buckets.rs +++ b/tokenizers/tk-encode/src/vocab/buckets.rs @@ -443,6 +443,9 @@ impl Buckets { pub fn id_to_token(&self, id: u32) -> Option { self.vocab.id_to_token(id) } + pub fn id_to_token_bytes(&self, id: u32) -> Option<&[u8]> { + self.vocab.id_to_token_bytes(id) + } } impl Default for Buckets { From d6d524cdfbacb107d580a2edaaf98005221df209 Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Fri, 24 Jul 2026 13:06:44 +0200 Subject: [PATCH 04/17] stub PipelineDecoder --- .../tk-encode/src/tokenizer/pipeline.rs | 56 +++++++++++++++++-- 1 file changed, 52 insertions(+), 4 deletions(-) diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index 474db894c..cd796d79a 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -14,6 +14,7 @@ use crate::utils::byte_level::GPT2_REGEX_STR; use crate::vocab::bucket_added_vocabulary::{ AddedToken as BucketAddedToken, AddedVocabulary as BucketAddedVocabulary, }; +use crate::{Decoder, DecoderWrapper}; use crate::{ ModelWrapper, PostProcessorWrapper, PreTokenizerWrapper, Token, Tokenizer, normalizers::NormalizerWrapper, @@ -266,6 +267,51 @@ impl TryFrom<&PostProcessorWrapper> for PipelinePostProcessor { } } +#[derive(Debug, Default)] +pub enum PipelineDecoder { + #[default] + None, +} + +impl TryFrom<&DecoderWrapper> for PipelineDecoder { + type Error = crate::Error; + + fn try_from(value: &DecoderWrapper) -> std::prelude::v1::Result { + match value { + DecoderWrapper::BPE(decoder) => { + Err(format!("Decoder {:?} not supported yet", decoder).into()) + } + DecoderWrapper::ByteFallback(decoder) => { + Err(format!("Decoder {:?} not supported yet", decoder).into()) + } + DecoderWrapper::ByteLevel(decoder) => { + Err(format!("Decoder {:?} not supported yet", decoder).into()) + } + DecoderWrapper::CTC(decoder) => { + Err(format!("Decoder {:?} not supported yet", decoder).into()) + } + DecoderWrapper::Fuse(decoder) => { + Err(format!("Decoder {:?} not supported yet", decoder).into()) + } + DecoderWrapper::Metaspace(decoder) => { + Err(format!("Decoder {:?} not supported yet", decoder).into()) + } + DecoderWrapper::Replace(decoder) => { + Err(format!("Decoder {:?} not supported yet", decoder).into()) + } + DecoderWrapper::Sequence(decoder) => { + Err(format!("Decoder {:?} not supported yet", decoder).into()) + } + DecoderWrapper::Strip(decoder) => { + Err(format!("Decoder {:?} not supported yet", decoder).into()) + } + DecoderWrapper::WordPiece(decoder) => { + Err(format!("Decoder {:?} not supported yet", decoder).into()) + } + } + } +} + /// An output token. Carries only the vocabulary `id` — offsets and the token /// string are dropped, which is all an encode-only caller needs. #[derive(Debug, Clone, Copy)] @@ -391,6 +437,7 @@ pub struct PipelineTokenizer { pre_tokenizer: PipelinePreTokenizer, model: PipelineModel, post_processor: PipelinePostProcessor, + decoder: PipelineDecoder, } impl TryFrom<&Tokenizer> for PipelineTokenizer { @@ -493,6 +540,11 @@ impl TryFrom<&Tokenizer> for PipelineTokenizer { .map(PipelinePostProcessor::try_from) .transpose()? .unwrap_or_default(), + decoder: tok + .get_decoder() + .map(PipelineDecoder::try_from) + .transpose()? + .unwrap_or_default(), }) } } @@ -536,10 +588,6 @@ impl PipelineTokenizer { } /// Decode token ids back to a `String`. - /// - /// Incomplete: it concatenates raw token bytes only. No decoder, no added- - /// vocab lookup, no `skip_special_tokens` — so the `pipeline_decode_oracle` - /// test fails on purpose until those land. pub fn decode(&self, ids: &[u32], skip_special_tokens: bool) -> Result { let mut output = Vec::with_capacity(ids.len()); for &id in ids { From 93a2067959808d4fc9ff3ced23be788729fc280a Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Fri, 24 Jul 2026 13:34:39 +0200 Subject: [PATCH 05/17] implement: no decoder, wordpiece decoder (partial) --- .../tk-encode/src/tokenizer/pipeline.rs | 93 +++++++++++++++---- 1 file changed, 75 insertions(+), 18 deletions(-) diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index cd796d79a..4b0e9b15e 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -4,6 +4,8 @@ use std::{borrow::Cow, convert::TryFrom}; use atomsplit::classify::classify; +use crate::DecoderWrapper; +use crate::decoders::wordpiece::WordPiece; use crate::models::bpe::{BpeScratch, PipelineBPE}; use crate::models::unigram::{Unigram, UnigramScratch}; use crate::models::wordlevel::WordLevel; @@ -14,7 +16,6 @@ use crate::utils::byte_level::GPT2_REGEX_STR; use crate::vocab::bucket_added_vocabulary::{ AddedToken as BucketAddedToken, AddedVocabulary as BucketAddedVocabulary, }; -use crate::{Decoder, DecoderWrapper}; use crate::{ ModelWrapper, PostProcessorWrapper, PreTokenizerWrapper, Token, Tokenizer, normalizers::NormalizerWrapper, @@ -267,26 +268,88 @@ impl TryFrom<&PostProcessorWrapper> for PipelinePostProcessor { } } +/// [`PipelineDecoder`] is responsible for turning a chunk of token ids back into human-readable text #[derive(Debug, Default)] pub enum PipelineDecoder { + WordPiece { + /// Tokens that are inside a word start with this prefix. + /// For example for bert, it's "##": "tokenization" would be tokenized as ["tok", "##eni", "##z", "##ation"] + /// To reconstruct the text from the tokens, we need to strip that prefix for every token + word_continuation_prefix: String, + /// Whether to perform a cleanup phase (see implementation for details) + cleanup: bool, + }, #[default] None, } +pub trait Decoder<'a> { + fn decode( + &'a self, + model: &'a PipelineModel, + added_vocabulary: &BucketAddedVocabulary, + token_id: u32, + decoded: &mut Vec, + skip_special_tokens: bool, + ) -> Result<()> { + if let Some(special_token) = added_vocabulary.simple_id_to_token_bytes(token_id) { + if !skip_special_tokens { + decoded.extend_from_slice(special_token); + } + return Ok(()); + } + let slice = self.decode_token_to_slice(model, token_id)?; + decoded.extend_from_slice(slice); + Ok(()) + } + + fn decode_token_to_slice(&'a self, model: &'a PipelineModel, token_id: u32) + -> Result<&'a [u8]>; +} + +impl<'a> Decoder<'a> for PipelineDecoder { + fn decode_token_to_slice( + &'a self, + model: &'a PipelineModel, + token_id: u32, + ) -> Result<&'a [u8]> { + let mut bytes = model + .id_to_token_bytes(token_id) + .ok_or::(format!("Invalid token id: {token_id}").into())?; + match self { + Self::None => Ok(bytes), + PipelineDecoder::WordPiece { + word_continuation_prefix, + .. + } => { + if bytes.starts_with(word_continuation_prefix.as_bytes()) { + // trim prefix + bytes = &bytes[word_continuation_prefix.len()..]; + } + // todo: cleanup + Ok(bytes) + } + } + } +} + impl TryFrom<&DecoderWrapper> for PipelineDecoder { type Error = crate::Error; fn try_from(value: &DecoderWrapper) -> std::prelude::v1::Result { match value { + // ByteLevel decoder is no longer needed as the vocabulary is stored as raw bytes + DecoderWrapper::ByteLevel(_) => Ok(Self::None), + DecoderWrapper::WordPiece(WordPiece { cleanup, prefix }) => Ok(Self::WordPiece { + word_continuation_prefix: prefix.clone(), + cleanup: *cleanup, + }), DecoderWrapper::BPE(decoder) => { Err(format!("Decoder {:?} not supported yet", decoder).into()) } DecoderWrapper::ByteFallback(decoder) => { Err(format!("Decoder {:?} not supported yet", decoder).into()) } - DecoderWrapper::ByteLevel(decoder) => { - Err(format!("Decoder {:?} not supported yet", decoder).into()) - } DecoderWrapper::CTC(decoder) => { Err(format!("Decoder {:?} not supported yet", decoder).into()) } @@ -305,9 +368,6 @@ impl TryFrom<&DecoderWrapper> for PipelineDecoder { DecoderWrapper::Strip(decoder) => { Err(format!("Decoder {:?} not supported yet", decoder).into()) } - DecoderWrapper::WordPiece(decoder) => { - Err(format!("Decoder {:?} not supported yet", decoder).into()) - } } } } @@ -590,18 +650,15 @@ impl PipelineTokenizer { /// Decode token ids back to a `String`. pub fn decode(&self, ids: &[u32], skip_special_tokens: bool) -> Result { let mut output = Vec::with_capacity(ids.len()); + for &id in ids { - if let Some(special_token) = self.added_vocabulary.simple_id_to_token_bytes(id) { - if !skip_special_tokens { - output.extend_from_slice(special_token); - } - continue; - } - let slice = self - .model - .id_to_token_bytes(id) - .ok_or::(format!("Invalid token id: {id}").into())?; - output.extend_from_slice(slice); + self.decoder.decode( + &self.model, + &self.added_vocabulary, + id, + &mut output, + skip_special_tokens, + )?; } Ok(String::from_utf8(output)?) } From 5689dc5d10a7c8db040122420d62ee5f39893a99 Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Fri, 24 Jul 2026 13:49:58 +0200 Subject: [PATCH 06/17] tests give where it fails + better implem (?) --- .../tk-encode/src/tokenizer/pipeline.rs | 46 +++++++++++------- .../tk-encode/tests/pipeline_decode_oracle.rs | 47 +++++++++++++++++-- 2 files changed, 72 insertions(+), 21 deletions(-) diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index 4b0e9b15e..5e6259be3 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -283,10 +283,10 @@ pub enum PipelineDecoder { None, } -pub trait Decoder<'a> { +pub trait Decoder { fn decode( - &'a self, - model: &'a PipelineModel, + &self, + model: &PipelineModel, added_vocabulary: &BucketAddedVocabulary, token_id: u32, decoded: &mut Vec, @@ -298,36 +298,46 @@ pub trait Decoder<'a> { } return Ok(()); } - let slice = self.decode_token_to_slice(model, token_id)?; - decoded.extend_from_slice(slice); + self.decode_token(model, token_id, decoded)?; Ok(()) } - fn decode_token_to_slice(&'a self, model: &'a PipelineModel, token_id: u32) - -> Result<&'a [u8]>; + fn decode_token( + &self, + model: &PipelineModel, + token_id: u32, + decoded: &mut Vec, + ) -> Result<()>; } -impl<'a> Decoder<'a> for PipelineDecoder { - fn decode_token_to_slice( - &'a self, - model: &'a PipelineModel, +impl Decoder for PipelineDecoder { + fn decode_token( + &self, + model: &PipelineModel, token_id: u32, - ) -> Result<&'a [u8]> { - let mut bytes = model + decoded: &mut Vec, + ) -> Result<()> { + let bytes = model .id_to_token_bytes(token_id) .ok_or::(format!("Invalid token id: {token_id}").into())?; match self { - Self::None => Ok(bytes), + Self::None => { + decoded.extend_from_slice(bytes); + Ok(()) + } PipelineDecoder::WordPiece { word_continuation_prefix, .. } => { if bytes.starts_with(word_continuation_prefix.as_bytes()) { // trim prefix - bytes = &bytes[word_continuation_prefix.len()..]; + decoded.extend_from_slice(&bytes[word_continuation_prefix.len()..]); + } else { + decoded.push(b' '); + decoded.extend_from_slice(bytes); } - // todo: cleanup - Ok(bytes) + // todo: cleanup phase + Ok(()) } } } @@ -649,7 +659,7 @@ impl PipelineTokenizer { /// Decode token ids back to a `String`. pub fn decode(&self, ids: &[u32], skip_special_tokens: bool) -> Result { - let mut output = Vec::with_capacity(ids.len()); + let mut output = Vec::with_capacity(2 * ids.len()); for &id in ids { self.decoder.decode( diff --git a/tokenizers/tk-encode/tests/pipeline_decode_oracle.rs b/tokenizers/tk-encode/tests/pipeline_decode_oracle.rs index 107821160..db0f4442f 100644 --- a/tokenizers/tk-encode/tests/pipeline_decode_oracle.rs +++ b/tokenizers/tk-encode/tests/pipeline_decode_oracle.rs @@ -113,12 +113,18 @@ fn check_model(tok_file: &str) { ); match pipeline.decode(&ids, skip_special_tokens) { Ok(got) if got == expected => {} - Ok(_) => failures.push(format!("{ctx}: decode mismatch")), + Ok(got) => failures.push(format!( + "{ctx}: decode mismatch, {}", + divergence(&expected, &got) + )), Err(e) => failures.push(format!("{ctx}: decode error: {e}")), } match stream_decode(&pipeline, &ids, skip_special_tokens) { Ok(got) if got == expected => {} - Ok(_) => failures.push(format!("{ctx}: decode_stream mismatch")), + Ok(got) => failures.push(format!( + "{ctx}: decode_stream mismatch, {}", + divergence(&expected, &got) + )), Err(e) => failures.push(format!("{ctx}: decode_stream error: {e}")), } } @@ -134,7 +140,13 @@ fn check_model(tok_file: &str) { let expected = released.decode_batch(&sentences, false).unwrap(); match pipeline.decode_batch(&sentences, false) { Ok(got) if got == expected => {} - Ok(_) => failures.push("decode_batch mismatch".into()), + Ok(got) => { + let i = expected.iter().zip(&got).position(|(e, g)| e != g).unwrap(); + failures.push(format!( + "decode_batch mismatch at sentence {i}, {}", + divergence(&expected[i], &got[i]) + )); + } Err(e) => failures.push(format!("decode_batch error: {e}")), } } @@ -149,6 +161,35 @@ fn check_model(tok_file: &str) { ); } +/// Show where `got` first diverges from `expected`, with nearby text from both +/// sides — fixture windows are kilobytes, so printing whole strings would bury +/// the interesting byte. +fn divergence(expected: &str, got: &str) -> String { + let byte = expected + .bytes() + .zip(got.bytes()) + .position(|(e, g)| e != g) + .unwrap_or_else(|| expected.len().min(got.len())); + let excerpt = |s: &str| { + let mut start = byte.saturating_sub(40); + while !s.is_char_boundary(start) { + start -= 1; + } + let mut end = (byte + 40).min(s.len()); + while !s.is_char_boundary(end) { + end += 1; + } + format!("…{:?}…", &s[start..end]) + }; + format!( + "first divergence at byte {byte}\n expected ({} B): {}\n got ({} B): {}", + expected.len(), + excerpt(expected), + got.len(), + excerpt(got), + ) +} + /// Feed `ids` through [`PipelineTokenizer::decode_stream`] one at a time and /// concatenate the emitted chunks — for a complete id sequence this must equal a /// one-shot `decode`, so the oracle can compare it against the release directly. From fb1ff9c2cfec7cfa5778204c781780aff7e9a4e2 Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Sat, 25 Jul 2026 19:07:53 +0200 Subject: [PATCH 07/17] fix: bert uncased --- .../tk-encode/src/decoders/wordpiece.rs | 36 +++++++- .../tk-encode/src/tokenizer/pipeline.rs | 85 ++++++++----------- 2 files changed, 70 insertions(+), 51 deletions(-) diff --git a/tokenizers/tk-encode/src/decoders/wordpiece.rs b/tokenizers/tk-encode/src/decoders/wordpiece.rs index a2da414c0..196e8d922 100644 --- a/tokenizers/tk-encode/src/decoders/wordpiece.rs +++ b/tokenizers/tk-encode/src/decoders/wordpiece.rs @@ -1,4 +1,7 @@ -use crate::tokenizer::{Decoder, Result}; +use crate::{ + pipeline, + tokenizer::{Decoder, Result}, +}; use serde::{Deserialize, Serialize}; @@ -28,6 +31,7 @@ impl Default for WordPiece { } } } + pub fn cleanup(dirty_input: &str) -> String { dirty_input .replace(" .", ".") @@ -61,6 +65,36 @@ impl Decoder for WordPiece { } } +const CLEANUP_LIST: [&'static [u8]; 10] = [ + b".", b"?", b"!", b",", b"n't", b"'m", b"do not", b"'s", b"'ve", b"'re" +]; + +impl pipeline::Decoder for WordPiece { + fn decode_token( + &self, + token_bytes: &[u8], + token_index: usize, + decoded: &mut Vec, + ) -> Result<()> { + if token_index == 0 { + decoded.extend_from_slice(token_bytes); + return Ok(()); + } + + if token_bytes.starts_with(self.prefix.as_bytes()) { + // trim prefix + decoded.extend_from_slice(&token_bytes[self.prefix.len()..]); + } else { + if !self.cleanup || !CLEANUP_LIST.contains(&token_bytes) { + decoded.push(b' '); + } + decoded.extend_from_slice(token_bytes); + } + + Ok(()) + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index 5e6259be3..162f8b510 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -271,14 +271,7 @@ impl TryFrom<&PostProcessorWrapper> for PipelinePostProcessor { /// [`PipelineDecoder`] is responsible for turning a chunk of token ids back into human-readable text #[derive(Debug, Default)] pub enum PipelineDecoder { - WordPiece { - /// Tokens that are inside a word start with this prefix. - /// For example for bert, it's "##": "tokenization" would be tokenized as ["tok", "##eni", "##z", "##ation"] - /// To reconstruct the text from the tokens, we need to strip that prefix for every token - word_continuation_prefix: String, - /// Whether to perform a cleanup phase (see implementation for details) - cleanup: bool, - }, + WordPiece(WordPiece), #[default] None, } @@ -288,24 +281,33 @@ pub trait Decoder { &self, model: &PipelineModel, added_vocabulary: &BucketAddedVocabulary, - token_id: u32, - decoded: &mut Vec, skip_special_tokens: bool, + ids: &[u32], + decoded: &mut Vec, ) -> Result<()> { - if let Some(special_token) = added_vocabulary.simple_id_to_token_bytes(token_id) { - if !skip_special_tokens { - decoded.extend_from_slice(special_token); + let mut token_index = 0; + for &token_id in ids { + if let Some(special) = added_vocabulary.simple_id_to_token_bytes(token_id) { + if skip_special_tokens { + continue; + } + self.decode_token(special, token_index, decoded)?; + } else { + let token_bytes = model + .id_to_token_bytes(token_id) + .ok_or::(format!("Invalid token id: {token_id}").into())?; + + self.decode_token(token_bytes, token_index, decoded)?; } - return Ok(()); + token_index += 1; } - self.decode_token(model, token_id, decoded)?; Ok(()) } fn decode_token( &self, - model: &PipelineModel, - token_id: u32, + token_bytes: &[u8], + token_index: usize, decoded: &mut Vec, ) -> Result<()>; } @@ -313,32 +315,21 @@ pub trait Decoder { impl Decoder for PipelineDecoder { fn decode_token( &self, - model: &PipelineModel, - token_id: u32, + token_bytes: &[u8], + token_index: usize, decoded: &mut Vec, ) -> Result<()> { - let bytes = model - .id_to_token_bytes(token_id) - .ok_or::(format!("Invalid token id: {token_id}").into())?; match self { Self::None => { - decoded.extend_from_slice(bytes); - Ok(()) - } - PipelineDecoder::WordPiece { - word_continuation_prefix, - .. - } => { - if bytes.starts_with(word_continuation_prefix.as_bytes()) { - // trim prefix - decoded.extend_from_slice(&bytes[word_continuation_prefix.len()..]); - } else { + if token_index != 0 { decoded.push(b' '); - decoded.extend_from_slice(bytes); } - // todo: cleanup phase + decoded.extend_from_slice(token_bytes); Ok(()) } + PipelineDecoder::WordPiece(decoder) => { + decoder.decode_token(token_bytes, token_index, decoded) + } } } } @@ -350,10 +341,7 @@ impl TryFrom<&DecoderWrapper> for PipelineDecoder { match value { // ByteLevel decoder is no longer needed as the vocabulary is stored as raw bytes DecoderWrapper::ByteLevel(_) => Ok(Self::None), - DecoderWrapper::WordPiece(WordPiece { cleanup, prefix }) => Ok(Self::WordPiece { - word_continuation_prefix: prefix.clone(), - cleanup: *cleanup, - }), + DecoderWrapper::WordPiece(decoder) => Ok(Self::WordPiece(decoder.clone())), DecoderWrapper::BPE(decoder) => { Err(format!("Decoder {:?} not supported yet", decoder).into()) } @@ -659,17 +647,14 @@ impl PipelineTokenizer { /// Decode token ids back to a `String`. pub fn decode(&self, ids: &[u32], skip_special_tokens: bool) -> Result { - let mut output = Vec::with_capacity(2 * ids.len()); - - for &id in ids { - self.decoder.decode( - &self.model, - &self.added_vocabulary, - id, - &mut output, - skip_special_tokens, - )?; - } + let mut output = Vec::with_capacity(4 * ids.len()); + self.decoder.decode( + &self.model, + &self.added_vocabulary, + skip_special_tokens, + ids, + &mut output, + )?; Ok(String::from_utf8(output)?) } From 7fa6783aa8701e5dee7034e772b5a310d1b87402 Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Sat, 25 Jul 2026 19:12:11 +0200 Subject: [PATCH 08/17] fmt --- tokenizers/tk-encode/src/decoders/wordpiece.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tokenizers/tk-encode/src/decoders/wordpiece.rs b/tokenizers/tk-encode/src/decoders/wordpiece.rs index 196e8d922..8376e9a7d 100644 --- a/tokenizers/tk-encode/src/decoders/wordpiece.rs +++ b/tokenizers/tk-encode/src/decoders/wordpiece.rs @@ -65,8 +65,8 @@ impl Decoder for WordPiece { } } -const CLEANUP_LIST: [&'static [u8]; 10] = [ - b".", b"?", b"!", b",", b"n't", b"'m", b"do not", b"'s", b"'ve", b"'re" +const CLEANUP_LIST: [&[u8]; 10] = [ + b".", b"?", b"!", b",", b"n't", b"'m", b"do not", b"'s", b"'ve", b"'re", ]; impl pipeline::Decoder for WordPiece { From 75e4517f1468997dcb0160328d4749fe95192e90 Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Sun, 26 Jul 2026 18:14:32 +0200 Subject: [PATCH 09/17] fix: fixture bench does not panic --- .../tk-encode/examples/fixture_bench.rs | 25 +++++++++++++------ 1 file changed, 18 insertions(+), 7 deletions(-) diff --git a/tokenizers/tk-encode/examples/fixture_bench.rs b/tokenizers/tk-encode/examples/fixture_bench.rs index 9ac6d89bf..ccb3191ca 100644 --- a/tokenizers/tk-encode/examples/fixture_bench.rs +++ b/tokenizers/tk-encode/examples/fixture_bench.rs @@ -670,7 +670,9 @@ fn bench_threads( // release. No released baseline → no decode oracle, so the whole phase is null. // `pipeline_ok` is the once-probed "can the pipeline decode yet" flag: while // `PipelineTokenizer::decode` is a loud stub it is false, so the pipeline series -// is `null` (rendered "pending") and only the baseline bar is drawn. +// is `null` (rendered "pending") and only the baseline bar is drawn. A decoder +// that passes that probe but `Err`s on real ids fails `text_match` and nulls +// its series for the affected fixtures/sweep — it never aborts the run. /// Encode every chunk with the released crate into its id stream (untimed input; /// specials included, so decode sees the frame tokens a real stream carries). @@ -724,22 +726,28 @@ fn bench_decode( }; let mbps = |secs: f64| dec_bytes as f64 / secs / 1e6; + // The `main` probe only decodes `[0]`, so a partial decoder can still `Err` + // on this fixture's real id streams — that fails the `text_match` gate and + // skips the pipeline timing (series stays null) instead of aborting the run. + let pipe_ok = pipeline_ok && ids.iter().all(|i| pipeline.decode(i, false).is_ok()); // Correctness gate (first 3 chunks): pipeline decode == released decode. let text_match = pipeline_ok.then(|| { - ids.iter() - .take(3) - .all(|i| pipeline.decode(i, false).unwrap() == baseline.decode(i, false).unwrap()) + pipe_ok + && ids + .iter() + .take(3) + .all(|i| pipeline.decode(i, false).unwrap() == baseline.decode(i, false).unwrap()) }); // Interleaved warm-up + REPS so thermal drift hits both equally. one_pass(&|i| baseline.decode(i, false).unwrap().len()); - if pipeline_ok { + if pipe_ok { one_pass(&|i| pipeline.decode(i, false).unwrap().len()); } let (mut base_s, mut pipe_s) = (Vec::new(), Vec::new()); for _ in 0..REPS { base_s.push(one_pass(&|i| baseline.decode(i, false).unwrap().len())); - if pipeline_ok { + if pipe_ok { pipe_s.push(one_pass(&|i| pipeline.decode(i, false).unwrap().len())); } } @@ -793,11 +801,14 @@ fn bench_decode_threads( .iter() .map(|i| baseline.decode(i, false).unwrap().len()) .sum(); + // Same tolerance as `bench_decode`: any `Err` over the corpus drops the + // pipeline series to null instead of panicking mid-sweep. + let pipe_ok = pipeline_ok && ids.iter().all(|i| pipeline.decode(i, false).is_ok()); let counts = thread_counts(); let (mut pipe, mut base) = (Vec::new(), Vec::new()); for &n in &counts { let b = par_decode_mbps(|i| baseline.decode(i, false).unwrap().len(), &ids, bytes, n); - let p = pipeline_ok + let p = pipe_ok .then(|| par_decode_mbps(|i| pipeline.decode(i, false).unwrap().len(), &ids, bytes, n)); eprintln!( " decode {n} thread(s): pipeline {}, baseline {b:.1} MB/s", From 21b9a49482431faa2525ccce223ebb1f0152bbf0 Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Sun, 26 Jul 2026 18:23:42 +0200 Subject: [PATCH 10/17] fix: differentiate None vs JoinWithSpaces --- tokenizers/tk-encode/src/tokenizer/pipeline.rs | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index 162f8b510..b18e7502d 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -272,8 +272,9 @@ impl TryFrom<&PostProcessorWrapper> for PipelinePostProcessor { #[derive(Debug, Default)] pub enum PipelineDecoder { WordPiece(WordPiece), - #[default] None, + #[default] + JoinWithSpaces, } pub trait Decoder { @@ -321,6 +322,10 @@ impl Decoder for PipelineDecoder { ) -> Result<()> { match self { Self::None => { + decoded.extend_from_slice(token_bytes); + Ok(()) + } + Self::JoinWithSpaces => { if token_index != 0 { decoded.push(b' '); } From e147433cce80d2bc140feee055963b6de4493d1b Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Sun, 26 Jul 2026 18:48:28 +0200 Subject: [PATCH 11/17] implement metaspace decoder --- .../tk-encode/src/pre_tokenizers/metaspace.rs | 24 +++++++++++- .../tk-encode/src/tokenizer/pipeline.rs | 39 +++++++------------ 2 files changed, 36 insertions(+), 27 deletions(-) diff --git a/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs b/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs index 2d7b7b0aa..0f7dc661b 100644 --- a/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs +++ b/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs @@ -1,4 +1,7 @@ -use crate::tokenizer::{Decoder, PreTokenizedString, PreTokenizer, Result, SplitDelimiterBehavior}; +use crate::{ + pipeline, + tokenizer::{Decoder, PreTokenizedString, PreTokenizer, Result, SplitDelimiterBehavior}, +}; use serde::{Deserialize, Deserializer, Serialize, de}; /// Enum representing options for the metaspace prepending scheme. @@ -172,6 +175,25 @@ impl Decoder for Metaspace { } } +impl pipeline::Decoder for Metaspace { + fn decode_token( + &self, + token_bytes: &[u8], + token_index: usize, + decoded: &mut Vec, + ) -> Result<()> { + if token_bytes.starts_with(self.str_rep.as_bytes()) { + if token_index == 0 && self.prepend_scheme != PrependScheme::Never { + decoded.push(b' '); + } + decoded.extend_from_slice(&token_bytes[self.replacement.len_utf8()..]); + } else { + decoded.extend_from_slice(token_bytes); + } + Ok(()) + } +} + #[cfg(test)] mod tests { use regex::Regex; diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index b18e7502d..17273feb6 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -5,6 +5,7 @@ use std::{borrow::Cow, convert::TryFrom}; use atomsplit::classify::classify; use crate::DecoderWrapper; +use crate::decoders::metaspace::Metaspace; use crate::decoders::wordpiece::WordPiece; use crate::models::bpe::{BpeScratch, PipelineBPE}; use crate::models::unigram::{Unigram, UnigramScratch}; @@ -272,6 +273,7 @@ impl TryFrom<&PostProcessorWrapper> for PipelinePostProcessor { #[derive(Debug, Default)] pub enum PipelineDecoder { WordPiece(WordPiece), + MetaSpace(Metaspace), None, #[default] JoinWithSpaces, @@ -332,9 +334,16 @@ impl Decoder for PipelineDecoder { decoded.extend_from_slice(token_bytes); Ok(()) } - PipelineDecoder::WordPiece(decoder) => { - decoder.decode_token(token_bytes, token_index, decoded) + Self::WordPiece(decoder) => { + decoder.decode_token(token_bytes, token_index, decoded)?; + if let Some(&last) = decoded.last() + && last == b' ' + { + decoded.pop(); + } + Ok(()) } + Self::MetaSpace(decoder) => decoder.decode_token(token_bytes, token_index, decoded), } } } @@ -347,30 +356,8 @@ impl TryFrom<&DecoderWrapper> for PipelineDecoder { // ByteLevel decoder is no longer needed as the vocabulary is stored as raw bytes DecoderWrapper::ByteLevel(_) => Ok(Self::None), DecoderWrapper::WordPiece(decoder) => Ok(Self::WordPiece(decoder.clone())), - DecoderWrapper::BPE(decoder) => { - Err(format!("Decoder {:?} not supported yet", decoder).into()) - } - DecoderWrapper::ByteFallback(decoder) => { - Err(format!("Decoder {:?} not supported yet", decoder).into()) - } - DecoderWrapper::CTC(decoder) => { - Err(format!("Decoder {:?} not supported yet", decoder).into()) - } - DecoderWrapper::Fuse(decoder) => { - Err(format!("Decoder {:?} not supported yet", decoder).into()) - } - DecoderWrapper::Metaspace(decoder) => { - Err(format!("Decoder {:?} not supported yet", decoder).into()) - } - DecoderWrapper::Replace(decoder) => { - Err(format!("Decoder {:?} not supported yet", decoder).into()) - } - DecoderWrapper::Sequence(decoder) => { - Err(format!("Decoder {:?} not supported yet", decoder).into()) - } - DecoderWrapper::Strip(decoder) => { - Err(format!("Decoder {:?} not supported yet", decoder).into()) - } + DecoderWrapper::Metaspace(decoder) => Ok(Self::MetaSpace(decoder.clone())), + decoder => Err(format!("Decoder {:?} not supported yet", decoder).into()), } } } From 0e9f6e757ad70940b10d6d3b02a27365d9e9522c Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Sun, 26 Jul 2026 20:28:04 +0200 Subject: [PATCH 12/17] stub: Sequence decoder --- tokenizers/tk-encode/src/tokenizer/pipeline.rs | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index 17273feb6..eb82e4e8a 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -272,6 +272,7 @@ impl TryFrom<&PostProcessorWrapper> for PipelinePostProcessor { /// [`PipelineDecoder`] is responsible for turning a chunk of token ids back into human-readable text #[derive(Debug, Default)] pub enum PipelineDecoder { + Sequence(Vec), WordPiece(WordPiece), MetaSpace(Metaspace), None, @@ -357,6 +358,22 @@ impl TryFrom<&DecoderWrapper> for PipelineDecoder { DecoderWrapper::ByteLevel(_) => Ok(Self::None), DecoderWrapper::WordPiece(decoder) => Ok(Self::WordPiece(decoder.clone())), DecoderWrapper::Metaspace(decoder) => Ok(Self::MetaSpace(decoder.clone())), + DecoderWrapper::Sequence(sequence) => { + let decoders = sequence.get_decoders(); + if decoders.len() == 0 { + return Ok(Self::default()); + } + if decoders.len() == 1 { + // SAFETY: safe to .unwrap() as the len is asserted to be 1 + return Self::try_from(decoders.first().unwrap()); + } + Ok(Self::Sequence( + decoders + .into_iter() + .map(Self::try_from) + .collect::>>()?, + )) + } decoder => Err(format!("Decoder {:?} not supported yet", decoder).into()), } } From 47ec010f2383dc59ea2e08f3026fb0d5b740912e Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Sun, 26 Jul 2026 21:44:07 +0200 Subject: [PATCH 13/17] wip: sequence --- .../tk-encode/src/decoders/byte_fallback.rs | 39 ++++- .../tk-encode/src/decoders/wordpiece.rs | 26 ++-- .../tk-encode/src/pre_tokenizers/metaspace.rs | 9 +- .../tk-encode/src/tokenizer/pipeline.rs | 133 +++++++++++++----- 4 files changed, 155 insertions(+), 52 deletions(-) diff --git a/tokenizers/tk-encode/src/decoders/byte_fallback.rs b/tokenizers/tk-encode/src/decoders/byte_fallback.rs index 57b7b63cd..b575da989 100644 --- a/tokenizers/tk-encode/src/decoders/byte_fallback.rs +++ b/tokenizers/tk-encode/src/decoders/byte_fallback.rs @@ -1,4 +1,8 @@ -use crate::tokenizer::{Decoder, Result}; +use crate::{ + pipeline::{self, DecoderState}, + tokenizer::{Decoder, Result}, +}; +use ahash::AHashMap; use monostate::MustBe; use serde::{Deserialize, Serialize}; @@ -11,12 +15,16 @@ use serde::{Deserialize, Serialize}; pub struct ByteFallback { #[serde(rename = "type")] type_: MustBe!("ByteFallback"), + /// Lookup mapping a token id to the raw byte it represents + /// todo: closed-addressing, can use ptrhash + fallback_lookup: AHashMap, } impl ByteFallback { - pub fn new() -> Self { + pub fn new(fallback_lookup: AHashMap) -> Self { Self { type_: MustBe!("ByteFallback"), + fallback_lookup, } } } @@ -62,13 +70,38 @@ impl Decoder for ByteFallback { } } +impl pipeline::Decoder for ByteFallback { + fn decode_token( + &self, + state: &mut pipeline::DecoderState, + token_id: u32, + token_bytes: &[u8], + decoded: &mut Vec, + ) -> Result<()> { + if let Some(&raw_byte) = self.fallback_lookup.get(&token_id) { + state.pending_buffer.push(raw_byte); + return Ok(()); + } + self.flush(state, decoded)?; + decoded.extend_from_slice(token_bytes); + Ok(()) + } + + fn flush(&self, state: &mut DecoderState, decoded: &mut Vec) -> Result<()> { + if state.pending_buffer.len() > 0 { + decoded.append(&mut state.pending_buffer); + } + Ok(()) + } +} + #[cfg(test)] mod tests { use super::*; #[test] fn decode() { - let decoder = ByteFallback::new(); + let decoder = ByteFallback::new(AHashMap::new()); let res = decoder .decode_chain(vec!["Hey".into(), "friend!".into()]) .unwrap(); diff --git a/tokenizers/tk-encode/src/decoders/wordpiece.rs b/tokenizers/tk-encode/src/decoders/wordpiece.rs index 8376e9a7d..05f9efb61 100644 --- a/tokenizers/tk-encode/src/decoders/wordpiece.rs +++ b/tokenizers/tk-encode/src/decoders/wordpiece.rs @@ -1,5 +1,7 @@ +use std::mem::replace; + use crate::{ - pipeline, + pipeline::{self, DecoderState}, tokenizer::{Decoder, Result}, }; @@ -72,25 +74,21 @@ const CLEANUP_LIST: [&[u8]; 10] = [ impl pipeline::Decoder for WordPiece { fn decode_token( &self, + state: &mut DecoderState, + _token_id: u32, token_bytes: &[u8], - token_index: usize, decoded: &mut Vec, ) -> Result<()> { - if token_index == 0 { - decoded.extend_from_slice(token_bytes); - return Ok(()); - } - - if token_bytes.starts_with(self.prefix.as_bytes()) { - // trim prefix + if !replace(&mut state.started, true) { + decoded.extend(token_bytes) + } else if token_bytes.starts_with(self.prefix.as_bytes()) { decoded.extend_from_slice(&token_bytes[self.prefix.len()..]); + } else if self.cleanup && CLEANUP_LIST.contains(&token_bytes) { + decoded.extend(token_bytes); } else { - if !self.cleanup || !CLEANUP_LIST.contains(&token_bytes) { - decoded.push(b' '); - } - decoded.extend_from_slice(token_bytes); + decoded.push(b' '); + decoded.extend(token_bytes); } - Ok(()) } } diff --git a/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs b/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs index 0f7dc661b..24aca962c 100644 --- a/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs +++ b/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs @@ -1,5 +1,7 @@ +use std::mem::replace; + use crate::{ - pipeline, + pipeline::{self, DecoderState}, tokenizer::{Decoder, PreTokenizedString, PreTokenizer, Result, SplitDelimiterBehavior}, }; use serde::{Deserialize, Deserializer, Serialize, de}; @@ -178,12 +180,13 @@ impl Decoder for Metaspace { impl pipeline::Decoder for Metaspace { fn decode_token( &self, + state: &mut DecoderState, + _token_id: u32, token_bytes: &[u8], - token_index: usize, decoded: &mut Vec, ) -> Result<()> { if token_bytes.starts_with(self.str_rep.as_bytes()) { - if token_index == 0 && self.prepend_scheme != PrependScheme::Never { + if !replace(&mut state.started, true) && self.prepend_scheme != PrependScheme::Never { decoded.push(b' '); } decoded.extend_from_slice(&token_bytes[self.replacement.len_utf8()..]); diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index eb82e4e8a..4505f8713 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -1,10 +1,14 @@ use std::cell::RefCell; use std::convert::TryInto; +use std::iter::zip; +use std::mem::{replace, swap}; +use std::u32; use std::{borrow::Cow, convert::TryFrom}; use atomsplit::classify::classify; use crate::DecoderWrapper; +use crate::decoders::byte_fallback::ByteFallback; use crate::decoders::metaspace::Metaspace; use crate::decoders::wordpiece::WordPiece; use crate::models::bpe::{BpeScratch, PipelineBPE}; @@ -275,11 +279,20 @@ pub enum PipelineDecoder { Sequence(Vec), WordPiece(WordPiece), MetaSpace(Metaspace), + ByteFallback(ByteFallback), None, #[default] JoinWithSpaces, } +#[derive(Default)] +pub struct DecoderState { + pub(crate) pending_buffer: Vec, + pub(crate) swap_buffers: [Vec; 2], + pub(crate) started: bool, + pub(crate) members: Vec, +} + pub trait Decoder { fn decode( &self, @@ -288,39 +301,41 @@ pub trait Decoder { skip_special_tokens: bool, ids: &[u32], decoded: &mut Vec, - ) -> Result<()> { - let mut token_index = 0; + ) -> crate::Result<()> { + let mut state = DecoderState::default(); for &token_id in ids { - if let Some(special) = added_vocabulary.simple_id_to_token_bytes(token_id) { - if skip_special_tokens { - continue; - } - self.decode_token(special, token_index, decoded)?; - } else { - let token_bytes = model + let token_bytes = match added_vocabulary.simple_id_to_token_bytes(token_id) { + Some(_) if skip_special_tokens => continue, + Some(bytes) => bytes, + None => model .id_to_token_bytes(token_id) - .ok_or::(format!("Invalid token id: {token_id}").into())?; - - self.decode_token(token_bytes, token_index, decoded)?; - } - token_index += 1; + .ok_or(format!("Invalid token id: {token_id}"))?, + }; + self.decode_token(&mut state, token_id, token_bytes, decoded)?; } + self.flush(&mut state, decoded)?; Ok(()) } fn decode_token( &self, + state: &mut DecoderState, + token_id: u32, token_bytes: &[u8], - token_index: usize, decoded: &mut Vec, ) -> Result<()>; + + fn flush(&self, _state: &mut DecoderState, _decoded: &mut Vec) -> Result<()> { + return Ok(()); + } } impl Decoder for PipelineDecoder { fn decode_token( &self, + state: &mut DecoderState, + token_id: u32, token_bytes: &[u8], - token_index: usize, decoded: &mut Vec, ) -> Result<()> { match self { @@ -329,22 +344,69 @@ impl Decoder for PipelineDecoder { Ok(()) } Self::JoinWithSpaces => { - if token_index != 0 { + if replace(&mut state.started, true) { decoded.push(b' '); } decoded.extend_from_slice(token_bytes); Ok(()) } - Self::WordPiece(decoder) => { - decoder.decode_token(token_bytes, token_index, decoded)?; - if let Some(&last) = decoded.last() - && last == b' ' - { - decoded.pop(); + Self::WordPiece(decoder) => decoder.decode_token(state, token_id, token_bytes, decoded), + Self::MetaSpace(decoder) => decoder.decode_token(state, token_id, token_bytes, decoded), + Self::ByteFallback(decoder) => { + decoder.decode_token(state, token_id, token_bytes, decoded) + } + + Self::Sequence(stages) => { + let DecoderState { + swap_buffers: [a, b], + members, + .. + } = state; + if members.is_empty() { + members.resize_with(stages.len(), Default::default); + } + a.copy_from_slice(token_bytes); + for (stage, member_state) in zip(stages, members) { + b.clear(); + stage.decode_token(member_state, token_id, a, b)?; + if b.is_empty() { + return Ok(()); + } + swap(a, b); + } + decoded.extend_from_slice(a); + Ok(()) + } + } + } + + fn flush(&self, state: &mut DecoderState, decoded: &mut Vec) -> Result<()> { + match self { + Self::Sequence(stages) => { + let DecoderState { + swap_buffers: [a, b], + members, + .. + } = state; + for i in 0..stages.len() { + a.clear(); + stages[i].flush(&mut members[i], a)?; + if a.is_empty() { + continue; + } + for (stage, member_state) in stages.iter().zip(members.iter_mut()).skip(i) { + b.clear(); + stage.decode_token(member_state, u32::MAX, a, b)?; + if b.is_empty() { + return Ok(()); + } + std::mem::swap(a, b); + } + decoded.extend_from_slice(a); } Ok(()) } - Self::MetaSpace(decoder) => decoder.decode_token(token_bytes, token_index, decoded), + decoder => decoder.flush(state, decoded), } } } @@ -359,20 +421,27 @@ impl TryFrom<&DecoderWrapper> for PipelineDecoder { DecoderWrapper::WordPiece(decoder) => Ok(Self::WordPiece(decoder.clone())), DecoderWrapper::Metaspace(decoder) => Ok(Self::MetaSpace(decoder.clone())), DecoderWrapper::Sequence(sequence) => { - let decoders = sequence.get_decoders(); + let mut decoders = sequence + .get_decoders() + .into_iter() + .map(Self::try_from) + .filter(|maybe_decoder| { + if let Ok(Self::None) = maybe_decoder { + false + } else { + true + } + }) + .collect::>>()?; if decoders.len() == 0 { return Ok(Self::default()); } if decoders.len() == 1 { // SAFETY: safe to .unwrap() as the len is asserted to be 1 - return Self::try_from(decoders.first().unwrap()); + let first = decoders.pop().unwrap(); + return Ok(first); } - Ok(Self::Sequence( - decoders - .into_iter() - .map(Self::try_from) - .collect::>>()?, - )) + Ok(Self::Sequence(decoders)) } decoder => Err(format!("Decoder {:?} not supported yet", decoder).into()), } From 600fca06e13ac56a0494e7589b88ff205ba57f60 Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Sun, 26 Jul 2026 22:16:30 +0200 Subject: [PATCH 14/17] wip(ai): fixes --- tokenizers/benches/ci_benchmark.rs | 2 +- .../tk-encode/src/decoders/byte_fallback.rs | 28 ++- tokenizers/tk-encode/src/decoders/strip.rs | 33 +++- tokenizers/tk-encode/src/models/bpe/model.rs | 13 ++ .../tk-encode/src/normalizers/replace.rs | 16 ++ .../tk-encode/src/pre_tokenizers/metaspace.rs | 45 ++++- .../tk-encode/src/tokenizer/pipeline.rs | 181 ++++++++++++++++-- .../src/vocab/bucket_added_vocabulary.rs | 8 + 8 files changed, 301 insertions(+), 25 deletions(-) diff --git a/tokenizers/benches/ci_benchmark.rs b/tokenizers/benches/ci_benchmark.rs index 35f6c881e..a4e8ce23f 100644 --- a/tokenizers/benches/ci_benchmark.rs +++ b/tokenizers/benches/ci_benchmark.rs @@ -267,7 +267,7 @@ fn bench_decode(c: &mut Criterion) { let mut sp_chain = Tokenizer::from_file("data/albert-base-v1-tokenizer.json").unwrap(); sp_chain.with_decoder(Some(Sequence::new(vec![ Replace::new("▁", " ").unwrap().into(), - ByteFallback::new().into(), + ByteFallback::default().into(), Fuse::new().into(), ]))); let lines = encode_lines(&sp_chain, &data); diff --git a/tokenizers/tk-encode/src/decoders/byte_fallback.rs b/tokenizers/tk-encode/src/decoders/byte_fallback.rs index b575da989..b9176860a 100644 --- a/tokenizers/tk-encode/src/decoders/byte_fallback.rs +++ b/tokenizers/tk-encode/src/decoders/byte_fallback.rs @@ -15,8 +15,10 @@ use serde::{Deserialize, Serialize}; pub struct ByteFallback { #[serde(rename = "type")] type_: MustBe!("ByteFallback"), - /// Lookup mapping a token id to the raw byte it represents + /// Lookup mapping a token id to the raw byte it represents. Built from + /// the model when the pipeline is assembled, never part of tokenizer.json. /// todo: closed-addressing, can use ptrhash + #[serde(skip)] fallback_lookup: AHashMap, } @@ -27,6 +29,18 @@ impl ByteFallback { fallback_lookup, } } + + /// Invert the model's encode-time byte -> id table into the id -> byte + /// lookup decoding needs. + pub(crate) fn from_byte_to_id(byte_to_id: &[u32; 256]) -> Self { + Self::new( + byte_to_id + .iter() + .enumerate() + .map(|(byte, &id)| (id, byte as u8)) + .collect(), + ) + } } impl Decoder for ByteFallback { @@ -88,8 +102,18 @@ impl pipeline::Decoder for ByteFallback { } fn flush(&self, state: &mut DecoderState, decoded: &mut Vec) -> Result<()> { - if state.pending_buffer.len() > 0 { + if state.pending_buffer.is_empty() { + return Ok(()); + } + if std::str::from_utf8(&state.pending_buffer).is_ok() { decoded.append(&mut state.pending_buffer); + } else { + // one '�' per byte token, like `decode_chain` above — not + // from_utf8_lossy, which merges maximal invalid subparts + for _ in 0..state.pending_buffer.len() { + decoded.extend_from_slice("�".as_bytes()); + } + state.pending_buffer.clear(); } Ok(()) } diff --git a/tokenizers/tk-encode/src/decoders/strip.rs b/tokenizers/tk-encode/src/decoders/strip.rs index 9aeffec64..e5eaf3c11 100644 --- a/tokenizers/tk-encode/src/decoders/strip.rs +++ b/tokenizers/tk-encode/src/decoders/strip.rs @@ -1,4 +1,7 @@ -use crate::tokenizer::{Decoder, Result}; +use crate::{ + pipeline, + tokenizer::{Decoder, Result}, +}; use serde::{Deserialize, Serialize}; @@ -59,6 +62,34 @@ impl Decoder for Strip { } } +impl pipeline::Decoder for Strip { + fn decode_token( + &self, + _state: &mut pipeline::DecoderState, + _token_id: u32, + token_bytes: &[u8], + decoded: &mut Vec, + ) -> Result<()> { + let mut pat_buf = [0u8; 4]; + let pat = self.content.encode_utf8(&mut pat_buf).as_bytes(); + let mut token = token_bytes; + for _ in 0..self.start { + match token.strip_prefix(pat) { + Some(rest) => token = rest, + None => break, + } + } + for _ in 0..self.stop { + match token.strip_suffix(pat) { + Some(rest) => token = rest, + None => break, + } + } + decoded.extend_from_slice(token); + Ok(()) + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/tokenizers/tk-encode/src/models/bpe/model.rs b/tokenizers/tk-encode/src/models/bpe/model.rs index 4cc9905c6..80277a429 100644 --- a/tokenizers/tk-encode/src/models/bpe/model.rs +++ b/tokenizers/tk-encode/src/models/bpe/model.rs @@ -763,6 +763,19 @@ impl PipelineBPE { }) } + /// The `<0xHH>` token ids indexed by byte, when this model encodes with + /// byte fallback. `Atoms::Bytes` also holds a byte -> id table, but it + /// maps byte-level atoms, not `<0xHH>` tokens, so it is not exposed here. + pub(crate) fn byte_fallback_ids(&self) -> Option<&[u32; 256]> { + match &self.atoms { + Atoms::Chars { + byte_fallback: Some(table), + .. + } => Some(table), + _ => None, + } + } + fn merge_word( &self, sequence: &str, diff --git a/tokenizers/tk-encode/src/normalizers/replace.rs b/tokenizers/tk-encode/src/normalizers/replace.rs index 48fc58180..24e5d710d 100644 --- a/tokenizers/tk-encode/src/normalizers/replace.rs +++ b/tokenizers/tk-encode/src/normalizers/replace.rs @@ -131,6 +131,22 @@ impl Decoder for Replace { } } +impl pipeline::Decoder for Replace { + fn decode_token( + &self, + state: &mut pipeline::DecoderState, + token_id: u32, + token_bytes: &[u8], + decoded: &mut Vec, + ) -> Result<()> { + match &self.pattern { + ReplacePattern::String(string) => {} + ReplacePattern::Regex(_) => {} + } + Ok(()) + } +} + // `Replace` needs a system-regex backend (SysRegex) for every test here. #[cfg(all(test, feature = "fancy-regex"))] mod tests { diff --git a/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs b/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs index 24aca962c..7b1c5cc0f 100644 --- a/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs +++ b/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs @@ -185,14 +185,20 @@ impl pipeline::Decoder for Metaspace { token_bytes: &[u8], decoded: &mut Vec, ) -> Result<()> { - if token_bytes.starts_with(self.str_rep.as_bytes()) { - if !replace(&mut state.started, true) && self.prepend_scheme != PrependScheme::Never { + // Every replacement char becomes ' ', except in the first token where + // prepend_scheme != Never drops it (it was prepended at encode time). + let first = !replace(&mut state.started, true); + let drop_replacement = first && self.prepend_scheme != PrependScheme::Never; + let pat = self.str_rep.as_bytes(); + let mut rest = token_bytes; + while let Some(pos) = rest.windows(pat.len()).position(|window| window == pat) { + decoded.extend_from_slice(&rest[..pos]); + if !drop_replacement { decoded.push(b' '); } - decoded.extend_from_slice(&token_bytes[self.replacement.len_utf8()..]); - } else { - decoded.extend_from_slice(token_bytes); + rest = &rest[pos + pat.len()..]; } + decoded.extend_from_slice(rest); Ok(()) } } @@ -204,6 +210,35 @@ mod tests { use super::*; use crate::{OffsetReferential, OffsetType}; + #[test] + fn pipeline_decode_token_matches_decode_chain() { + let tokens = ["▁Hey", "▁▁friend", "▁", "what", "▁▁"]; + for scheme in [ + PrependScheme::Always, + PrependScheme::First, + PrependScheme::Never, + ] { + let decoder = Metaspace::new('▁', scheme, true); + let expected = decoder + .decode_chain(tokens.iter().map(|t| t.to_string()).collect()) + .unwrap() + .concat(); + let mut state = DecoderState::default(); + let mut out = Vec::new(); + for token in tokens { + pipeline::Decoder::decode_token( + &decoder, + &mut state, + 0, + token.as_bytes(), + &mut out, + ) + .unwrap(); + } + assert_eq!(out, expected.as_bytes(), "prepend_scheme {scheme:?}"); + } + } + #[test] fn serialization() { let metaspace = Metaspace::new('_', PrependScheme::Always, true); diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index 4505f8713..7bc27182f 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -10,11 +10,13 @@ use atomsplit::classify::classify; use crate::DecoderWrapper; use crate::decoders::byte_fallback::ByteFallback; use crate::decoders::metaspace::Metaspace; +use crate::decoders::strip::Strip; use crate::decoders::wordpiece::WordPiece; use crate::models::bpe::{BpeScratch, PipelineBPE}; use crate::models::unigram::{Unigram, UnigramScratch}; use crate::models::wordlevel::WordLevel; use crate::models::wordpiece::{PipelineWordPiece, WordPieceScratch}; +use crate::normalizers::Replace; use crate::processors::bert::BertProcessing; use crate::processors::roberta::RobertaProcessing; use crate::utils::byte_level::GPT2_REGEX_STR; @@ -280,6 +282,8 @@ pub enum PipelineDecoder { WordPiece(WordPiece), MetaSpace(Metaspace), ByteFallback(ByteFallback), + Replace(Replace), + Strip(Strip), None, #[default] JoinWithSpaces, @@ -305,7 +309,9 @@ pub trait Decoder { let mut state = DecoderState::default(); for &token_id in ids { let token_bytes = match added_vocabulary.simple_id_to_token_bytes(token_id) { - Some(_) if skip_special_tokens => continue, + Some(_) if skip_special_tokens && added_vocabulary.is_special(token_id) => { + continue; + } Some(bytes) => bytes, None => model .id_to_token_bytes(token_id) @@ -350,6 +356,8 @@ impl Decoder for PipelineDecoder { decoded.extend_from_slice(token_bytes); Ok(()) } + Self::Strip(decoder) => decoder.decode_token(state, token_id, token_bytes, decoded), + Self::Replace(decoder) => decoder.decode_token(state, token_id, token_bytes, decoded), Self::WordPiece(decoder) => decoder.decode_token(state, token_id, token_bytes, decoded), Self::MetaSpace(decoder) => decoder.decode_token(state, token_id, token_bytes, decoded), Self::ByteFallback(decoder) => { @@ -365,7 +373,8 @@ impl Decoder for PipelineDecoder { if members.is_empty() { members.resize_with(stages.len(), Default::default); } - a.copy_from_slice(token_bytes); + a.clear(); + a.extend_from_slice(token_bytes); for (stage, member_state) in zip(stages, members) { b.clear(); stage.decode_token(member_state, token_id, a, b)?; @@ -394,7 +403,8 @@ impl Decoder for PipelineDecoder { if a.is_empty() { continue; } - for (stage, member_state) in stages.iter().zip(members.iter_mut()).skip(i) { + // stage i's flush is one token for the stages after it + for (stage, member_state) in stages.iter().zip(members.iter_mut()).skip(i + 1) { b.clear(); stage.decode_token(member_state, u32::MAX, a, b)?; if b.is_empty() { @@ -406,25 +416,40 @@ impl Decoder for PipelineDecoder { } Ok(()) } - decoder => decoder.flush(state, decoded), + Self::Strip(decoder) => decoder.flush(state, decoded), + Self::Replace(decoder) => decoder.flush(state, decoded), + Self::WordPiece(decoder) => decoder.flush(state, decoded), + Self::MetaSpace(decoder) => decoder.flush(state, decoded), + Self::ByteFallback(decoder) => decoder.flush(state, decoded), + Self::None | Self::JoinWithSpaces => Ok(()), } } } -impl TryFrom<&DecoderWrapper> for PipelineDecoder { - type Error = crate::Error; - - fn try_from(value: &DecoderWrapper) -> std::prelude::v1::Result { +impl PipelineDecoder { + /// Build from a legacy decoder. Takes the already-built `model` because + /// ByteFallback recognizes byte tokens by id, through the model's byte + /// fallback table. + fn from_decoder(value: &DecoderWrapper, model: &PipelineModel) -> Result { match value { // ByteLevel decoder is no longer needed as the vocabulary is stored as raw bytes DecoderWrapper::ByteLevel(_) => Ok(Self::None), DecoderWrapper::WordPiece(decoder) => Ok(Self::WordPiece(decoder.clone())), DecoderWrapper::Metaspace(decoder) => Ok(Self::MetaSpace(decoder.clone())), + DecoderWrapper::ByteFallback(_) => { + let byte_to_id = model.byte_fallback_ids().ok_or( + "ByteFallback decoder requires a model byte fallback table to map \ + token ids to bytes; only BPE models with `byte_fallback: true` build one", + )?; + Ok(Self::ByteFallback(ByteFallback::from_byte_to_id( + byte_to_id, + ))) + } DecoderWrapper::Sequence(sequence) => { let mut decoders = sequence .get_decoders() - .into_iter() - .map(Self::try_from) + .iter() + .map(|decoder| Self::from_decoder(decoder, model)) .filter(|maybe_decoder| { if let Ok(Self::None) = maybe_decoder { false @@ -666,6 +691,12 @@ impl TryFrom<&Tokenizer> for PipelineTokenizer { ModelWrapper::WordPiece(model) => PipelineModel::WordPiece(model.try_into()?), }; + let decoder = tok + .get_decoder() + .map(|decoder| PipelineDecoder::from_decoder(decoder, &model)) + .transpose()? + .unwrap_or_default(); + Ok(Self { added_vocabulary, normalizer: tok.get_normalizer().cloned(), @@ -676,11 +707,7 @@ impl TryFrom<&Tokenizer> for PipelineTokenizer { .map(PipelinePostProcessor::try_from) .transpose()? .unwrap_or_default(), - decoder: tok - .get_decoder() - .map(PipelineDecoder::try_from) - .transpose()? - .unwrap_or_default(), + decoder, }) } } @@ -733,7 +760,13 @@ impl PipelineTokenizer { ids, &mut output, )?; - Ok(String::from_utf8(output)?) + // A byte-level id sequence can end mid-character (a decode_stream + // prefix does constantly); the released ByteLevel decoder is lossy + // there, so match it instead of erroring. + Ok(match String::from_utf8(output) { + Ok(text) => text, + Err(err) => String::from_utf8_lossy(err.as_bytes()).into_owned(), + }) } /// Decode several id sequences at once, one `String` per input. Mirrors the @@ -1173,6 +1206,17 @@ impl Model for PipelineModel { } } +impl PipelineModel { + /// The model's encode-time byte fallback table (`<0xHH>` token ids indexed + /// by byte), when it has one. + fn byte_fallback_ids(&self) -> Option<&[u32; 256]> { + match self { + Self::BPE(model) => model.byte_fallback_ids(), + _ => None, + } + } +} + pub enum PipelineModelScratch { BPE(BpeScratch), WordLevel(()), @@ -1267,6 +1311,111 @@ mod tests { assert!(err.contains("not supported with model"), "{}", err); } + /// "h" and "e" plus all 256 `<0xHH>` tokens, each byte token's id being + /// its byte value. + fn byte_fallback_tokenizer(byte_fallback: bool) -> Tokenizer { + use crate::decoders::byte_fallback::ByteFallback; + use crate::models::bpe::{BpeBuilder, Vocab}; + + let mut vocab: Vocab = [("h".to_string(), 300), ("e".to_string(), 301)] + .into_iter() + .collect(); + vocab.extend((0..=255u8).map(|b| (format!("<0x{b:02X}>"), u32::from(b)))); + let bpe = BpeBuilder::default() + .vocab_and_merges(vocab, vec![]) + .byte_fallback(byte_fallback) + .build() + .unwrap(); + let mut tok = Tokenizer::new(bpe); + tok.with_decoder(Some(ByteFallback::default())); + tok + } + + #[test] + fn byte_fallback_decoder_decodes_byte_tokens_by_id() { + let pipeline = PipelineTokenizer::try_from(&byte_fallback_tokenizer(true)).unwrap(); + // 0xE5 0x8F 0xAB is the UTF-8 encoding of '叫'; the trailing byte run + // exercises the end-of-ids flush + assert_eq!( + pipeline + .decode(&[300, 301, 0xE5, 0x8F, 0xAB], false) + .unwrap(), + "he叫" + ); + } + + #[test] + fn conversion_rejects_byte_fallback_decoder_without_model_table() { + let err = conversion_error(&byte_fallback_tokenizer(false)); + assert!(err.contains("byte fallback table"), "{}", err); + } + + #[test] + fn byte_fallback_decoder_replaces_invalid_byte_runs() { + let pipeline = PipelineTokenizer::try_from(&byte_fallback_tokenizer(true)).unwrap(); + // 0xE5 0x8F alone is an incomplete UTF-8 sequence; legacy emits one + // '�' per byte token of the run + assert_eq!(pipeline.decode(&[0xE5, 0x8F, 300], false).unwrap(), "��h"); + assert_eq!(pipeline.decode(&[300, 0xE5, 0x8F], false).unwrap(), "h��"); + } + + #[test] + fn skip_special_tokens_keeps_non_special_added_tokens() { + let mut tok = byte_fallback_tokenizer(true); + tok.add_tokens([crate::AddedToken::from("", false)]) + .unwrap(); + tok.add_special_tokens([crate::AddedToken::from("", true)]) + .unwrap(); + let think = tok.token_to_id("").unwrap(); + let end = tok.token_to_id("").unwrap(); + let pipeline = PipelineTokenizer::try_from(&tok).unwrap(); + assert_eq!( + pipeline.decode(&[think, end, 300], false).unwrap(), + "h" + ); + assert_eq!( + pipeline.decode(&[think, end, 300], true).unwrap(), + "h" + ); + } + + #[test] + fn sequence_decoder_chains_stages_per_token() { + let decoder = PipelineDecoder::Sequence(vec![ + PipelineDecoder::Strip(Strip::new('#', 1, 0)), + PipelineDecoder::JoinWithSpaces, + ]); + let mut state = DecoderState::default(); + let mut out = Vec::new(); + for token in [b"#hey".as_slice(), b"#you"] { + decoder + .decode_token(&mut state, 0, token, &mut out) + .unwrap(); + } + decoder.flush(&mut state, &mut out).unwrap(); + assert_eq!(out, b"hey you"); + } + + #[test] + fn sequence_flushes_held_byte_run_through_later_stages() { + let mut lookup = ahash::AHashMap::new(); + lookup.insert(7u32, b'a'); + let decoder = PipelineDecoder::Sequence(vec![ + PipelineDecoder::ByteFallback(ByteFallback::new(lookup)), + PipelineDecoder::JoinWithSpaces, + ]); + let mut state = DecoderState::default(); + let mut out = Vec::new(); + decoder + .decode_token(&mut state, 1, b"hey", &mut out) + .unwrap(); + decoder + .decode_token(&mut state, 7, b"<0x61>", &mut out) + .unwrap(); + decoder.flush(&mut state, &mut out).unwrap(); + assert_eq!(out, b"hey a"); + } + fn wordlevel_tokenizer( vocab: Vec<(&str, u32)>, post_processor: Option, diff --git a/tokenizers/tk-encode/src/vocab/bucket_added_vocabulary.rs b/tokenizers/tk-encode/src/vocab/bucket_added_vocabulary.rs index 6ff84fbb7..19f80ae2b 100644 --- a/tokenizers/tk-encode/src/vocab/bucket_added_vocabulary.rs +++ b/tokenizers/tk-encode/src/vocab/bucket_added_vocabulary.rs @@ -262,6 +262,14 @@ impl AddedVocabulary { .or_else(|| self.normalized_vocab.id_to_token_bytes(id)) } + /// Whether `id` is a *special* added token — the kind `skip_special_tokens` + /// drops at decode time, as opposed to user-added regular tokens. + pub fn is_special(&self, id: u32) -> bool { + self.token_metadata + .get(id as usize) + .is_some_and(|metadata| metadata.special) + } + // pub fn set_encode_special_tokens(&mut self, value: bool) { self.encode_special_tokens = value; From e913e51e8a862f8c1f02ccbb4ce8c1bc8678c2a4 Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Sun, 26 Jul 2026 22:21:33 +0200 Subject: [PATCH 15/17] replace impl --- .../tk-encode/src/normalizers/replace.rs | 67 +++++++++++++++++-- .../tk-encode/src/tokenizer/pipeline.rs | 18 +++++ 2 files changed, 80 insertions(+), 5 deletions(-) diff --git a/tokenizers/tk-encode/src/normalizers/replace.rs b/tokenizers/tk-encode/src/normalizers/replace.rs index 24e5d710d..1b8df983b 100644 --- a/tokenizers/tk-encode/src/normalizers/replace.rs +++ b/tokenizers/tk-encode/src/normalizers/replace.rs @@ -134,16 +134,39 @@ impl Decoder for Replace { impl pipeline::Decoder for Replace { fn decode_token( &self, - state: &mut pipeline::DecoderState, - token_id: u32, + _state: &mut pipeline::DecoderState, + _token_id: u32, token_bytes: &[u8], decoded: &mut Vec, ) -> Result<()> { match &self.pattern { - ReplacePattern::String(string) => {} - ReplacePattern::Regex(_) => {} + // Plain byte search: no regex machinery on the hot path, and raw + // (non-UTF-8) token bytes pass through unharmed. + ReplacePattern::String(pattern) if !pattern.is_empty() => { + let pat = pattern.as_bytes(); + let mut rest = token_bytes; + while let Some(pos) = rest.windows(pat.len()).position(|window| window == pat) { + decoded.extend_from_slice(&rest[..pos]); + decoded.extend_from_slice(self.content.as_bytes()); + rest = &rest[pos + pat.len()..]; + } + decoded.extend_from_slice(rest); + Ok(()) + } + _ => { + let token = std::str::from_utf8(token_bytes).map_err( + |_| "Replace decoder with a regex pattern requires valid UTF-8 tokens", + )?; + let mut last_end = 0; + for (start, end) in self.regex.find_iter(token) { + decoded.extend_from_slice(&token.as_bytes()[last_end..start]); + decoded.extend_from_slice(self.content.as_bytes()); + last_end = end; + } + decoded.extend_from_slice(&token.as_bytes()[last_end..]); + Ok(()) + } } - Ok(()) } } @@ -200,6 +223,40 @@ mod tests { ); } + #[test] + fn pipeline_decode_token_matches_decode_chain() { + let cases = vec![ + ( + Replace::new("▁", " ").unwrap(), + vec!["▁Hey", "▁▁friend", "no_meta", "▁", ""], + ), + (Replace::new("ab", "X").unwrap(), vec!["aabb", "abab", "b"]), + ( + Replace::new(ReplacePattern::Regex(r"\s+".into()), " ").unwrap(), + vec!["a b", " x ", "y"], + ), + ]; + for (replace, tokens) in cases { + let expected = replace + .decode_chain(tokens.iter().map(|t| t.to_string()).collect()) + .unwrap() + .concat(); + let mut state = pipeline::DecoderState::default(); + let mut out = Vec::new(); + for token in &tokens { + pipeline::Decoder::decode_token( + &replace, + &mut state, + 0, + token.as_bytes(), + &mut out, + ) + .unwrap(); + } + assert_eq!(out, expected.as_bytes(), "pattern {:?}", replace.pattern); + } + } + #[test] fn pipeline_replace_matches_legacy() { let n = Replace::new("''", "\"").unwrap(); diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index 7bc27182f..b78f48905 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -436,6 +436,7 @@ impl PipelineDecoder { DecoderWrapper::ByteLevel(_) => Ok(Self::None), DecoderWrapper::WordPiece(decoder) => Ok(Self::WordPiece(decoder.clone())), DecoderWrapper::Metaspace(decoder) => Ok(Self::MetaSpace(decoder.clone())), + DecoderWrapper::Replace(decoder) => Ok(Self::Replace(decoder.clone())), DecoderWrapper::ByteFallback(_) => { let byte_to_id = model.byte_fallback_ids().ok_or( "ByteFallback decoder requires a model byte fallback table to map \ @@ -1396,6 +1397,23 @@ mod tests { assert_eq!(out, b"hey you"); } + // llama-2's chain in miniature: Replace rewrites plain tokens, ByteFallback + // holds byte runs, and the run flushes through the rest of the chain. + #[cfg(feature = "fancy-regex")] + #[test] + fn sequence_replace_then_byte_fallback_decodes() { + let mut tok = byte_fallback_tokenizer(true); + tok.with_decoder(Some(crate::decoders::sequence::Sequence::new(vec![ + DecoderWrapper::Replace(crate::normalizers::Replace::new("h", "H").unwrap()), + DecoderWrapper::ByteFallback(crate::decoders::byte_fallback::ByteFallback::default()), + ]))); + let pipeline = PipelineTokenizer::try_from(&tok).unwrap(); + assert_eq!( + pipeline.decode(&[300, 0xE5, 0x8F, 0xAB], false).unwrap(), + "H叫" + ); + } + #[test] fn sequence_flushes_held_byte_run_through_later_stages() { let mut lookup = ahash::AHashMap::new(); From c8b55d306011863e818237357eb3eb2966b7a022 Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Sun, 26 Jul 2026 22:24:36 +0200 Subject: [PATCH 16/17] fix bindings --- bindings/node/src/decoders.rs | 3 ++- bindings/python/src/decoders.rs | 3 ++- tokenizers/tk-encode/src/tokenizer/pipeline.rs | 12 +++--------- 3 files changed, 7 insertions(+), 11 deletions(-) diff --git a/bindings/node/src/decoders.rs b/bindings/node/src/decoders.rs index 51126c0dc..afccf3ef2 100644 --- a/bindings/node/src/decoders.rs +++ b/bindings/node/src/decoders.rs @@ -1,4 +1,5 @@ use crate::arc_rwlock_serde; +use ahash::AHashMap; use serde::{Deserialize, Serialize}; extern crate tokenizers as tk; use napi::bindgen_prelude::*; @@ -58,7 +59,7 @@ pub fn bpe_decoder(suffix: Option) -> Decoder { pub fn byte_fallback_decoder() -> Decoder { Decoder { decoder: Some(Arc::new(RwLock::new( - tk::decoders::byte_fallback::ByteFallback::new().into(), + tk::decoders::byte_fallback::ByteFallback::new(AHashMap::new()).into(), ))), } } diff --git a/bindings/python/src/decoders.rs b/bindings/python/src/decoders.rs index 46dbab4b0..2d825efb4 100644 --- a/bindings/python/src/decoders.rs +++ b/bindings/python/src/decoders.rs @@ -3,6 +3,7 @@ use std::sync::{Arc, RwLock}; use crate::pre_tokenizers::from_string; use crate::tokenizer::PyTokenizer; use crate::utils::PyPattern; +use ahash::AHashMap; use pyo3::exceptions; use pyo3::prelude::*; use pyo3::types::*; @@ -299,7 +300,7 @@ impl PyByteFallbackDec { #[new] #[pyo3(signature = (), text_signature = "(self)")] fn new() -> PyClassInitializer { - PyClassInitializer::::from(PyDecoder::from(ByteFallback::new())) + PyClassInitializer::::from(PyDecoder::from(ByteFallback::new(AHashMap::new()))) .add_subclass(PyByteFallbackDec {}) } } diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index b78f48905..5678eae31 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -332,7 +332,7 @@ pub trait Decoder { ) -> Result<()>; fn flush(&self, _state: &mut DecoderState, _decoded: &mut Vec) -> Result<()> { - return Ok(()); + Ok(()) } } @@ -451,15 +451,9 @@ impl PipelineDecoder { .get_decoders() .iter() .map(|decoder| Self::from_decoder(decoder, model)) - .filter(|maybe_decoder| { - if let Ok(Self::None) = maybe_decoder { - false - } else { - true - } - }) + .filter(|maybe_decoder| !matches!(maybe_decoder, Ok(Self::None))) .collect::>>()?; - if decoders.len() == 0 { + if decoders.is_empty() { return Ok(Self::default()); } if decoders.len() == 1 { From 5845e8eb41cdc293bc21e3365a2e087ffaebf80c Mon Sep 17 00:00:00 2001 From: SBrandeis <33657802+SBrandeis@users.noreply.github.com> Date: Mon, 27 Jul 2026 12:12:47 +0200 Subject: [PATCH 17/17] fix: llama-2 --- .../tk-encode/src/tokenizer/pipeline.rs | 142 +++++++++++++++--- .../src/vocab/bucket_added_vocabulary.rs | 14 +- tokenizers/tk-encode/src/vocab/buckets.rs | 5 + 3 files changed, 132 insertions(+), 29 deletions(-) diff --git a/tokenizers/tk-encode/src/tokenizer/pipeline.rs b/tokenizers/tk-encode/src/tokenizer/pipeline.rs index 5678eae31..1e738c296 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -2,7 +2,6 @@ use std::cell::RefCell; use std::convert::TryInto; use std::iter::zip; use std::mem::{replace, swap}; -use std::u32; use std::{borrow::Cow, convert::TryFrom}; use atomsplit::classify::classify; @@ -283,7 +282,13 @@ pub enum PipelineDecoder { MetaSpace(Metaspace), ByteFallback(ByteFallback), Replace(Replace), - Strip(Strip), + /// A `Strip` with no `Fuse` before it: strips each token. + StripToken(Strip), + /// A `Strip` after a `Fuse` (llama-2 style): `Fuse` concatenates all + /// tokens into one, so the strip applies once to the whole decoded output, + /// in [`decode_whole`](Decoder::decode_whole). `decode_token` passes bytes + /// through untouched. + StripWhole(Strip), None, #[default] JoinWithSpaces, @@ -309,7 +314,7 @@ pub trait Decoder { let mut state = DecoderState::default(); for &token_id in ids { let token_bytes = match added_vocabulary.simple_id_to_token_bytes(token_id) { - Some(_) if skip_special_tokens && added_vocabulary.is_special(token_id) => { + Some(_) if skip_special_tokens && added_vocabulary.skip_on_decode(token_id) => { continue; } Some(bytes) => bytes, @@ -320,6 +325,7 @@ pub trait Decoder { self.decode_token(&mut state, token_id, token_bytes, decoded)?; } self.flush(&mut state, decoded)?; + self.decode_whole(&mut state, decoded)?; Ok(()) } @@ -334,6 +340,13 @@ pub trait Decoder { fn flush(&self, _state: &mut DecoderState, _decoded: &mut Vec) -> Result<()> { Ok(()) } + + /// Runs once at the end of `decode`, on the fully decoded output. This is + /// where stages placed after a `Fuse` apply: `Fuse` concatenates all + /// tokens into one, so later stages see the whole text as one token. + fn decode_whole(&self, _state: &mut DecoderState, _decoded: &mut Vec) -> Result<()> { + Ok(()) + } } impl Decoder for PipelineDecoder { @@ -345,7 +358,7 @@ impl Decoder for PipelineDecoder { decoded: &mut Vec, ) -> Result<()> { match self { - Self::None => { + Self::None | Self::StripWhole(_) => { decoded.extend_from_slice(token_bytes); Ok(()) } @@ -356,7 +369,9 @@ impl Decoder for PipelineDecoder { decoded.extend_from_slice(token_bytes); Ok(()) } - Self::Strip(decoder) => decoder.decode_token(state, token_id, token_bytes, decoded), + Self::StripToken(decoder) => { + decoder.decode_token(state, token_id, token_bytes, decoded) + } Self::Replace(decoder) => decoder.decode_token(state, token_id, token_bytes, decoded), Self::WordPiece(decoder) => decoder.decode_token(state, token_id, token_bytes, decoded), Self::MetaSpace(decoder) => decoder.decode_token(state, token_id, token_bytes, decoded), @@ -397,6 +412,9 @@ impl Decoder for PipelineDecoder { members, .. } = state; + if members.is_empty() { + members.resize_with(stages.len(), Default::default); + } for i in 0..stages.len() { a.clear(); stages[i].flush(&mut members[i], a)?; @@ -416,12 +434,31 @@ impl Decoder for PipelineDecoder { } Ok(()) } - Self::Strip(decoder) => decoder.flush(state, decoded), + Self::StripToken(decoder) => decoder.flush(state, decoded), Self::Replace(decoder) => decoder.flush(state, decoded), Self::WordPiece(decoder) => decoder.flush(state, decoded), Self::MetaSpace(decoder) => decoder.flush(state, decoded), Self::ByteFallback(decoder) => decoder.flush(state, decoded), - Self::None | Self::JoinWithSpaces => Ok(()), + Self::None | Self::StripWhole(_) | Self::JoinWithSpaces => Ok(()), + } + } + + fn decode_whole(&self, state: &mut DecoderState, decoded: &mut Vec) -> Result<()> { + match self { + Self::Sequence(stages) => { + for stage in stages { + stage.decode_whole(state, decoded)?; + } + Ok(()) + } + Self::StripWhole(strip) => { + let [scratch, _] = &mut state.swap_buffers; + scratch.clear(); + strip.decode_token(&mut DecoderState::default(), u32::MAX, decoded, scratch)?; + swap(decoded, scratch); + Ok(()) + } + _ => Ok(()), } } } @@ -432,8 +469,10 @@ impl PipelineDecoder { /// fallback table. fn from_decoder(value: &DecoderWrapper, model: &PipelineModel) -> Result { match value { - // ByteLevel decoder is no longer needed as the vocabulary is stored as raw bytes - DecoderWrapper::ByteLevel(_) => Ok(Self::None), + // ByteLevel is not needed as the vocabulary is stored as raw bytes; + // a standalone Fuse only concatenates, which decoding does anyway + DecoderWrapper::ByteLevel(_) | DecoderWrapper::Fuse(_) => Ok(Self::None), + DecoderWrapper::Strip(decoder) => Ok(Self::StripToken(decoder.clone())), DecoderWrapper::WordPiece(decoder) => Ok(Self::WordPiece(decoder.clone())), DecoderWrapper::Metaspace(decoder) => Ok(Self::MetaSpace(decoder.clone())), DecoderWrapper::Replace(decoder) => Ok(Self::Replace(decoder.clone())), @@ -447,21 +486,35 @@ impl PipelineDecoder { ))) } DecoderWrapper::Sequence(sequence) => { - let mut decoders = sequence - .get_decoders() - .iter() - .map(|decoder| Self::from_decoder(decoder, model)) - .filter(|maybe_decoder| !matches!(maybe_decoder, Ok(Self::None))) - .collect::>>()?; - if decoders.is_empty() { - return Ok(Self::default()); + // Fuse concatenates all tokens into one, so a Strip after it + // strips the whole decoded output, not each token. No other + // decoder is supported after a Fuse yet. + let mut fused = false; + let mut decoders = Vec::new(); + for stage in sequence.get_decoders() { + let decoder = match stage { + DecoderWrapper::Fuse(_) => { + fused = true; + continue; + } + DecoderWrapper::Strip(strip) if fused => Self::StripWhole(strip.clone()), + _ if fused => { + return Err(format!( + "only Strip decoders may follow Fuse, got {stage:?}" + ) + .into()); + } + _ => Self::from_decoder(stage, model)?, + }; + if !matches!(decoder, Self::None) { + decoders.push(decoder); + } } - if decoders.len() == 1 { - // SAFETY: safe to .unwrap() as the len is asserted to be 1 - let first = decoders.pop().unwrap(); - return Ok(first); + match decoders.len() { + 0 => Ok(Self::None), + 1 => Ok(decoders.pop().unwrap()), + _ => Ok(Self::Sequence(decoders)), } - Ok(Self::Sequence(decoders)) } decoder => Err(format!("Decoder {:?} not supported yet", decoder).into()), } @@ -609,7 +662,7 @@ impl TryFrom<&Tokenizer> for PipelineTokenizer { let pre_tokenizer: PipelinePreTokenizer = tok .get_pre_tokenizer() .cloned() - .map(TryInto::try_into) + .map(PipelinePreTokenizer::try_from) .transpose()? .unwrap_or(PipelinePreTokenizer::None); @@ -1374,10 +1427,26 @@ mod tests { ); } + #[test] + fn skip_special_tokens_keeps_specials_rewritten_by_normalization() { + use crate::normalizers::prepend::Prepend; + + let mut tok = byte_fallback_tokenizer(true); + tok.with_normalizer(Some(Prepend::new("▁".into()))).unwrap(); + tok.add_special_tokens([crate::AddedToken::from("", true).normalized(true)]) + .unwrap(); + let end = tok.token_to_id("").unwrap(); + let pipeline = PipelineTokenizer::try_from(&tok).unwrap(); + // The released crate compares the decode-time (normalized) string + // against raw special contents: "▁" never matches "", so the + // special survives the skip — llama-2's `` under `Prepend("▁")`. + assert_eq!(pipeline.decode(&[end, 300], true).unwrap(), "▁h"); + } + #[test] fn sequence_decoder_chains_stages_per_token() { let decoder = PipelineDecoder::Sequence(vec![ - PipelineDecoder::Strip(Strip::new('#', 1, 0)), + PipelineDecoder::StripToken(Strip::new('#', 1, 0)), PipelineDecoder::JoinWithSpaces, ]); let mut state = DecoderState::default(); @@ -1428,6 +1497,31 @@ mod tests { assert_eq!(out, b"hey a"); } + #[test] + fn stages_after_fuse_decode_the_whole_output() { + let mut tok = byte_fallback_tokenizer(true); + // llama-2's trailing Strip(' ', 1, 0) comes after Fuse: it must remove + // one leading space from the whole text, not one from every token + tok.with_decoder(Some(crate::decoders::sequence::Sequence::new(vec![ + DecoderWrapper::ByteFallback(crate::decoders::byte_fallback::ByteFallback::default()), + DecoderWrapper::Fuse(crate::decoders::fuse::Fuse::new()), + DecoderWrapper::Strip(Strip::new(' ', 1, 0)), + ]))); + let pipeline = PipelineTokenizer::try_from(&tok).unwrap(); + assert_eq!(pipeline.decode(&[32, 300, 32, 301], false).unwrap(), "h e"); + } + + #[test] + fn conversion_rejects_fuse_followed_by_non_strip() { + let mut tok = byte_fallback_tokenizer(true); + tok.with_decoder(Some(crate::decoders::sequence::Sequence::new(vec![ + DecoderWrapper::Fuse(crate::decoders::fuse::Fuse::new()), + DecoderWrapper::WordPiece(crate::decoders::wordpiece::WordPiece::default()), + ]))); + let err = conversion_error(&tok); + assert!(err.contains("follow Fuse"), "{}", err); + } + fn wordlevel_tokenizer( vocab: Vec<(&str, u32)>, post_processor: Option, diff --git a/tokenizers/tk-encode/src/vocab/bucket_added_vocabulary.rs b/tokenizers/tk-encode/src/vocab/bucket_added_vocabulary.rs index 19f80ae2b..f43887124 100644 --- a/tokenizers/tk-encode/src/vocab/bucket_added_vocabulary.rs +++ b/tokenizers/tk-encode/src/vocab/bucket_added_vocabulary.rs @@ -134,6 +134,7 @@ impl From<&AddedToken> for AddedTokenFlags { single_word: token.single_word, lstrip: token.lstrip, rstrip: token.rstrip, + skip_on_decode: false, } } } @@ -262,12 +263,12 @@ impl AddedVocabulary { .or_else(|| self.normalized_vocab.id_to_token_bytes(id)) } - /// Whether `id` is a *special* added token — the kind `skip_special_tokens` - /// drops at decode time, as opposed to user-added regular tokens. - pub fn is_special(&self, id: u32) -> bool { + /// Whether `skip_special_tokens` drops `id` at decode time — see + /// [`AddedTokenFlags::skip_on_decode`] for why this is not `special`. + pub fn skip_on_decode(&self, id: u32) -> bool { self.token_metadata .get(id as usize) - .is_some_and(|metadata| metadata.special) + .is_some_and(|metadata| metadata.skip_on_decode) } // @@ -332,7 +333,7 @@ impl AddedVocabulary { ignored += 1; continue; } - let flags = AddedTokenFlags::from(&token); + let mut flags = AddedTokenFlags::from(&token); let is_norm = flags.normalized; let norm_form: String = match normalizer { Some(n) => { @@ -343,6 +344,9 @@ impl AddedVocabulary { } None => token.content.clone(), }; + // a special is skipped only if normalization left its content + // unchanged — see `AddedTokenFlags::skip_on_decode` + flags.skip_on_decode = flags.special && (!is_norm || norm_form == token.content); let form = if is_norm { norm_form.clone().into_bytes() } else { diff --git a/tokenizers/tk-encode/src/vocab/buckets.rs b/tokenizers/tk-encode/src/vocab/buckets.rs index 4db9e0202..34f814446 100644 --- a/tokenizers/tk-encode/src/vocab/buckets.rs +++ b/tokenizers/tk-encode/src/vocab/buckets.rs @@ -7,6 +7,11 @@ pub struct AddedTokenFlags { pub single_word: bool, pub lstrip: bool, pub rstrip: bool, + /// Whether `skip_special_tokens` drops this token. Not the same as + /// `special`: `Tokenizer::decode` compares the decode-time (normalized) + /// string against raw special contents, so a special whose content the + /// normalizer rewrote (llama-2's ``, stored as `▁`) survives the skip. + pub skip_on_decode: bool, } /// The key to have a fast byte matching alrogithm is to skip failures fast and reject quickly