use serde_json::Value;
use splintr::pretrained::from_pretrained;
use splintr::{AnyTokenizer, Backend, PretrainedVocab, SpecialDecode, Tokenize, TokenizeError};
use std::fs;
use std::path::Path;
const EXTRA_TEXTS: [&str; 18] = [
"日本語のテキスト",
"한국어와中文混合",
"Здравствуй, мир",
"café résumé naïve",
"😀🎉🚀",
"👨👩👧👦",
"🏳️🌈",
"emoji😀letterRun日本語",
"𐍈𐌰𐌱",
"🜁🜂🜃🜄",
"\u{10FFFF}",
"\u{FFFD}",
"",
" ",
"\t\n\r\n",
" \t\n ",
"config.getValue()#hashtagWord",
"a\tb\nc d",
];
const SPECIAL_SPELLINGS: [&str; 8] = [
"<|endoftext|>",
"<|end_of_text|>",
"</s>",
"<s>",
"[INST]",
"<|im_start|>",
"<|pad|>",
"<|think|>",
];
const SPECIAL_NEIGHBOURS: [(&str, &str); 3] = [("abc", "def"), ("", "日本語"), ("🎉", "")];
fn escape(text: &str) -> String {
let mut out = String::with_capacity(text.len());
for ch in text.chars() {
match ch {
'\n' => out.push_str("\\n"),
'\r' => out.push_str("\\r"),
'\t' => out.push_str("\\t"),
'"' => out.push_str("\\\""),
'\\' => out.push_str("\\\\"),
' ' => out.push('\u{2423}'),
c if c.is_control() => out.push_str(&format!("\\u{{{:04x}}}", c as u32)),
c => out.push(c),
}
}
out
}
fn format_ids(ids: &[u32]) -> String {
ids.iter().map(u32::to_string).collect::<Vec<_>>().join(" ")
}
fn fixture_inputs() -> Vec<String> {
let fixtures_dir = Path::new(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/pretrained");
assert!(
fixtures_dir.is_dir(),
"decode_agreement: fixtures directory {} does not exist -- \
run scripts/extract_reference_cases.py to (re)generate it",
fixtures_dir.display()
);
let mut paths: Vec<_> = fs::read_dir(&fixtures_dir)
.unwrap_or_else(|e| {
panic!(
"decode_agreement: failed to read {}: {e}",
fixtures_dir.display()
)
})
.filter_map(|entry| entry.ok())
.map(|entry| entry.path())
.filter(|path| path.extension().is_some_and(|ext| ext == "json"))
.collect();
paths.sort();
assert!(
!paths.is_empty(),
"decode_agreement: no .json fixtures found in {} -- \
this test must never silently run on a smaller corpus",
fixtures_dir.display()
);
let mut inputs: Vec<String> = Vec::new();
for path in &paths {
let text = fs::read_to_string(path)
.unwrap_or_else(|e| panic!("decode_agreement: failed to read {}: {e}", path.display()));
let json: Value = serde_json::from_str(&text).unwrap_or_else(|e| {
panic!("decode_agreement: failed to parse {}: {e}", path.display())
});
let cases = json
.get("cases")
.and_then(Value::as_array)
.unwrap_or_else(|| panic!("decode_agreement: {}: missing `cases`", path.display()));
for case in cases {
let input = case
.get("input")
.and_then(Value::as_str)
.unwrap_or_else(|| {
panic!("decode_agreement: {}: case without `input`", path.display())
});
if !inputs.iter().any(|seen| seen == input) {
inputs.push(input.to_owned());
}
}
}
inputs
}
fn id_lists(tokenizer: &AnyTokenizer, corpus: &[String]) -> Vec<(String, Vec<u32>)> {
let mut lists: Vec<(String, Vec<u32>)> = Vec::new();
for text in corpus.iter().map(String::as_str).chain(EXTRA_TEXTS) {
lists.push((format!("text \"{}\"", escape(text)), tokenizer.encode(text)));
}
for text in ["日本語です。", "👨👩👧👦", "𐍈𐌰𐌱"] {
let ids = tokenizer.encode(text);
for len in 0..ids.len() {
lists.push((
format!("prefix {len} of \"{}\"", escape(text)),
ids[..len].to_vec(),
));
}
}
for spelling in SPECIAL_SPELLINGS {
let Some(id) = tokenizer.special_token_id(spelling) else {
continue;
};
for (before, after) in SPECIAL_NEIGHBOURS {
let mut ids = tokenizer.encode(before);
ids.push(id);
ids.extend(tokenizer.encode(after));
lists.push((
format!("special {spelling:?} between {before:?} and {after:?}"),
ids,
));
}
}
lists
}
fn stream_strict<T: Tokenize>(
tokenizer: &T,
ids: &[u32],
chunk: usize,
) -> Result<String, TokenizeError> {
let mut decoder = Tokenize::streaming_decoder(tokenizer)?;
let mut out = String::new();
for part in ids.chunks(chunk.max(1)) {
if let Some(text) = decoder.add_tokens(part)? {
out.push_str(&text);
}
}
out.push_str(&decoder.flush());
Ok(out)
}
fn stream_lossy<T: Tokenize>(
tokenizer: &T,
ids: &[u32],
chunk: usize,
) -> Result<String, TokenizeError> {
let mut decoder = Tokenize::streaming_decoder(tokenizer)?;
let mut out = String::new();
for part in ids.chunks(chunk.max(1)) {
if let Some(text) = decoder.add_tokens_lossy(part) {
out.push_str(&text);
}
}
out.push_str(&decoder.flush());
Ok(out)
}
fn stream_rendering_specials<T: Tokenize>(
tokenizer: &T,
ids: &[u32],
chunk: usize,
) -> Result<String, TokenizeError> {
let mut decoder = Tokenize::streaming_decoder_with(tokenizer, SpecialDecode::Render)?;
let mut out = String::new();
for part in ids.chunks(chunk.max(1)) {
if let Some(text) = decoder.add_tokens(part)? {
out.push_str(&text);
}
}
out.push_str(&decoder.flush());
Ok(out)
}
fn check_one<T: Tokenize>(
label: &str,
origin: &str,
tokenizer: &T,
ids: &[u32],
failures: &mut Vec<String>,
) {
let whole = Tokenize::decode(tokenizer, ids);
let whole_lossy = tokenizer.decode_lossy(ids);
let whole_rendered = Tokenize::decode_with(tokenizer, ids, SpecialDecode::Render).ok();
for chunk in 1..=ids.len().max(1) {
let strict = match stream_strict(tokenizer, ids, chunk) {
Ok(text) => text,
Err(e) => {
failures.push(format!(
"{label}: {origin}: strict stream at chunk {chunk} failed: {e}\n ids: {}",
format_ids(ids)
));
continue;
}
};
let expected_strict = match &whole {
Ok(text) => text,
Err(_) => &whole_lossy,
};
if &strict != expected_strict {
failures.push(format!(
"{label}: {origin}: strict stream at chunk {chunk} disagrees with decode\
\n ids: {}\n expected: \"{}\"\n streamed: \"{}\"",
format_ids(ids),
escape(expected_strict),
escape(&strict),
));
}
let lossy = match stream_lossy(tokenizer, ids, chunk) {
Ok(text) => text,
Err(e) => {
failures.push(format!(
"{label}: {origin}: lossy stream at chunk {chunk} failed: {e}\n ids: {}",
format_ids(ids)
));
continue;
}
};
if lossy != whole_lossy {
failures.push(format!(
"{label}: {origin}: lossy stream at chunk {chunk} disagrees with decode_lossy\
\n ids: {}\n expected: \"{}\"\n streamed: \"{}\"",
format_ids(ids),
escape(&whole_lossy),
escape(&lossy),
));
}
let Some(whole_rendered) = &whole_rendered else {
continue;
};
match stream_rendering_specials(tokenizer, ids, chunk) {
Ok(rendered) if &rendered == whole_rendered => {}
Ok(rendered) => failures.push(format!(
"{label}: {origin}: rendering-specials stream at chunk {chunk} disagrees with \
decode_with(Render)\n ids: {}\n expected: \"{}\"\n streamed: \"{}\"",
format_ids(ids),
escape(whole_rendered),
escape(&rendered),
)),
Err(e) => failures.push(format!(
"{label}: {origin}: rendering-specials stream at chunk {chunk} failed: {e}\
\n ids: {}",
format_ids(ids)
)),
}
}
}
fn check_reset<T: Tokenize>(
label: &str,
origin: &str,
tokenizer: &T,
ids: &[u32],
failures: &mut Vec<String>,
) {
let Ok(fresh) = stream_lossy(tokenizer, ids, 1) else {
failures.push(format!(
"{label}: {origin}: could not build a streaming decoder"
));
return;
};
for prefix in 0..=ids.len() {
let mut decoder = match Tokenize::streaming_decoder(tokenizer) {
Ok(decoder) => decoder,
Err(e) => {
failures.push(format!("{label}: {origin}: streaming_decoder failed: {e}"));
return;
}
};
for &id in &ids[..prefix] {
let _ = decoder.add_token_lossy(id);
}
decoder.reset();
if decoder.has_pending() || decoder.pending_bytes() != 0 {
failures.push(format!(
"{label}: {origin}: reset after {prefix} id(s) left {} pending byte(s)",
decoder.pending_bytes()
));
}
let mut reused = String::new();
for &id in ids {
if let Some(text) = decoder.add_token_lossy(id) {
reused.push_str(&text);
}
}
reused.push_str(&decoder.flush());
if reused != fresh {
failures.push(format!(
"{label}: {origin}: reset after {prefix} id(s) does not match a fresh decoder\
\n ids: {}\n fresh: \"{}\"\n reset: \"{}\"",
format_ids(ids),
escape(&fresh),
escape(&reused),
));
}
}
}
#[test]
fn streaming_decode_agrees_with_whole_sequence_decode() {
let corpus = fixture_inputs();
let mut vocabs: Vec<(&str, PretrainedVocab)> = Vec::new();
for &name in PretrainedVocab::supported_names() {
let vocab = PretrainedVocab::from_name(name).unwrap_or_else(|| {
panic!(
"decode_agreement: PretrainedVocab::from_name({name:?}) \
returned None despite being listed as supported"
)
});
if !vocabs.iter().any(|(_, seen)| *seen == vocab) {
vocabs.push((name, vocab));
}
}
let mut failures: Vec<String> = Vec::new();
for &(name, _) in &vocabs {
let handle = from_pretrained(name)
.unwrap_or_else(|e| panic!("decode_agreement: from_pretrained({name:?}) failed: {e}"));
let lists = id_lists(&handle, &corpus);
let handle_label = format!("{name}/AnyTokenizer");
for (origin, ids) in &lists {
check_one(&handle_label, origin, &handle, ids, &mut failures);
check_reset(&handle_label, origin, &handle, ids, &mut failures);
}
let backend = from_pretrained(name)
.unwrap_or_else(|e| panic!("decode_agreement: from_pretrained({name:?}) failed: {e}"))
.into_backend();
let backend_label = format!("{name}/{}", backend_name(&backend));
for (origin, ids) in &lists {
match &backend {
Backend::Bpe(t) => {
check_one(&backend_label, origin, t, ids, &mut failures);
check_reset(&backend_label, origin, t, ids, &mut failures);
}
Backend::Spm(t) => {
check_one(&backend_label, origin, t, ids, &mut failures);
check_reset(&backend_label, origin, t, ids, &mut failures);
}
Backend::Unigram(t) => {
check_one(&backend_label, origin, t, ids, &mut failures);
check_reset(&backend_label, origin, t, ids, &mut failures);
}
Backend::WordPiece(t) => {
check_one(&backend_label, origin, t, ids, &mut failures);
check_reset(&backend_label, origin, t, ids, &mut failures);
}
}
}
}
let shown = failures.len().min(20);
assert!(
failures.is_empty(),
"decode_agreement: {} streaming/decode disagreement(s), first {shown}:\n\n{}",
failures.len(),
failures[..shown].join("\n\n")
);
}
fn backend_name(backend: &Backend) -> &'static str {
match backend {
Backend::Bpe(_) => "Bpe",
Backend::Spm(_) => "Spm",
Backend::Unigram(_) => "Unigram",
Backend::WordPiece(_) => "WordPiece",
}
}