use splintr::{
from_json_bytes, AnyTokenizer, ByteFallback, NormOp, Normalizer, PreTokStage, PreTokenizer,
SentencePieceTokenizer, SplitBehavior, SplitPattern, SpmTokenizer, StreamingDecoder, Tokenize,
TokenizeError, Tokenizer, WordPieceTokenizer,
};
fn digit_encoder() -> splintr::FxHashMap<Vec<u8>, u32> {
[
(b"a".to_vec(), 0u32),
(b"1".to_vec(), 1u32),
(b"2".to_vec(), 2u32),
(b"12".to_vec(), 3u32),
]
.into_iter()
.collect()
}
fn digit_tokenizer() -> Tokenizer {
Tokenizer::new(digit_encoder(), splintr::FxHashMap::default(), r"\S+|\s+")
.expect("tokenizer construction")
}
#[test]
fn normalizer_can_be_built_and_attached_from_outside_the_crate() {
let normalizer = Normalizer::new(vec![NormOp::Lowercase]);
assert!(!normalizer.is_empty());
assert_eq!(normalizer.normalize("MiXeD"), "mixed");
let encoder = [(b"hello".to_vec(), 0u32), (b"HELLO".to_vec(), 1u32)]
.into_iter()
.collect();
let tokenizer = Tokenizer::new(encoder, splintr::FxHashMap::default(), r"\S+|\s+")
.expect("tokenizer construction")
.with_normalizer(normalizer);
assert_eq!(tokenizer.encode("HELLO"), vec![0]);
}
#[test]
fn empty_normalizer_is_a_valid_no_op() {
let normalizer = Normalizer::new(vec![]);
assert!(normalizer.is_empty());
assert_eq!(normalizer.normalize("Unchanged"), "Unchanged");
}
#[test]
fn normalizer_regex_op_is_constructible_from_outside_the_crate() {
let op = NormOp::replace_regex(r"\s+", "_".to_string()).expect("regex builds");
assert_eq!(Normalizer::new(vec![op]).normalize("a b"), "a_b");
assert!(NormOp::replace_regex("(", "_".to_string()).is_none());
}
#[test]
fn pre_tokenizer_can_be_built_and_attached_from_outside_the_crate() {
let pt =
PreTokenizer::new(vec![PreTokStage::Digits { individual: true }]).expect("pipeline builds");
assert!(!pt.is_empty());
assert!(!pt.byte_level());
assert_eq!(pt.stages(), [PreTokStage::Digits { individual: true }]);
assert_eq!(pt.split("a12"), vec!["a", "1", "2"]);
assert_eq!(digit_tokenizer().encode("a12"), vec![0, 3]);
assert_eq!(
digit_tokenizer().with_pre_tokenizer(pt).encode("a12"),
vec![0, 1, 2]
);
}
#[test]
fn byte_fallback_can_be_built_and_attached_from_outside_the_crate() {
let mut byte_ids = [None; 256];
byte_ids[0x62] = Some(999);
let fallback = ByteFallback::new(byte_ids, None);
let mut encoder = splintr::FxHashMap::default();
encoder.insert(b"a".to_vec(), 1u32);
encoder.insert(b"c".to_vec(), 2u32);
let tokenizer = Tokenizer::new(encoder, splintr::FxHashMap::default(), r"\S+|\s+")
.expect("tokenizer construction")
.with_byte_fallback(Some(fallback));
assert!(tokenizer.has_byte_fallback());
assert_eq!(tokenizer.encode("abc"), vec![1, 999, 2]);
}
#[test]
fn split_stage_reports_an_invalid_pattern() {
let err: splintr::TokenizerError = PreTokenizer::new(vec![PreTokStage::Split {
pattern: SplitPattern::Regex("(".to_string()),
behavior: SplitBehavior::Isolated,
invert: false,
}])
.expect_err("unbalanced group must not compile");
assert!(matches!(err, splintr::TokenizerError::RegexrError(_)));
}
#[test]
fn empty_pre_tokenizer_is_a_valid_no_op() {
let pt = PreTokenizer::new(vec![]).expect("empty pipeline builds");
assert!(pt.is_empty());
assert!(pt.stages().is_empty());
let baseline = digit_tokenizer().encode("a12");
assert_eq!(
digit_tokenizer().with_pre_tokenizer(pt).encode("a12"),
baseline
);
}
#[test]
fn every_split_behavior_is_nameable() {
let split = |behavior| {
PreTokenizer::new(vec![PreTokStage::Split {
pattern: SplitPattern::Regex(r"\s".to_string()),
behavior,
invert: false,
}])
.expect("pipeline builds")
.split("a b")
};
assert_eq!(split(SplitBehavior::Isolated), vec!["a", " ", " ", "b"]);
assert_eq!(split(SplitBehavior::Removed), vec!["a", "b"]);
assert_eq!(split(SplitBehavior::MergedWithPrevious), vec!["a ", "b"]);
assert_eq!(split(SplitBehavior::MergedWithNext), vec!["a", " b"]);
assert_eq!(split(SplitBehavior::Contiguous), vec!["a", " ", "b"]);
assert_eq!(SplitBehavior::default(), SplitBehavior::Isolated);
}
#[test]
fn streaming_decoder_is_nameable_and_only_reachable_through_the_factory() {
let mut decoder: StreamingDecoder = digit_tokenizer().streaming_decoder();
let ids = digit_tokenizer().encode("a12");
let mut streamed = String::new();
for id in &ids {
if let Some(text) = decoder.add_token(*id).expect("ids come from encode") {
streamed.push_str(&text);
}
}
streamed.push_str(&decoder.flush());
assert_eq!(streamed, "a12");
assert_eq!(
streamed,
digit_tokenizer()
.decode(&ids)
.expect("ids come from encode")
);
assert!(matches!(
decoder.add_token(9_999),
Err(splintr::TokenizeError::InvalidTokenId(9_999))
));
decoder.reset();
assert_eq!(decoder.add_token_lossy(9_999), None);
assert!(!decoder.has_pending());
assert_eq!(decoder.pending_bytes(), 0);
}
fn spm_tokenizer() -> SpmTokenizer {
let tokens = ["<unk>", "▁", "▁hello", "▁world"]
.iter()
.map(|s| (*s).to_string())
.collect();
SpmTokenizer::new(tokens, vec![], None, None).expect("vocabulary builds")
}
#[test]
fn spm_streaming_decoder_is_reachable_and_agrees_with_decode() {
let spm = spm_tokenizer();
let ids = [2u32, 3];
let mut decoder: StreamingDecoder = spm.streaming_decoder();
let mut streamed = String::new();
for id in ids {
if let Some(text) = decoder.add_token(id).expect("ids are in the vocabulary") {
streamed.push_str(&text);
}
}
streamed.push_str(&decoder.flush());
assert_eq!(streamed, "hello world");
assert_eq!(
streamed,
spm.decode(&ids).expect("ids are in the vocabulary")
);
}
#[test]
fn spm_streaming_decoder_outlives_its_tokenizer_and_keeps_the_prefix_strip() {
let mut decoder = {
let spm = spm_tokenizer();
spm.streaming_decoder()
};
assert_eq!(decoder.add_token(0).expect("known id"), None);
assert_eq!(
decoder.add_tokens(&[2, 3]).expect("known ids"),
Some("hello world".to_string())
);
}
fn unigram_tokenizer() -> SentencePieceTokenizer {
let tokens = ["<unk>", "</s>", "▁hello", "▁world", "<0x21>"]
.iter()
.map(|s| (*s).to_string())
.collect();
SentencePieceTokenizer::new(tokens, vec![], None, 1).expect("vocabulary builds")
}
#[test]
fn unigram_streaming_decoder_is_reachable_and_agrees_with_decode() {
let unigram = unigram_tokenizer();
let ids = [2u32, 3, 4];
let mut decoder: StreamingDecoder = unigram.streaming_decoder();
let mut streamed = String::new();
for id in ids {
if let Some(text) = decoder.add_token(id).expect("ids are in the vocabulary") {
streamed.push_str(&text);
}
}
streamed.push_str(&decoder.flush());
assert_eq!(streamed, "hello world!");
assert_eq!(
streamed,
unigram.decode(&ids).expect("ids are in the vocabulary")
);
}
#[test]
fn unigram_streaming_decoder_outlives_its_tokenizer_and_keeps_the_prefix_strip() {
let mut decoder = {
let unigram = unigram_tokenizer();
unigram.streaming_decoder()
};
assert_eq!(decoder.add_token(0).expect("known id"), None);
assert_eq!(
decoder.add_tokens(&[2, 3]).expect("known ids"),
Some("hello world".to_string())
);
}
fn wordpiece_tokenizer() -> WordPieceTokenizer {
let vocab = ["[PAD]", "[UNK]", "[CLS]", "[SEP]", "hello", "##ing", ","]
.iter()
.map(|s| (*s).to_string())
.collect();
WordPieceTokenizer::new(vocab, 1, 200, true)
}
#[test]
fn wordpiece_streaming_decoder_is_reachable_and_agrees_with_decode() {
let wordpiece = wordpiece_tokenizer();
let ids = [2u32, 4, 5, 6, 3];
let mut decoder: StreamingDecoder = wordpiece.streaming_decoder();
let mut streamed = String::new();
for id in ids {
if let Some(text) = decoder.add_token(id).expect("ids are in the vocabulary") {
streamed.push_str(&text);
}
}
streamed.push_str(&decoder.flush());
assert_eq!(streamed, "helloing,");
assert_eq!(
streamed,
wordpiece.decode(&ids).expect("ids are in the vocabulary")
);
}
#[test]
fn wordpiece_streaming_decoder_outlives_its_tokenizer_and_keeps_the_word_separator() {
let mut decoder = {
let wordpiece = wordpiece_tokenizer();
wordpiece.streaming_decoder()
};
assert_eq!(decoder.add_token(2).expect("known id"), None);
assert_eq!(
decoder.add_token(4).expect("known id"),
Some("hello".to_string())
);
}
#[test]
fn any_tokenizer_streaming_decoder_and_its_error_are_reachable_downstream() {
let streamable = r#"{
"added_tokens": [{"id": 1, "content": "<s>", "special": true}],
"pre_tokenizer": {"type": "Metaspace", "prepend_scheme": "first"},
"decoder": {"type": "Sequence", "decoders": [
{"type": "Replace", "pattern": {"String": "▁"}, "content": " "},
{"type": "ByteFallback"},
{"type": "Fuse"},
{"type": "Strip", "content": " ", "start": 1, "stop": 0}
]},
"model": {"type": "BPE", "byte_fallback": true, "unk_token": "<unk>",
"vocab": {"<unk>": 0, "<s>": 1, "▁Hi": 2,
"<0xE2>": 3, "<0x82>": 4, "<0xAC>": 5},
"merges": []}
}"#;
let tok: AnyTokenizer = from_json_bytes(streamable.as_bytes()).expect("the document loads");
let ids = [1u32, 2, 3, 4, 5];
let mut decoder: StreamingDecoder = tok
.streaming_decoder()
.expect("this pipeline is incrementally computable");
let mut streamed = String::new();
for id in ids {
if let Some(text) = decoder.add_token(id).expect("ids are in the vocabulary") {
streamed.push_str(&text);
}
}
streamed.push_str(&decoder.flush());
assert_eq!(streamed, "Hi\u{20ac}");
assert_eq!(streamed, tok.decode(&ids).expect("the ids decode"));
let refused = r#"{
"decoder": {"type": "BPEDecoder", "suffix": "</w>"},
"model": {"type": "BPE", "vocab": {"hello</w>": 0, "world</w>": 1}, "merges": []}
}"#;
let tok = from_json_bytes(refused.as_bytes()).expect("the document loads");
let err: TokenizeError = tok
.streaming_decoder()
.err()
.expect("a BPEDecoder pipeline cannot stream");
assert!(
matches!(err, TokenizeError::UnstreamableDecoder("BPEDecoder")),
"unexpected error: {err}"
);
assert_eq!(tok.decode(&[0, 1]).expect("decodes"), "hello world");
}
#[test]
fn streaming_decoder_outlives_the_tokenizer_that_built_it() {
let mut decoder = {
let tokenizer = digit_tokenizer();
tokenizer.streaming_decoder()
};
assert_eq!(
decoder.add_tokens(&[0, 3]).expect("known ids"),
Some("a12".to_string())
);
}
#[test]
fn both_split_patterns_are_nameable() {
let split = |pattern| {
PreTokenizer::new(vec![PreTokStage::Split {
pattern,
behavior: SplitBehavior::Removed,
invert: false,
}])
.expect("pipeline builds")
.split("a.b c")
};
assert_eq!(
split(SplitPattern::Literal(".".to_string())),
vec!["a", "b c"]
);
assert!(split(SplitPattern::Regex(".".to_string())).is_empty());
}
fn per_token_through_the_trait<T: Tokenize>(tokenizer: &T, ids: &[u32]) -> (Vec<u8>, String) {
let joined: Vec<u8> = ids
.iter()
.flat_map(|&id| tokenizer.decode_token_bytes(id).expect("id is known"))
.collect();
for &id in ids {
match tokenizer.decode_token(id) {
Ok(text) => assert_eq!(
tokenizer.decode_token_bytes(id).expect("id is known"),
text.into_bytes()
),
Err(TokenizeError::Utf8Error) => {}
Err(other) => panic!("unexpected error: {other}"),
}
}
let mut decoder: StreamingDecoder = tokenizer
.streaming_decoder()
.expect("this tokenizer can stream");
let mut streamed = decoder.add_tokens_lossy(ids).unwrap_or_default();
streamed.push_str(&decoder.flush());
assert_eq!(streamed, tokenizer.decode_lossy(ids));
(joined, streamed)
}
#[test]
fn bpe_per_token_decoding_is_reachable_through_the_trait() {
let (joined, streamed) = per_token_through_the_trait(&digit_tokenizer(), &[0, 3]);
assert_eq!(joined, b"a12".to_vec());
assert_eq!(streamed, "a12");
}
#[test]
fn wordpiece_per_token_decoding_is_reachable_through_the_trait() {
let vocab = ["[UNK]", "[CLS]", "hello", "##ing"]
.into_iter()
.map(String::from)
.collect();
let tokenizer = WordPieceTokenizer::new(vocab, 0, 200, true);
let (joined, streamed) = per_token_through_the_trait(&tokenizer, &[1, 2, 3]);
assert_eq!(joined, b"helloing".to_vec());
assert_eq!(streamed, "helloing");
}
#[test]
fn unigram_per_token_decoding_is_reachable_through_the_trait() {
let tokens = ["<unk>", "</s>", "▁hello", "▁world"]
.into_iter()
.map(String::from)
.collect();
let tokenizer = SentencePieceTokenizer::new(tokens, vec![], None, 1)
.expect("the vocabulary is well formed")
.with_prefix_space(false);
let (joined, streamed) = per_token_through_the_trait(&tokenizer, &[2, 3, 1]);
assert_eq!(joined, b" hello world".to_vec());
assert_eq!(streamed, " hello world");
}
#[test]
fn spm_per_token_decoding_is_reachable_through_the_trait() {
let mut tokens: Vec<String> = ["<unk>", "</s>", "▁hi"]
.into_iter()
.map(String::from)
.collect();
for b in 0..=255u32 {
tokens.push(format!("<0x{b:02X}>"));
}
let n = tokens.len();
let tokenizer = SpmTokenizer::new(tokens, (0..n).map(|i| -(i as f32)).collect(), None, Some(1))
.expect("the vocabulary is well formed")
.with_prefix_space(false);
let em_dash: Vec<u32> = [0xE2u32, 0x80, 0x94].into_iter().map(|b| 3 + b).collect();
assert!(matches!(
tokenizer.decode_token(em_dash[0]),
Err(TokenizeError::Utf8Error)
));
let ids = [vec![2], em_dash].concat();
let (joined, streamed) = per_token_through_the_trait(&tokenizer, &ids);
assert_eq!(String::from_utf8(joined).expect("valid together"), " hi—");
assert_eq!(streamed, " hi—");
}
#[test]
fn any_tokenizer_per_token_decoding_is_reachable_through_the_trait() {
let json = r#"{
"added_tokens": [{"id": 1, "content": "<s>", "special": true}],
"pre_tokenizer": {"type": "Metaspace", "prepend_scheme": "first"},
"decoder": {"type": "Sequence", "decoders": [
{"type": "Replace", "pattern": {"String": "▁"}, "content": " "},
{"type": "ByteFallback"},
{"type": "Fuse"},
{"type": "Strip", "content": " ", "start": 1, "stop": 0}
]},
"model": {"type": "BPE", "byte_fallback": true, "unk_token": "<unk>",
"vocab": {"<unk>": 0, "<s>": 1, "▁Hi": 2,
"<0xE2>": 3, "<0x82>": 4, "<0xAC>": 5},
"merges": []}
}"#;
let tokenizer: AnyTokenizer = from_json_bytes(json.as_bytes()).expect("the document loads");
assert!(tokenizer.declares_decoder());
let (joined, streamed) = per_token_through_the_trait(&tokenizer, &[1, 2, 3, 4, 5]);
assert_eq!(joined, " Hi\u{20ac}".as_bytes().to_vec());
assert_eq!(streamed, "Hi\u{20ac}");
}