use serde_json::Value;
use splintr::pretrained::from_pretrained;
use splintr::AnyTokenizer;
use std::fs;
use std::path::Path;
#[derive(Clone, Copy)]
enum Mode {
Raw,
Template,
}
impl Mode {
const ALL: [Mode; 2] = [Mode::Raw, Mode::Template];
fn name(self) -> &'static str {
match self {
Mode::Raw => "encode_raw",
Mode::Template => "encode",
}
}
fn run(self, tokenizer: &AnyTokenizer, text: &str) -> Vec<u32> {
match self {
Mode::Raw => tokenizer.encode_raw(text),
Mode::Template => tokenizer.encode(text),
}
}
}
struct Case {
input: String,
expected: Vec<u32>,
decoded: String,
pieces: Option<Vec<String>>,
normalized: Option<String>,
}
struct Fixture {
vocab: String,
cases: Vec<Case>,
}
fn load_fixture(path: &Path) -> Fixture {
let text = fs::read_to_string(path)
.unwrap_or_else(|e| panic!("reference_parity: failed to read {}: {e}", path.display()));
let json: Value = serde_json::from_str(&text)
.unwrap_or_else(|e| panic!("reference_parity: failed to parse {}: {e}", path.display()));
let vocab = json
.get("vocab")
.and_then(Value::as_str)
.unwrap_or_else(|| {
panic!(
"reference_parity: {}: missing `vocab` string",
path.display()
)
})
.to_owned();
let array = json
.get("cases")
.and_then(Value::as_array)
.unwrap_or_else(|| {
panic!(
"reference_parity: {}: missing `cases` array",
path.display()
)
});
let cases = array
.iter()
.enumerate()
.map(|(index, entry)| {
let input = entry
.get("input")
.and_then(Value::as_str)
.unwrap_or_else(|| {
panic!(
"reference_parity: {}: case {index}: `input` missing or not a string",
path.display()
)
})
.to_owned();
let expected: Vec<u32> = entry
.get("expected")
.and_then(Value::as_array)
.unwrap_or_else(|| {
panic!(
"reference_parity: {}: case {index}: `expected` missing or not an array",
path.display()
)
})
.iter()
.enumerate()
.map(|(id_index, v)| {
v.as_u64()
.and_then(|n| u32::try_from(n).ok())
.unwrap_or_else(|| {
panic!(
"reference_parity: {}: case {index}: expected[{id_index}] is not a valid token id",
path.display()
)
})
})
.collect();
let decoded = entry
.get("decoded")
.and_then(Value::as_str)
.unwrap_or_else(|| {
panic!(
"reference_parity: {}: case {index}: `decoded` missing or not a string -- \
regenerate this fixture with scripts/extract_reference_cases.py",
path.display()
)
})
.to_owned();
let pieces = entry.get("pieces").map(|value| {
value
.as_array()
.unwrap_or_else(|| {
panic!(
"reference_parity: {}: case {index}: `pieces` is not an array",
path.display()
)
})
.iter()
.enumerate()
.map(|(piece_index, v)| {
v.as_str()
.unwrap_or_else(|| {
panic!(
"reference_parity: {}: case {index}: pieces[{piece_index}] is not a string",
path.display()
)
})
.to_owned()
})
.collect()
});
let normalized = entry.get("normalized").map(|value| {
value
.as_str()
.unwrap_or_else(|| {
panic!(
"reference_parity: {}: case {index}: `normalized` is not a string",
path.display()
)
})
.to_owned()
});
Case {
input,
expected,
decoded,
pieces,
normalized,
}
})
.collect();
Fixture { vocab, cases }
}
fn first_diff<T: PartialEq>(expected: &[T], actual: &[T]) -> Option<usize> {
let shared = expected.len().min(actual.len());
for i in 0..shared {
if expected[i] != actual[i] {
return Some(i);
}
}
if expected.len() == actual.len() {
None
} else {
Some(shared)
}
}
fn at_or_end(ids: &[u32], index: usize) -> String {
match ids.get(index) {
Some(id) => id.to_string(),
None => "<end>".to_owned(),
}
}
fn format_ids(ids: &[u32]) -> String {
ids.iter().map(u32::to_string).collect::<Vec<_>>().join(" ")
}
fn piece_at_or_end(pieces: &[String], index: usize) -> String {
match pieces.get(index) {
Some(piece) => format!("\"{}\"", escape(piece)),
None => "<end>".to_owned(),
}
}
fn format_pieces(pieces: &[String]) -> String {
pieces
.iter()
.map(|piece| format!("\"{}\"", escape(piece)))
.collect::<Vec<_>>()
.join(" ")
}
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
}
#[test]
fn pretrained_vocabularies_match_reference_tokenizers() {
let fixtures_dir = Path::new(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/pretrained");
assert!(
fixtures_dir.is_dir(),
"reference_parity: fixtures directory {} does not exist -- \
run scripts/extract_reference_cases.py to (re)generate it",
fixtures_dir.display()
);
let mut fixture_paths: Vec<_> = fs::read_dir(&fixtures_dir)
.unwrap_or_else(|e| {
panic!(
"reference_parity: 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();
fixture_paths.sort();
assert!(
!fixture_paths.is_empty(),
"reference_parity: no .json fixtures found in {} -- \
this test must never silently pass with zero fixtures",
fixtures_dir.display()
);
let mut failure_reports: Vec<String> = Vec::new();
for path in &fixture_paths {
let fixture = load_fixture(path);
let tokenizer = from_pretrained(&fixture.vocab).unwrap_or_else(|e| {
panic!(
"reference_parity: {}: from_pretrained({:?}) failed: {e}",
path.display(),
fixture.vocab
)
});
let mut per_mode: Vec<(Mode, Vec<usize>)> = Vec::new();
for mode in Mode::ALL {
let failures: Vec<usize> = fixture
.cases
.iter()
.enumerate()
.filter(|(_, case)| mode.run(&tokenizer, &case.input) != case.expected)
.map(|(index, _)| index)
.collect();
per_mode.push((mode, failures));
}
let (best_mode, best_failures) = per_mode
.into_iter()
.min_by_key(|(_, failures)| failures.len())
.expect("Mode::ALL is non-empty");
let decode_failures: Vec<(usize, Result<String, String>)> = fixture
.cases
.iter()
.enumerate()
.filter_map(|(index, case)| match tokenizer.decode(&case.expected) {
Ok(text) if text == case.decoded => None,
Ok(text) => Some((index, Ok(text))),
Err(e) => Some((index, Err(e.to_string()))),
})
.collect();
if !decode_failures.is_empty() {
let mut detail = format!(
"vocab {:?} ({}): {}/{} cases decoded differently from the reference:",
fixture.vocab,
path.display(),
decode_failures.len(),
fixture.cases.len(),
);
for (index, outcome) in &decode_failures {
let case = &fixture.cases[*index];
let actual = match outcome {
Ok(text) => format!("\"{}\"", escape(text)),
Err(message) => format!("<error: {message}>"),
};
detail.push_str(&format!(
"\n [case {index}] input: \"{}\"\n ids: {}\
\n expected: \"{}\"\n actual: {actual}",
escape(&case.input),
format_ids(&case.expected),
escape(&case.decoded),
));
}
failure_reports.push(detail);
}
let normalize_failures: Vec<(usize, Option<String>)> = fixture
.cases
.iter()
.enumerate()
.filter_map(|(index, case)| {
let expected = case.normalized.as_ref()?;
match tokenizer.normalize(&case.input) {
Some(actual) if actual == *expected => None,
outcome => Some((index, outcome)),
}
})
.collect();
if !normalize_failures.is_empty() {
let mut detail = format!(
"vocab {:?} ({}): {}/{} cases normalized differently from the reference:",
fixture.vocab,
path.display(),
normalize_failures.len(),
fixture.cases.len(),
);
for (index, outcome) in &normalize_failures {
let case = &fixture.cases[*index];
let actual = match outcome {
Some(text) => format!("\"{}\"", escape(text)),
None => "<this backend exposes no normalization stage>".to_owned(),
};
detail.push_str(&format!(
"\n [case {index}] input: \"{}\"\n expected: \"{}\"\n actual: {actual}",
escape(&case.input),
escape(case.normalized.as_deref().unwrap_or_default()),
));
}
failure_reports.push(detail);
}
let piece_failures: Vec<(usize, Option<Vec<String>>)> = fixture
.cases
.iter()
.enumerate()
.filter_map(|(index, case)| {
let expected = case.pieces.as_ref()?;
let text = case.normalized.as_deref().unwrap_or(&case.input);
match tokenizer.pre_tokenize(text) {
Some(actual) if actual == *expected => None,
outcome => Some((index, outcome)),
}
})
.collect();
if !piece_failures.is_empty() {
let mut detail = format!(
"vocab {:?} ({}): {}/{} cases pre-tokenized differently from the reference:",
fixture.vocab,
path.display(),
piece_failures.len(),
fixture.cases.len(),
);
for (index, outcome) in &piece_failures {
let case = &fixture.cases[*index];
let expected = case.pieces.as_deref().unwrap_or(&[]);
let text = case.normalized.as_deref().unwrap_or(&case.input);
detail.push_str(&format!(
"\n [case {index}] input: \"{}\"\n expected ({:>3}): {}",
escape(text),
expected.len(),
format_pieces(expected),
));
match outcome {
Some(actual) => {
detail.push_str(&format!(
"\n actual ({:>3}): {}",
actual.len(),
format_pieces(actual),
));
match first_diff(expected, actual) {
Some(at) => detail.push_str(&format!(
"\n first differs at index {at}: expected {}, got {}",
piece_at_or_end(expected, at),
piece_at_or_end(actual, at)
)),
None => {
detail.push_str("\n (sequences compare equal -- unreachable)")
}
}
}
None => detail
.push_str("\n actual: <this backend exposes no pre-tokenizer split>"),
}
}
failure_reports.push(detail);
}
if best_failures.is_empty() {
continue;
}
let mut detail = format!(
"vocab {:?} ({}): {}/{} cases mismatched under best mode `{}`:",
fixture.vocab,
path.display(),
best_failures.len(),
fixture.cases.len(),
best_mode.name(),
);
for &index in &best_failures {
let case = &fixture.cases[index];
let actual = best_mode.run(&tokenizer, &case.input);
let diff_at = first_diff(&case.expected, &actual);
detail.push_str(&format!(
"\n [case {index}] input: \"{}\"\n expected ({:>3}): {}\n actual ({:>3}): {}",
escape(&case.input),
case.expected.len(),
format_ids(&case.expected),
actual.len(),
format_ids(&actual),
));
match diff_at {
Some(at) => detail.push_str(&format!(
"\n first differs at index {at}: expected {}, got {}",
at_or_end(&case.expected, at),
at_or_end(&actual, at)
)),
None => detail.push_str("\n (sequences compare equal -- unreachable)"),
}
}
failure_reports.push(detail);
}
assert!(
failure_reports.is_empty(),
"reference_parity: splintr disagrees with the reference tokenizer:\n\n{}",
failure_reports.join("\n\n")
);
}