diff --git a/bindings/node/src/decoders.rs b/bindings/node/src/decoders.rs index 51126c0dcf..afccf3ef23 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 46dbab4b0d..2d825efb42 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/benches/ci_benchmark.rs b/tokenizers/benches/ci_benchmark.rs index 35f6c881e1..a4e8ce23fb 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/examples/fixture_bench.rs b/tokenizers/tk-encode/examples/fixture_bench.rs index 9ac6d89bf3..ccb3191caa 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", diff --git a/tokenizers/tk-encode/src/decoders/byte_fallback.rs b/tokenizers/tk-encode/src/decoders/byte_fallback.rs index 57b7b63cd7..b9176860a7 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,14 +15,32 @@ use serde::{Deserialize, Serialize}; pub struct ByteFallback { #[serde(rename = "type")] type_: MustBe!("ByteFallback"), + /// 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, } impl ByteFallback { - pub fn new() -> Self { + pub fn new(fallback_lookup: AHashMap) -> Self { Self { type_: MustBe!("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 { @@ -62,13 +84,48 @@ 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.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(()) + } +} + #[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/strip.rs b/tokenizers/tk-encode/src/decoders/strip.rs index 9aeffec647..e5eaf3c11d 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/decoders/wordpiece.rs b/tokenizers/tk-encode/src/decoders/wordpiece.rs index a2da414c0a..05f9efb61a 100644 --- a/tokenizers/tk-encode/src/decoders/wordpiece.rs +++ b/tokenizers/tk-encode/src/decoders/wordpiece.rs @@ -1,4 +1,9 @@ -use crate::tokenizer::{Decoder, Result}; +use std::mem::replace; + +use crate::{ + pipeline::{self, DecoderState}, + tokenizer::{Decoder, Result}, +}; use serde::{Deserialize, Serialize}; @@ -28,6 +33,7 @@ impl Default for WordPiece { } } } + pub fn cleanup(dirty_input: &str) -> String { dirty_input .replace(" .", ".") @@ -61,6 +67,32 @@ impl Decoder for WordPiece { } } +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 { + fn decode_token( + &self, + state: &mut DecoderState, + _token_id: u32, + token_bytes: &[u8], + decoded: &mut Vec, + ) -> Result<()> { + 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 { + decoded.push(b' '); + decoded.extend(token_bytes); + } + 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 fa7bcec23f..80277a4294 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, @@ -855,6 +868,10 @@ impl pipeline::Model for PipelineBPE { skip: Vec::new(), } } + + fn id_to_token_bytes(&self, id: u32) -> Option<&[u8]> { + self.vocab.id_to_token_bytes(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 1cc3f4b987..b0d9daa512 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: u32) -> Option<&[u8]> { + self.token_to_ids.id_to_token_bytes(id) + } } #[cfg(test)] diff --git a/tokenizers/tk-encode/src/models/wordlevel/mod.rs b/tokenizers/tk-encode/src/models/wordlevel/mod.rs index e26f16387a..03e73b3b5d 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: u32) -> Option<&[u8]> { + self.vocab_r.get(&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 a1286fe953..c8aa7be665 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: 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]) + } } #[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 _; + + 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(id).unwrap(); + assert_eq!(bytes, token.as_bytes()); + } + assert_eq!(model.id_to_token_bytes(3), None); + } } diff --git a/tokenizers/tk-encode/src/normalizers/replace.rs b/tokenizers/tk-encode/src/normalizers/replace.rs index 48fc581801..1b8df983b9 100644 --- a/tokenizers/tk-encode/src/normalizers/replace.rs +++ b/tokenizers/tk-encode/src/normalizers/replace.rs @@ -131,6 +131,45 @@ 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 { + // 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(()) + } + } + } +} + // `Replace` needs a system-regex backend (SysRegex) for every test here. #[cfg(all(test, feature = "fancy-regex"))] mod tests { @@ -184,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/pre_tokenizers/metaspace.rs b/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs index 2d7b7b0aad..7b1c5cc0f0 100644 --- a/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs +++ b/tokenizers/tk-encode/src/pre_tokenizers/metaspace.rs @@ -1,4 +1,9 @@ -use crate::tokenizer::{Decoder, PreTokenizedString, PreTokenizer, Result, SplitDelimiterBehavior}; +use std::mem::replace; + +use crate::{ + pipeline::{self, DecoderState}, + tokenizer::{Decoder, PreTokenizedString, PreTokenizer, Result, SplitDelimiterBehavior}, +}; use serde::{Deserialize, Deserializer, Serialize, de}; /// Enum representing options for the metaspace prepending scheme. @@ -172,6 +177,32 @@ impl Decoder for Metaspace { } } +impl pipeline::Decoder for Metaspace { + fn decode_token( + &self, + state: &mut DecoderState, + _token_id: u32, + token_bytes: &[u8], + decoded: &mut Vec, + ) -> Result<()> { + // 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' '); + } + rest = &rest[pos + pat.len()..]; + } + decoded.extend_from_slice(rest); + Ok(()) + } +} + #[cfg(test)] mod tests { use regex::Regex; @@ -179,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 efb12eb679..1e738c296f 100644 --- a/tokenizers/tk-encode/src/tokenizer/pipeline.rs +++ b/tokenizers/tk-encode/src/tokenizer/pipeline.rs @@ -1,13 +1,21 @@ use std::cell::RefCell; use std::convert::TryInto; +use std::iter::zip; +use std::mem::{replace, swap}; 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::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; @@ -266,6 +274,253 @@ 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), + ByteFallback(ByteFallback), + Replace(Replace), + /// 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, +} + +#[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, + model: &PipelineModel, + added_vocabulary: &BucketAddedVocabulary, + skip_special_tokens: bool, + ids: &[u32], + decoded: &mut Vec, + ) -> crate::Result<()> { + 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.skip_on_decode(token_id) => { + continue; + } + Some(bytes) => bytes, + None => model + .id_to_token_bytes(token_id) + .ok_or(format!("Invalid token id: {token_id}"))?, + }; + self.decode_token(&mut state, token_id, token_bytes, decoded)?; + } + self.flush(&mut state, decoded)?; + self.decode_whole(&mut state, decoded)?; + Ok(()) + } + + fn decode_token( + &self, + state: &mut DecoderState, + token_id: u32, + token_bytes: &[u8], + decoded: &mut Vec, + ) -> Result<()>; + + 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 { + fn decode_token( + &self, + state: &mut DecoderState, + token_id: u32, + token_bytes: &[u8], + decoded: &mut Vec, + ) -> Result<()> { + match self { + Self::None | Self::StripWhole(_) => { + decoded.extend_from_slice(token_bytes); + Ok(()) + } + Self::JoinWithSpaces => { + if replace(&mut state.started, true) { + decoded.push(b' '); + } + decoded.extend_from_slice(token_bytes); + Ok(()) + } + 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), + 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.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)?; + 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; + 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)?; + if a.is_empty() { + continue; + } + // 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() { + return Ok(()); + } + std::mem::swap(a, b); + } + decoded.extend_from_slice(a); + } + Ok(()) + } + 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::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(()), + } + } +} + +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 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())), + 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) => { + // 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); + } + } + match decoders.len() { + 0 => Ok(Self::None), + 1 => Ok(decoders.pop().unwrap()), + _ => Ok(Self::Sequence(decoders)), + } + } + 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 +646,7 @@ pub struct PipelineTokenizer { pre_tokenizer: PipelinePreTokenizer, model: PipelineModel, post_processor: PipelinePostProcessor, + decoder: PipelineDecoder, } impl TryFrom<&Tokenizer> for PipelineTokenizer { @@ -406,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); @@ -483,6 +739,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(), @@ -493,6 +755,7 @@ impl TryFrom<&Tokenizer> for PipelineTokenizer { .map(PipelinePostProcessor::try_from) .transpose()? .unwrap_or_default(), + decoder, }) } } @@ -536,14 +799,22 @@ 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: &[u32], _skip_special_tokens: bool) -> Result { - Err("PipelineTokenizer::decode is not implemented yet".into()) + pub fn decode(&self, ids: &[u32], skip_special_tokens: bool) -> Result { + let mut output = Vec::with_capacity(4 * ids.len()); + self.decoder.decode( + &self.model, + &self.added_vocabulary, + skip_special_tokens, + ids, + &mut 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 @@ -923,6 +1194,8 @@ pub trait Model { ) -> Result<()>; fn init_scratch(&self) -> Self::Scratch; + + fn id_to_token_bytes(&self, id: u32) -> Option<&[u8]>; } #[allow( @@ -970,6 +1243,26 @@ impl Model for PipelineModel { Self::Unigram(unigram) => Self::Scratch::Unigram(unigram.init_scratch()), } } + + 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), + Self::WordPiece(model) => model.id_to_token_bytes(id), + Self::Unigram(model) => model.id_to_token_bytes(id), + } + } +} + +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 { @@ -1066,6 +1359,169 @@ 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 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::StripToken(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"); + } + + // 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(); + 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"); + } + + #[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 5bc0be6ddc..f43887124b 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, } } } @@ -249,10 +250,25 @@ 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)) + } + + /// 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.skip_on_decode) } // @@ -317,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) => { @@ -328,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 9a78a8cec4..34f8144469 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 @@ -443,6 +448,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 { diff --git a/tokenizers/tk-encode/tests/pipeline_decode_oracle.rs b/tokenizers/tk-encode/tests/pipeline_decode_oracle.rs index 1078211609..db0f4442f3 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.