use super::*;
fn tokenizer_with(post_processor: &str) -> Tokenizer {
let json = format!(
r#"{{"version":"1.0","truncation":null,"padding":null,"added_tokens":[],"normalizer":null,"pre_tokenizer":{{"type":"Whitespace"}},"post_processor":{post_processor},"decoder":null,"model":{{"type":"WordLevel","vocab":{{"<pad>":0,"a":1,"b":2}},"unk_token":"<pad>"}}}}"#
);
Tokenizer::from_bytes(json.as_bytes()).expect("the fixture tokenizer must parse")
}
fn template(single: &str, special_tokens: &str) -> String {
format!(
r#"{{"type":"TemplateProcessing","single":{single},"pair":[{{"Sequence":{{"id":"A","type_id":0}}}}],"special_tokens":{special_tokens}}}"#
)
}
const SPECIAL_A: &str = r#"{"a":{"id":"a","ids":[1],"tokens":["a"]}}"#;
const PIECE_A: &str = r#"{"Sequence":{"id":"A","type_id":0}}"#;
const PIECE_B: &str = r#"{"Sequence":{"id":"B","type_id":1}}"#;
const PIECE_SPECIAL_A: &str = r#"{"SpecialToken":{"id":"a","type_id":0}}"#;
const PIECE_SPECIAL_MISSING: &str = r#"{"SpecialToken":{"id":"<s>","type_id":0}}"#;
const BYTE_LEVEL: &str =
r#"{"type":"ByteLevel","add_prefix_space":true,"trim_offsets":true,"use_regex":true}"#;
const ROBERTA: &str = r#"{"type":"RobertaProcessing","sep":["b",2],"cls":["a",1],"trim_offsets":true,"add_prefix_space":false}"#;
fn sequence(members: &[&str]) -> String {
format!(
r#"{{"type":"Sequence","processors":[{}]}}"#,
members.join(",")
)
}
fn template_with_pair(single: &str, pair: &str, special_tokens: &str) -> String {
format!(
r#"{{"type":"TemplateProcessing","single":{single},"pair":{pair},"special_tokens":{special_tokens}}}"#
)
}
fn template_of_n_inputs(n: usize) -> String {
let pieces: Vec<&str> = std::iter::repeat_n(PIECE_A, n).collect();
template(&format!("[{}]", pieces.join(",")), "{}")
}
fn cls_a_sep() -> String {
template(
&format!("[{PIECE_SPECIAL_A},{PIECE_A},{PIECE_SPECIAL_A}]"),
SPECIAL_A,
)
}
#[test]
fn the_defective_fixtures_are_really_defective() {
for single in [
format!("[{PIECE_SPECIAL_MISSING},{PIECE_A}]"),
format!("[{PIECE_A},{PIECE_B}]"),
] {
let tokenizer = tokenizer_with(&template(&single, "{}"));
let panicked = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let _ = tokenizer.encode("a b", true);
}))
.is_err();
assert!(panicked, "`{single}` must panic inside the dependency");
}
let tokenizer = tokenizer_with(&template(&format!("[{PIECE_SPECIAL_A}]"), SPECIAL_A));
let ids = tokenizer
.encode("a b a b", true)
.expect("a template with no $A does not panic — it drops the text")
.get_ids()
.to_vec();
assert_eq!(
ids,
vec![1],
"the text is gone; only the special token remains"
);
}
#[test]
fn a_tokenizer_without_a_post_processor_passes() {
assert_eq!(check_post_processor(&tokenizer_with("null")), Ok(()));
}
#[test]
fn template_free_post_processors_pass() {
for post in [
r#"{"type":"RobertaProcessing","sep":["</s>",2],"cls":["<s>",0],"trim_offsets":true,"add_prefix_space":false}"#,
r#"{"type":"BertProcessing","sep":["[SEP]",2],"cls":["[CLS]",0]}"#,
r#"{"type":"ByteLevel","add_prefix_space":true,"trim_offsets":true,"use_regex":true}"#,
] {
assert_eq!(
check_post_processor(&tokenizer_with(post)),
Ok(()),
"{post} carries no template"
);
}
}
#[test]
fn a_well_formed_single_template_passes() {
let post = template(&format!("[{PIECE_SPECIAL_A},{PIECE_A}]"), SPECIAL_A);
assert_eq!(check_post_processor(&tokenizer_with(&post)), Ok(()));
}
#[test]
fn an_undeclared_special_token_is_refused_and_named() {
let post = template(&format!("[{PIECE_SPECIAL_MISSING},{PIECE_A}]"), "{}");
assert_eq!(
check_post_processor(&tokenizer_with(&post)),
Err(PostProcessorTemplate::UndeclaredSpecialToken(
"<s>".to_string()
))
);
}
#[test]
fn an_undeclared_id_beside_a_declared_one_is_refused() {
let post = template(
&format!("[{PIECE_SPECIAL_A},{PIECE_SPECIAL_MISSING},{PIECE_A}]"),
SPECIAL_A,
);
assert_eq!(
check_post_processor(&tokenizer_with(&post)),
Err(PostProcessorTemplate::UndeclaredSpecialToken(
"<s>".to_string()
))
);
}
#[test]
fn a_pair_sequence_in_the_single_template_is_refused() {
let post = template(&format!("[{PIECE_A},{PIECE_B}]"), "{}");
assert_eq!(
check_post_processor(&tokenizer_with(&post)),
Err(PostProcessorTemplate::PairSequenceInSingleTemplate)
);
}
#[test]
fn a_single_template_that_never_places_the_text_is_refused() {
let post = template(&format!("[{PIECE_SPECIAL_A}]"), SPECIAL_A);
assert_eq!(
check_post_processor(&tokenizer_with(&post)),
Err(PostProcessorTemplate::NoInputSequenceInSingleTemplate)
);
}
#[test]
fn an_empty_single_template_is_refused() {
let post = template("[]", "{}");
assert_eq!(
check_post_processor(&tokenizer_with(&post)),
Err(PostProcessorTemplate::NoInputSequenceInSingleTemplate)
);
}
#[test]
fn the_pair_template_is_not_judged() {
let post = format!(
r#"{{"type":"TemplateProcessing","single":[{PIECE_A}],"pair":[{PIECE_SPECIAL_MISSING},{PIECE_B}],"special_tokens":{{}}}}"#
);
assert_eq!(check_post_processor(&tokenizer_with(&post)), Ok(()));
}
#[test]
fn a_defective_template_nested_in_a_sequence_is_refused() {
let inner = template(&format!("[{PIECE_SPECIAL_MISSING},{PIECE_A}]"), "{}");
let post = format!(r#"{{"type":"Sequence","processors":[{inner}]}}"#);
assert_eq!(
check_post_processor(&tokenizer_with(&post)),
Err(PostProcessorTemplate::UndeclaredSpecialToken(
"<s>".to_string()
))
);
}
#[test]
fn a_sequence_of_sound_processors_passes() {
let inner = template(&format!("[{PIECE_SPECIAL_A},{PIECE_A}]"), SPECIAL_A);
let post = format!(
r#"{{"type":"Sequence","processors":[{{"type":"ByteLevel","add_prefix_space":true,"trim_offsets":true,"use_regex":true}},{inner}]}}"#
);
assert_eq!(check_post_processor(&tokenizer_with(&post)), Ok(()));
}
#[test]
fn every_reason_displays_distinctly() {
let rendered = [
PostProcessorTemplate::UndeclaredSpecialToken("<s>".to_string()).to_string(),
PostProcessorTemplate::PairSequenceInSingleTemplate.to_string(),
PostProcessorTemplate::NoInputSequenceInSingleTemplate.to_string(),
PostProcessorTemplate::RepeatedInputSequence(2).to_string(),
PostProcessorTemplate::UnsupportedEncodingCount(3).to_string(),
PostProcessorTemplate::Unreadable.to_string(),
];
assert!(rendered[0].contains("`<s>`"), "{}", rendered[0]);
assert!(rendered[1].contains("$B"), "{}", rendered[1]);
assert!(rendered[2].contains("$A"), "{}", rendered[2]);
assert!(rendered[3].contains('2'), "{}", rendered[3]);
assert!(rendered[4].contains('3'), "{}", rendered[4]);
for (i, a) in rendered.iter().enumerate() {
for b in &rendered[i + 1..] {
assert_ne!(a, b, "each reason must render distinctly");
}
}
}
fn truncated_at(tokenizer: &Tokenizer, window: usize) -> Tokenizer {
use tokenizers::{TruncationDirection, TruncationParams, TruncationStrategy};
let mut tokenizer = tokenizer.clone();
tokenizer
.with_truncation(Some(TruncationParams {
max_length: window,
strategy: TruncationStrategy::LongestFirst,
stride: 0,
direction: TruncationDirection::Right,
}))
.expect("the fixtures here never over-fill the window");
tokenizer.with_padding(None);
tokenizer
}
fn declared_overhead(tokenizer: &Tokenizer) -> usize {
tokenizers::PostProcessor::added_tokens(
tokenizer.get_post_processor().expect("the fixture has one"),
false,
)
}
#[test]
fn the_cardinality_fixtures_are_really_defective() {
for chain in [
sequence(&[&cls_a_sep(), &template(&format!("[{PIECE_A}]"), "{}")]),
sequence(&[&template_of_n_inputs(3), &cls_a_sep()]),
] {
let tokenizer = tokenizer_with(&chain);
let hook = std::panic::take_hook();
std::panic::set_hook(Box::new(|_| {}));
let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let _ = tokenizer.encode("a b", true);
}));
std::panic::set_hook(hook);
let payload = outcome.expect_err(&format!("`{chain}` must panic inside the dependency"));
let message = payload
.downcast_ref::<String>()
.cloned()
.or_else(|| payload.downcast_ref::<&str>().map(|s| (*s).to_string()))
.unwrap_or_default();
assert_eq!(message, "not yet implemented", "for `{chain}`");
}
}
#[test]
fn a_template_fed_more_than_two_encodings_is_refused_with_the_count() {
let chain = sequence(&[&cls_a_sep(), &template(&format!("[{PIECE_A}]"), "{}")]);
assert_eq!(
check_post_processor(&tokenizer_with(&chain)),
Err(PostProcessorTemplate::UnsupportedEncodingCount(3)),
"for `{chain}`"
);
}
#[test]
fn a_template_fed_no_encodings_is_refused_with_the_count() {
let template: serde_json::Value =
serde_json::from_str(&cls_a_sep()).expect("the fixture is JSON");
assert_eq!(
check_value(&template, 0),
Err(PostProcessorTemplate::UnsupportedEncodingCount(0))
);
}
#[test]
fn a_repeated_input_sequence_really_doubles_the_text() {
let tokenizer = tokenizer_with(&template_of_n_inputs(2));
assert_eq!(
declared_overhead(&tokenizer),
0,
"a `Sequence` piece counts as no overhead, however many times it appears"
);
assert_eq!(
tokenizer.encode("a b a", true).expect("encode").get_ids(),
&[1, 2, 1, 1, 2, 1],
"three text tokens in, six out"
);
let ids = truncated_at(&tokenizer, 4)
.encode("a b a", true)
.expect("encode")
.get_ids()
.to_vec();
assert_eq!(
ids,
vec![1, 2, 1, 1, 2, 1],
"truncation is sized for one copy, so it does not fire and the window is blown"
);
assert!(ids.len() > 4, "past a four-token window: {} ids", ids.len());
}
#[test]
fn a_doubled_text_folded_by_a_pair_template_erases_it() {
let doubler = template_of_n_inputs(2);
for (pair, expected) in [
("[]".to_string(), Vec::new()),
(
format!("[{PIECE_SPECIAL_A},{PIECE_SPECIAL_A}]"),
vec![1u32, 1],
),
] {
let chain = sequence(&[
&doubler,
&template_with_pair(&format!("[{PIECE_A}]"), &pair, SPECIAL_A),
]);
let tokenizer = tokenizer_with(&chain);
for text in ["some words", "a b", "a b a b a b"] {
assert_eq!(
tokenizer.encode(text, true).expect("encode").get_ids(),
expected.as_slice(),
"every input encodes identically under pair `{pair}`"
);
}
}
}
#[test]
fn a_single_template_that_places_the_text_more_than_once_is_refused() {
for n in [2usize, 3, 5] {
assert_eq!(
check_post_processor(&tokenizer_with(&template_of_n_inputs(n))),
Err(PostProcessorTemplate::RepeatedInputSequence(n)),
);
}
}
#[test]
fn a_chain_whose_first_template_repeats_the_text_is_refused_there() {
let doubler = template_of_n_inputs(2);
for pair in [
"[]".to_string(),
format!("[{PIECE_SPECIAL_A},{PIECE_SPECIAL_A}]"),
] {
let chain = sequence(&[
&doubler,
&template_with_pair(&format!("[{PIECE_A}]"), &pair, SPECIAL_A),
]);
assert_eq!(
check_post_processor(&tokenizer_with(&chain)),
Err(PostProcessorTemplate::RepeatedInputSequence(2)),
"for `{chain}`"
);
}
}
#[test]
fn the_pair_mode_fixtures_are_really_defective() {
let a_then_special = template(&format!("[{PIECE_A},{PIECE_SPECIAL_A}]"), SPECIAL_A);
for (pair, expected) in [
("[]".to_string(), Vec::new()),
(
format!("[{PIECE_SPECIAL_A},{PIECE_SPECIAL_A}]"),
vec![1u32, 1],
),
] {
let chain = sequence(&[
&a_then_special,
&template_with_pair(&format!("[{PIECE_A}]"), &pair, SPECIAL_A),
]);
let tokenizer = tokenizer_with(&chain);
for text in ["some words", "a b", "a b a b a b"] {
assert_eq!(
tokenizer.encode(text, true).expect("encode").get_ids(),
expected.as_slice(),
"every input encodes identically under pair `{pair}`"
);
}
}
let chain = sequence(&[
&a_then_special,
&template_with_pair(
&format!("[{PIECE_A}]"),
&format!("[{PIECE_SPECIAL_A},{PIECE_A},{PIECE_B},{PIECE_SPECIAL_A}]"),
SPECIAL_A,
),
]);
let tokenizer = tokenizer_with(&chain);
assert_eq!(
declared_overhead(&tokenizer),
1,
"the sum of the members' `added_single`, whichever mode each really runs in"
);
let ids = truncated_at(&tokenizer, 4)
.encode("a b a b a b", true)
.expect("encode")
.get_ids()
.to_vec();
assert_eq!(
ids,
vec![1, 1, 2, 1, 1, 1],
"sized on an overhead of 1, the pair template adds 3"
);
assert!(ids.len() > 4, "past a four-token window: {} ids", ids.len());
}
#[test]
fn a_template_fed_two_encodings_is_refused_with_the_count() {
let a_then_special = template(&format!("[{PIECE_A},{PIECE_SPECIAL_A}]"), SPECIAL_A);
for pair in [
"[]".to_string(),
format!("[{PIECE_SPECIAL_A},{PIECE_SPECIAL_A}]"),
format!("[{PIECE_SPECIAL_A},{PIECE_A},{PIECE_B},{PIECE_SPECIAL_A}]"),
] {
let chain = sequence(&[
&a_then_special,
&template_with_pair(&format!("[{PIECE_A}]"), &pair, SPECIAL_A),
]);
assert_eq!(
check_post_processor(&tokenizer_with(&chain)),
Err(PostProcessorTemplate::UnsupportedEncodingCount(2)),
"for pair `{pair}`"
);
}
}
#[test]
fn the_token_adding_passthroughs_under_count_their_overhead() {
const BERT: &str = r#"{"type":"BertProcessing","sep":["b",2],"cls":["a",1]}"#;
for (passthrough, declared, length) in [(ROBERTA, 4, 12), (BERT, 4, 10), (BYTE_LEVEL, 2, 8)] {
let chain = sequence(&[&cls_a_sep(), passthrough]);
let tokenizer = tokenizer_with(&chain);
assert_eq!(declared_overhead(&tokenizer), declared, "for {passthrough}");
let ids = truncated_at(&tokenizer, 8)
.encode("a b a b a b", true)
.expect("encode")
.get_ids()
.to_vec();
assert_eq!(ids.len(), length, "for {passthrough}: {ids:?}");
}
}
#[test]
fn the_token_adding_passthroughs_are_refused_past_one_encoding() {
const BERT: &str = r#"{"type":"BertProcessing","sep":["b",2],"cls":["a",1]}"#;
for passthrough in [ROBERTA, BERT] {
let chain = sequence(&[&cls_a_sep(), passthrough]);
assert_eq!(
check_post_processor(&tokenizer_with(&chain)),
Err(PostProcessorTemplate::UnsupportedEncodingCount(3)),
"for {passthrough}"
);
}
let chain = sequence(&[&cls_a_sep(), BYTE_LEVEL]);
assert_eq!(
check_post_processor(&tokenizer_with(&chain)),
Ok(()),
"`ByteLevel` adds nothing at any count"
);
}
#[test]
fn byte_level_propagates_the_count() {
let chain = sequence(&[
&template(&format!("[{PIECE_A},{PIECE_SPECIAL_A}]"), SPECIAL_A),
BYTE_LEVEL,
&template(&format!("[{PIECE_A}]"), "{}"),
]);
assert_eq!(
check_post_processor(&tokenizer_with(&chain)),
Err(PostProcessorTemplate::UnsupportedEncodingCount(2))
);
}
#[test]
fn the_template_free_kinds_pass_a_sound_chain() {
const BERT: &str = r#"{"type":"BertProcessing","sep":["b",2],"cls":["a",1]}"#;
for passthrough in [ROBERTA, BERT, BYTE_LEVEL] {
for chain in [
sequence(&[passthrough]),
sequence(&[passthrough, &template(&format!("[{PIECE_A}]"), "{}")]),
sequence(&[&template(&format!("[{PIECE_A}]"), "{}"), passthrough]),
] {
assert_eq!(
check_post_processor(&tokenizer_with(&chain)),
Ok(()),
"for `{chain}`"
);
}
}
}
#[test]
fn a_nested_sequence_hands_its_count_to_its_sibling() {
let inner = sequence(&[&template(
&format!("[{PIECE_A},{PIECE_SPECIAL_A}]"),
SPECIAL_A,
)]);
let chain = sequence(&[&inner, &template(&format!("[{PIECE_A}]"), "{}")]);
assert_eq!(
check_post_processor(&tokenizer_with(&chain)),
Err(PostProcessorTemplate::UnsupportedEncodingCount(2))
);
}
#[test]
fn a_chain_that_keeps_the_count_at_one_is_accepted_and_encodes() {
let alone = tokenizer_with(&cls_a_sep());
assert_eq!(check_post_processor(&alone), Ok(()));
let expected = alone
.encode("a b", true)
.expect("encode")
.get_ids()
.to_vec();
assert_eq!(
expected,
vec![1, 1, 2, 1],
"the special token, the text, the special token — merged into one encoding"
);
for chain in [
sequence(&[&template(&format!("[{PIECE_A}]"), "{}"), &cls_a_sep()]),
sequence(&[BYTE_LEVEL, &cls_a_sep()]),
sequence(&[&cls_a_sep()]),
] {
let tokenizer = tokenizer_with(&chain);
assert_eq!(check_post_processor(&tokenizer), Ok(()), "for `{chain}`");
assert_eq!(
tokenizer.encode("a b", true).expect("encode").get_ids(),
expected.as_slice(),
"for `{chain}`"
);
}
}
#[test]
fn an_unrecognized_kind_passes_and_keeps_the_count() {
for count in [0, 1, 2, 3] {
let unknown = serde_json::json!({"type": "SomeProcessorAddedLater"});
assert_eq!(check_value(&unknown, count), Ok(count));
}
let chain = serde_json::json!({
"type": "Sequence",
"processors": [
{"type": "SomeProcessorAddedLater"},
serde_json::from_str::<serde_json::Value>(&template(
&format!("[{PIECE_A},{PIECE_SPECIAL_A}]"),
SPECIAL_A,
))
.expect("json"),
{"type": "SomeProcessorAddedLater"},
serde_json::from_str::<serde_json::Value>(&template(&format!("[{PIECE_A}]"), "{}"))
.expect("json"),
],
});
assert_eq!(
check_value(&chain, 1),
Err(PostProcessorTemplate::UnsupportedEncodingCount(2))
);
}