use super::*;
use crate::core::added::AddedTokenSet;
use crate::core::normalizer::{NormOp, Normalizer};
use crate::core::policy::SpecialMode;
use rustc_hash::FxHashMap;
fn make_test_tokenizer() -> Tokenizer {
let mut encoder = FxHashMap::default();
for b in 32u8..=126 {
encoder.insert(vec![b], b as u32);
}
encoder.insert(b"Hello".to_vec(), 200);
encoder.insert(b"World".to_vec(), 201);
encoder.insert(b" World".to_vec(), 202);
let mut special_tokens = FxHashMap::default();
special_tokens.insert("<|endoftext|>".to_string(), 50256);
let pattern = r"\S+|\s+";
Tokenizer::new(encoder, special_tokens, pattern).unwrap()
}
#[test]
fn test_encode_decode() {
let tokenizer = make_test_tokenizer();
let text = "Hello World";
let tokens = tokenizer.encode(text);
let decoded = tokenizer.decode(&tokens).unwrap();
assert_eq!(decoded, text);
}
#[test]
fn byte_fallback_emits_fallback_id_for_unresolved_byte() {
let mut encoder = FxHashMap::default();
encoder.insert(b"a".to_vec(), 1);
encoder.insert(b"c".to_vec(), 2);
let mut byte_ids = [None; 256];
byte_ids[0x62] = Some(999);
let pattern = r"\S+|\s+";
let tokenizer = Tokenizer::new(encoder, FxHashMap::default(), pattern)
.unwrap()
.with_byte_fallback(Some(ByteFallback::new(byte_ids, None)));
assert_eq!(tokenizer.encode("abc"), vec![1, 999, 2]);
}
#[test]
fn byte_fallback_emits_one_id_per_byte_for_a_multi_byte_run() {
let mut encoder = FxHashMap::default();
encoder.insert(b"a".to_vec(), 1);
encoder.insert(b"e".to_vec(), 2);
let mut byte_ids = [None; 256];
byte_ids[0x62] = Some(900);
byte_ids[0x63] = Some(901);
byte_ids[0x64] = Some(902);
let pattern = r"\S+|\s+";
let tokenizer = Tokenizer::new(encoder, FxHashMap::default(), pattern)
.unwrap()
.with_byte_fallback(Some(ByteFallback::new(byte_ids, None)));
assert_eq!(tokenizer.encode("abcde"), vec![1, 900, 901, 902, 2]);
}
#[test]
fn no_byte_fallback_still_drops_the_unresolved_byte() {
let mut encoder = FxHashMap::default();
encoder.insert(b"a".to_vec(), 1);
encoder.insert(b"c".to_vec(), 2);
let pattern = r"\S+|\s+";
let tokenizer = Tokenizer::new(encoder, FxHashMap::default(), pattern).unwrap();
assert!(!tokenizer.has_byte_fallback());
assert_eq!(tokenizer.encode("abc"), vec![1, 2]);
}
#[test]
fn full_coverage_vocab_never_emits_fallback_and_round_trips() {
let mut encoder = FxHashMap::default();
let mut byte_ids = [None; 256];
for b in 0u32..256 {
encoder.insert(vec![b as u8], b);
byte_ids[b as usize] = Some(b);
}
let pattern = r"\S+|\s+";
let tokenizer = Tokenizer::new(encoder, FxHashMap::default(), pattern)
.unwrap()
.with_byte_fallback(Some(ByteFallback::new(byte_ids, Some(7777))));
for text in ["hello world", "你好,世界", "emoji test 😀🎉", ""] {
let tokens = tokenizer.encode(text);
assert_eq!(tokens.len(), text.len());
let decoded = tokenizer.decode(&tokens).unwrap();
assert_eq!(decoded, text);
}
}
fn partial_fallback_tokenizer(byte_62: Option<u32>, unk_id: Option<u32>) -> Tokenizer {
let mut encoder = FxHashMap::default();
encoder.insert(b"a".to_vec(), 1);
encoder.insert(b"c".to_vec(), 2);
let mut byte_ids = [None; 256];
byte_ids[0x78] = Some(3);
byte_ids[0x62] = byte_62;
let pattern = r"\S+|\s+";
Tokenizer::new(encoder, FxHashMap::default(), pattern)
.expect("the test pattern compiles")
.with_byte_fallback(Some(ByteFallback::new(byte_ids, unk_id)))
}
#[test]
fn byte_fallback_resolves_each_byte_separately_falling_back_to_unk() {
let tokenizer = partial_fallback_tokenizer(None, Some(0));
assert_eq!(tokenizer.encode("abxbc"), vec![1, 3, 0, 0, 2]);
assert_eq!(tokenizer.encode("abc"), vec![1, 0, 2]);
}
#[test]
fn declaring_the_byte_token_flips_that_byte_from_unk_to_its_own_id() {
let tokenizer = partial_fallback_tokenizer(Some(4), Some(0));
assert_eq!(tokenizer.encode("abxbc"), vec![1, 4, 3, 4, 2]);
assert_eq!(tokenizer.encode("abc"), vec![1, 4, 2]);
}
#[test]
fn byte_with_neither_a_byte_token_nor_an_unk_is_dropped() {
let tokenizer = partial_fallback_tokenizer(None, None);
assert_eq!(tokenizer.encode("abxbc"), vec![1, 3, 2]);
assert_eq!(tokenizer.encode("abc"), vec![1, 2]);
}
#[test]
fn multi_byte_char_falls_back_whole_unless_all_its_bytes_are_declared() {
let mut encoder = FxHashMap::default();
encoder.insert(b"a".to_vec(), 1);
encoder.insert(b"c".to_vec(), 2);
let pattern = r"\S+|\s+";
let build = |byte_ids: [Option<u32>; 256]| {
Tokenizer::new(encoder.clone(), FxHashMap::default(), pattern)
.expect("the test pattern compiles")
.with_byte_fallback(Some(ByteFallback::new(byte_ids, Some(0))))
};
let mut byte_ids = [None; 256];
byte_ids[0xC3] = Some(5);
assert_eq!(build(byte_ids).encode("aéc"), vec![1, 0, 2]);
byte_ids[0xA9] = Some(6);
assert_eq!(build(byte_ids).encode("aéc"), vec![1, 5, 6, 2]);
}
fn byte_operand_merge_tokenizer() -> Tokenizer {
let encoder: FxHashMap<Vec<u8>, u32> = [
("<unk>", 0),
("a", 1),
("b", 2),
("<0x7A>", 3),
("<0x7A>b", 4),
("a<0x7A>", 5),
("<0x7A><0x7A>", 6),
]
.iter()
.map(|(token, id)| (token.as_bytes().to_vec(), *id))
.collect();
let merge_ranks: FxHashMap<Vec<u8>, u32> = ["<0x7A>b", "a<0x7A>", "<0x7A><0x7A>"]
.iter()
.enumerate()
.map(|(rank, key)| (key.as_bytes().to_vec(), rank as u32))
.collect();
let mut byte_ids = [None; 256];
byte_ids[0x7A] = Some(3);
Tokenizer::new(encoder, FxHashMap::default(), r"\S+|\s+")
.expect("the test pattern compiles")
.with_merge_ranks(crate::core::encoder::encoder_from_owned(merge_ranks))
.with_byte_fallback(Some(ByteFallback::new(byte_ids, Some(0))))
}
#[test]
fn a_byte_fallback_token_merges_with_its_neighbour_as_huggingface_does() {
let tokenizer = byte_operand_merge_tokenizer();
assert_eq!(tokenizer.encode("zb"), vec![4]);
assert_eq!(tokenizer.encode("az"), vec![5]);
assert_eq!(tokenizer.encode("zz"), vec![6]);
assert_eq!(tokenizer.encode("zbz"), vec![4, 3]);
assert_eq!(tokenizer.encode("z"), vec![3]);
assert_eq!(tokenizer.encode("ab"), vec![1, 2]);
}
#[test]
fn a_multi_byte_char_s_fallback_tokens_merge_with_each_other() {
let encoder: FxHashMap<Vec<u8>, u32> = [
("<unk>", 0),
("a", 1),
("<0xC3>", 2),
("<0xA9>", 3),
("<0xC3><0xA9>", 4),
("a<0xC3>", 5),
]
.iter()
.map(|(token, id)| (token.as_bytes().to_vec(), *id))
.collect();
let merge_ranks: FxHashMap<Vec<u8>, u32> = ["<0xC3><0xA9>", "a<0xC3>"]
.iter()
.enumerate()
.map(|(rank, key)| (key.as_bytes().to_vec(), rank as u32))
.collect();
let mut byte_ids = [None; 256];
byte_ids[0xC3] = Some(2);
byte_ids[0xA9] = Some(3);
let tokenizer = Tokenizer::new(encoder, FxHashMap::default(), r"\S+|\s+")
.expect("the test pattern compiles")
.with_merge_ranks(crate::core::encoder::encoder_from_owned(merge_ranks))
.with_byte_fallback(Some(ByteFallback::new(byte_ids, Some(0))));
assert_eq!(tokenizer.encode("é"), vec![4]);
assert_eq!(tokenizer.encode("aé"), vec![1, 4]);
assert_eq!(tokenizer.encode("éa"), vec![4, 1]);
}
#[test]
fn resolving_first_leaves_a_vocabulary_without_byte_operand_merges_alone() {
let encoder: FxHashMap<Vec<u8>, u32> = [("<unk>", 0), ("a", 1), ("c", 2), ("<0x78>", 3)]
.iter()
.map(|(token, id)| (token.as_bytes().to_vec(), *id))
.collect();
let mut byte_ids = [None; 256];
byte_ids[0x78] = Some(3);
let merge_ranks: FxHashMap<Vec<u8>, u32> = [(b"ac".to_vec(), 0u32)].into_iter().collect();
let tokenizer = Tokenizer::new(encoder, FxHashMap::default(), r"\S+|\s+")
.expect("the test pattern compiles")
.with_merge_ranks(crate::core::encoder::encoder_from_owned(merge_ranks))
.with_byte_fallback(Some(ByteFallback::new(byte_ids, Some(0))));
assert_eq!(tokenizer.encode("abxbc"), vec![1, 3, 0, 0, 2]);
assert_eq!(tokenizer.encode("abc"), vec![1, 0, 2]);
}
#[test]
fn fuse_unk_still_holds_when_the_fallback_is_resolved_first() {
let encoder: FxHashMap<Vec<u8>, u32> =
[("<unk>", 0), ("a", 1), ("<0x7A>", 2), ("b", 3), ("ab", 4)]
.iter()
.map(|(token, id)| (token.as_bytes().to_vec(), *id))
.collect();
let merge_ranks: FxHashMap<Vec<u8>, u32> = [(b"ab".to_vec(), 0u32)].into_iter().collect();
let mut byte_ids = [None; 256];
byte_ids[0x7A] = Some(2);
let build = |fuse_unk: bool| {
Tokenizer::new(encoder.clone(), FxHashMap::default(), r"\S+|\s+")
.expect("the test pattern compiles")
.with_merge_ranks(crate::core::encoder::encoder_from_owned(
merge_ranks.clone(),
))
.with_byte_fallback(Some(
ByteFallback::new(byte_ids, Some(0)).with_fuse_unk(fuse_unk),
))
};
assert_eq!(build(false).encode("xxzxx"), vec![0, 2, 0, 0, 0]);
assert_eq!(build(true).encode("xxzxx"), vec![2, 0]);
assert_eq!(build(false).encode("xax"), vec![0, 1, 0]);
assert_eq!(build(true).encode("xax"), vec![0, 1, 0]);
}
#[test]
fn a_fallback_id_with_no_vocabulary_spelling_keeps_the_after_merge_answer() {
let encoder: FxHashMap<Vec<u8>, u32> = [("<unk>", 0), ("a", 1), ("c", 2)]
.iter()
.map(|(token, id)| (token.as_bytes().to_vec(), *id))
.collect();
let merge_ranks: FxHashMap<Vec<u8>, u32> = [(b"ac".to_vec(), 0u32)].into_iter().collect();
let mut byte_ids = [None; 256];
byte_ids[0x78] = Some(3);
let tokenizer = Tokenizer::new(encoder, FxHashMap::default(), r"\S+|\s+")
.expect("the test pattern compiles")
.with_merge_ranks(crate::core::encoder::encoder_from_owned(merge_ranks))
.with_byte_fallback(Some(ByteFallback::new(byte_ids, Some(0))));
assert_eq!(tokenizer.encode("abxbc"), vec![1, 3, 0, 0, 2]);
}
#[test]
fn decode_of_unknown_id_errors_with_invalid_token_id() {
let tokenizer = make_test_tokenizer();
let err = tokenizer.decode(&[7_000_000]).unwrap_err();
assert!(matches!(err, TokenizerError::InvalidTokenId(7_000_000)));
}
#[test]
fn decode_lossy_skips_unknown_ids() {
let tokenizer = make_test_tokenizer();
let mut tokens = tokenizer.encode("Hello");
tokens.push(7_000_000);
tokens.extend(tokenizer.encode(" World"));
let decoded = tokenizer.decode_lossy(&tokens);
assert_eq!(decoded, "Hello World");
}
#[test]
fn decode_is_strict_about_utf8_where_decode_lossy_substitutes() {
let mut encoder = FxHashMap::default();
for b in 0u8..=255 {
encoder.insert(vec![b], b as u32);
}
let tokenizer = Tokenizer::new(encoder, FxHashMap::default(), r".").unwrap();
for ids in [
vec![0xFFu32],
vec![0xE4],
vec![0xE4, 0xB8],
vec![0x61, 0xFF, 0x62],
] {
assert!(
matches!(tokenizer.decode(&ids), Err(TokenizerError::Utf8Error)),
"decode must stay strict for {ids:02X?}"
);
let bytes: Vec<u8> = ids.iter().map(|&id| id as u8).collect();
let lossy = tokenizer.decode_lossy(&ids);
assert!(
lossy.contains('\u{FFFD}'),
"decode_lossy must substitute for {ids:02X?}"
);
assert_eq!(lossy, String::from_utf8_lossy(&bytes));
}
}
#[test]
fn decode_batch_propagates_invalid_token_id() {
let tokenizer = make_test_tokenizer();
let good = tokenizer.encode("Hello");
let bad = vec![7_000_000u32];
let err = tokenizer.decode_batch(&[good, bad]).unwrap_err();
assert!(matches!(err, TokenizerError::InvalidTokenId(7_000_000)));
}
#[test]
fn decode_encode_round_trip_does_not_misclassify_known_ids() {
let tokenizer = make_test_tokenizer();
let text = "Hello<|endoftext|>World";
let tokens = tokenizer.encode_with_special(text);
let decoded = tokenizer.decode(&tokens).unwrap();
assert_eq!(decoded, text);
}
#[test]
fn test_encode_with_special() {
let tokenizer = make_test_tokenizer();
let text = "Hello<|endoftext|>World";
let tokens = tokenizer.encode_with_special(text);
assert!(tokens.contains(&50256));
}
#[test]
fn encode_with_all_vs_ordinary_diverge_on_a_special_token() {
let tokenizer = make_test_tokenizer().with_added_token_matching(true);
let text = "Hello<|endoftext|>World";
let all_ids = tokenizer.encode_with(text, &SpecialMode::All).unwrap();
let ordinary_ids = tokenizer.encode_with(text, &SpecialMode::Ordinary).unwrap();
assert_ne!(all_ids, ordinary_ids);
assert!(all_ids.contains(&50256));
assert!(!ordinary_ids.contains(&50256));
let decoded = tokenizer.decode(&ordinary_ids).unwrap();
assert_eq!(decoded, text);
}
#[test]
fn test_batch_encode() {
let tokenizer = make_test_tokenizer();
let texts = vec!["Hello".to_string(), "World".to_string()];
let batch_tokens = tokenizer.encode_batch(&texts);
assert_eq!(batch_tokens.len(), 2);
}
#[test]
fn test_vocab_size() {
let tokenizer = make_test_tokenizer();
assert!(tokenizer.vocab_size() > 0);
}
#[test]
fn test_cache_works() {
let tokenizer = make_test_tokenizer();
let text = "HelloWorld";
let tokens1 = tokenizer.encode(text);
let tokens2 = tokenizer.encode(text);
assert_eq!(tokens1, tokens2);
assert!(tokenizer.cache_len() > 0);
}
#[test]
fn test_clear_cache() {
let tokenizer = make_test_tokenizer();
tokenizer.encode("HelloWorld");
assert!(tokenizer.cache_len() > 0);
tokenizer.clear_cache();
assert_eq!(tokenizer.cache_len(), 0);
}
#[test]
fn cache_hit_returns_ids_for_the_queried_chunk_not_a_different_one() {
let tokenizer = make_test_tokenizer();
let texts = ["abc", "abcd", "xyz", "Hello World", "foobar", "zzz"];
let first_pass: Vec<Vec<u32>> = texts.iter().map(|t| tokenizer.encode(t)).collect();
let second_pass: Vec<Vec<u32>> = texts.iter().map(|t| tokenizer.encode(t)).collect();
assert_eq!(first_pass, second_pass);
for (text, ids) in texts.iter().zip(first_pass.iter()) {
let fresh = make_test_tokenizer();
assert_eq!(&fresh.encode(text), ids, "mismatch for {text:?}");
}
}
#[test]
fn prefix_chunk_gets_its_own_cache_entry() {
let tokenizer = make_test_tokenizer();
let short = tokenizer.encode("abc");
let len_after_short = tokenizer.cache_len();
let long = tokenizer.encode("abcd");
assert!(tokenizer.cache_len() > len_after_short);
assert_ne!(short, long);
assert_eq!(tokenizer.encode("abc"), short);
assert_eq!(tokenizer.encode("abcd"), long);
}
#[cfg(feature = "pcre2")]
#[test]
fn test_pcre2_backend() {
let tokenizer = make_test_tokenizer().pcre2(true).unwrap();
let text = "Hello World";
let tokens = tokenizer.encode(text);
let decoded = tokenizer.decode(&tokens).unwrap();
assert_eq!(decoded, text);
}
#[cfg(not(feature = "pcre2"))]
#[test]
fn test_pcre2_not_enabled() {
let tokenizer = make_test_tokenizer();
let result = tokenizer.pcre2(true);
assert!(result.is_err());
}
#[test]
fn test_jit_disable() {
let tokenizer = make_test_tokenizer().jit(false).unwrap();
let text = "Hello World";
let tokens = tokenizer.encode(text);
let decoded = tokenizer.decode(&tokens).unwrap();
assert_eq!(decoded, text);
}
#[test]
fn test_jit_enable() {
let tokenizer = make_test_tokenizer().jit(true).unwrap();
let text = "Hello World";
let tokens = tokenizer.encode(text);
let decoded = tokenizer.decode(&tokens).unwrap();
assert_eq!(decoded, text);
}
#[cfg(feature = "pcre2")]
#[test]
fn test_pcre2_switch_back_to_regexr() {
let tokenizer = make_test_tokenizer()
.pcre2(true)
.unwrap()
.pcre2(false)
.unwrap();
let text = "Hello World";
let tokens = tokenizer.encode(text);
let decoded = tokenizer.decode(&tokens).unwrap();
assert_eq!(decoded, text);
}
#[cfg(feature = "pcre2")]
#[test]
fn test_pcre2_with_jit_disabled() {
let tokenizer = make_test_tokenizer()
.jit(false)
.unwrap()
.pcre2(true)
.unwrap();
let text = "Hello World";
let tokens = tokenizer.encode(text);
let decoded = tokenizer.decode(&tokens).unwrap();
assert_eq!(decoded, text);
}
fn pieces(patterns: &[&str], text: &str) -> Vec<String> {
let tokenizer =
Tokenizer::new_byte_level_chain(FxHashMap::default(), AddedTokenSet::new(), patterns)
.expect("patterns compile");
tokenizer
.split_chunks(text)
.into_iter()
.filter_map(|(s, e)| text.get(s..e).map(str::to_owned))
.collect()
}
#[test]
fn single_expression_list_keeps_the_original_split() {
let one = Tokenizer::new_byte_level_chain(
FxHashMap::default(),
AddedTokenSet::new(),
&[GPT2_PATTERN],
)
.expect("compiles");
assert!(
one.chain.is_empty(),
"a one-expression list must not engage the chained path"
);
let plain = Tokenizer::new_byte_level(FxHashMap::default(), AddedTokenSet::new(), GPT2_PATTERN)
.expect("compiles");
let text = "Hello, world! 1234\n\n trailing";
assert_eq!(one.split_chunks(text), plain.split_chunks(text));
}
#[test]
fn later_pass_subdivides_earlier_pieces_and_cannot_re_merge() {
assert_eq!(pieces(&[GPT2_PATTERN], "abc 123"), vec!["abc", " 123"]);
assert_eq!(
pieces(&[r"\p{N}", GPT2_PATTERN], "abc 123"),
vec!["abc", " ", "1", "2", "3"],
);
}
#[test]
fn unmatched_gaps_are_kept_and_still_subdivided() {
assert_eq!(
pieces(&[r"\p{N}+", r"\p{L}+"], "ab12cd"),
vec!["ab", "12", "cd"],
);
assert_eq!(
pieces(&[r"\p{N}+", r"\p{N}+"], "ab12cd"),
vec!["ab", "12", "cd"]
);
}
#[test]
fn each_pass_matches_within_a_span_not_across_the_text() {
assert_eq!(
pieces(&[r"\p{N}+", r"^."], "ab12cd"),
vec!["a", "b", "1", "2", "c", "d"],
);
}
#[test]
fn falcon_three_pass_composition() {
let falcon = [r"[\p{P}\$\+<=>\^~\|`]+", GPT2_PATTERN, r"[0-9][0-9][0-9]"];
assert_eq!(pieces(&falcon, "a=1234"), vec!["a", "=", "123", "4"]);
assert_eq!(
pieces(&[r"[\p{P}\$\+<=>\^~\|`]+|'s| ?\p{L}+| ?\p{N}+"], "a=1234"),
vec!["a", "=", "1234"],
);
}
#[test]
fn empty_pattern_list_is_refused() {
assert!(matches!(
Tokenizer::new_byte_level_chain(FxHashMap::default(), AddedTokenSet::new(), &[]),
Err(TokenizerError::EmptyPatternList)
));
}
#[test]
fn toggling_jit_preserves_a_chained_split() {
let patterns = [r"\p{N}", GPT2_PATTERN];
let tokenizer =
Tokenizer::new_byte_level_chain(FxHashMap::default(), AddedTokenSet::new(), &patterns)
.expect("compiles");
let text = "abc 123";
let before = tokenizer.split_chunks(text);
let tokenizer = tokenizer.jit(false).expect("recompiles");
assert_eq!(tokenizer.chain.len(), 1);
assert_eq!(tokenizer.split_chunks(text), before);
}
#[test]
fn cloning_preserves_a_chained_split() {
let patterns = [r"\p{N}", GPT2_PATTERN];
let tokenizer =
Tokenizer::new_byte_level_chain(FxHashMap::default(), AddedTokenSet::new(), &patterns)
.expect("compiles");
let text = "abc 123";
assert_eq!(
tokenizer.clone().split_chunks(text),
tokenizer.split_chunks(text)
);
}
const _: () = {
assert!(super::cl100k_agent_tokens::SYSTEM > 100276);
assert!(super::cl100k_agent_tokens::SUMMARY_END == 100330);
assert!(super::o200k_agent_tokens::SYSTEM > 200018);
assert!(super::o200k_agent_tokens::SUMMARY_END == 200072);
assert!(super::cl100k_agent_tokens::USER == super::cl100k_agent_tokens::SYSTEM + 1);
assert!(super::o200k_agent_tokens::USER == super::o200k_agent_tokens::SYSTEM + 1);
};
fn make_full_byte_tokenizer() -> Tokenizer {
let mut encoder = FxHashMap::default();
for b in 0u16..=255 {
encoder.insert(vec![b as u8], b as u32);
}
let pattern = r"\S+|\s+";
Tokenizer::new(encoder, FxHashMap::default(), pattern).unwrap()
}
#[test]
fn encode_rayon_matches_encode_with_added_tokens_in_input() {
let mut encoder = FxHashMap::default();
for b in 32u8..=126 {
encoder.insert(vec![b], b as u32);
}
let mut special_tokens = FxHashMap::default();
special_tokens.insert("<|s|>".to_string(), 1000);
let tokenizer = Tokenizer::new(encoder, special_tokens, r"\S+|\s+")
.unwrap()
.with_added_token_matching(true);
let text = "a<|s|>b";
let expected = vec![97u32, 1000, 98];
assert_eq!(tokenizer.encode(text), expected);
assert_eq!(tokenizer.encode_rayon(text), expected);
}
#[test]
fn encode_rayon_matches_encode_with_normalizer() {
let tokenizer = make_full_byte_tokenizer().with_normalizer(Normalizer::new(vec![NormOp::Nfc]));
let decomposed = "e\u{0301}";
let precomposed = "\u{e9}";
assert_eq!(tokenizer.encode(decomposed), tokenizer.encode(precomposed));
assert_eq!(
tokenizer.encode_rayon(decomposed),
tokenizer.encode(decomposed)
);
}
#[test]
fn encode_rayon_matches_encode_for_metaspace_tokenizer() {
let mut encoder = FxHashMap::default();
for b in 0u16..=255 {
encoder.insert(vec![b as u8], b as u32);
}
let tokenizer =
Tokenizer::new_with_metaspace_decoder(encoder, FxHashMap::default(), r"\S+|\s+").unwrap();
let text = " hello world\tfoo bar ";
assert_eq!(tokenizer.encode_rayon(text), tokenizer.encode(text));
}
#[test]
fn encode_rayon_matches_encode_for_large_input() {
let tokenizer = make_full_byte_tokenizer();
let sentence = "Hello World, this is a test of the rayon parallel encoding path. ";
let repeats = 1 + (1_048_576 / sentence.len());
let text = sentence.repeat(repeats);
assert!(text.len() > 1_048_576);
assert_eq!(tokenizer.encode_rayon(&text), tokenizer.encode(&text));
}
#[test]
fn encode_rayon_matches_encode_for_plain_tokenizer() {
let tokenizer = make_full_byte_tokenizer();
for text in ["", " ", "你好世界", "😀🎉", "Hello World"] {
assert_eq!(
tokenizer.encode_rayon(text),
tokenizer.encode(text),
"mismatch for {text:?}"
);
}
}
fn byte_fallback_tokenizer() -> Tokenizer {
let mut encoder = FxHashMap::default();
encoder.insert(b"a".to_vec(), 1);
encoder.insert(b"c".to_vec(), 2);
for (i, b) in [0xF0u8, 0x90, 0x8D, 0x88].into_iter().enumerate() {
encoder.insert(format!("<0x{b:02X}>").into_bytes(), 10 + i as u32);
}
let byte_fallback = Tokenizer::byte_fallback_from(
|spelling| crate::core::encoder::encoder_from_owned(encoder.clone()).get(spelling),
None,
true,
);
Tokenizer::new(encoder, FxHashMap::default(), r"\S+|\s+")
.expect("the test pattern compiles")
.with_byte_fallback(byte_fallback)
}
#[test]
fn byte_fallback_ids_decode_to_the_bytes_they_denote() {
let tokenizer = byte_fallback_tokenizer();
let ids = tokenizer.encode("a𐍈c");
assert_eq!(ids, vec![1, 10, 11, 12, 13, 2]);
assert_eq!(tokenizer.decode(&ids).unwrap(), "a𐍈c");
assert_eq!(tokenizer.decode_lossy(&ids), "a𐍈c");
}
#[test]
fn without_byte_fallback_a_byte_token_surface_decodes_literally() {
let mut encoder = FxHashMap::default();
encoder.insert(b"a".to_vec(), 1);
encoder.insert(b"<0x41>".to_vec(), 2);
let tokenizer = Tokenizer::new(encoder, FxHashMap::default(), r"\S+|\s+")
.expect("the test pattern compiles");
assert!(!tokenizer.has_byte_fallback());
assert_eq!(tokenizer.decode(&[1, 2]).unwrap(), "a<0x41>");
}
fn skipping_tokenizer() -> Tokenizer {
let mut encoder = FxHashMap::default();
encoder.insert(b"Hello".to_vec(), 200);
encoder.insert(b" world".to_vec(), 201);
let mut special_tokens = FxHashMap::default();
special_tokens.insert("<|endoftext|>".to_string(), 50256);
Tokenizer::new(encoder, special_tokens, r"\S+|\s+")
.expect("the test pattern compiles")
.with_special_decode_ids([50256].into_iter().collect())
}
#[test]
fn decode_token_bytes_separates_content_skip_and_unknown() {
use crate::core::tokenize::{Tokenize, TokenizeError};
let tokenizer = skipping_tokenizer();
assert_eq!(
tokenizer.decode_token_bytes(200).unwrap(),
b"Hello".to_vec()
);
assert_eq!(tokenizer.decode_token(200).unwrap(), "Hello");
assert_eq!(
tokenizer.decode_token_bytes(50256).unwrap(),
Vec::<u8>::new()
);
assert_eq!(tokenizer.decode_token(50256).unwrap(), "");
assert!(matches!(
tokenizer.decode_token_bytes(4242),
Err(TokenizeError::InvalidTokenId(4242))
));
assert!(matches!(
tokenizer.decode_token(4242),
Err(TokenizeError::InvalidTokenId(4242))
));
}
#[test]
fn a_byte_fallback_id_has_bytes_but_no_text_of_its_own() {
use crate::core::tokenize::{Tokenize, TokenizeError};
let tokenizer = byte_fallback_tokenizer();
let ids = tokenizer.encode("a𐍈c");
assert_eq!(ids, vec![1, 10, 11, 12, 13, 2]);
for (id, byte) in ids[1..5].iter().zip([0xF0, 0x90, 0x8D, 0x88]) {
assert_eq!(tokenizer.decode_token_bytes(*id).unwrap(), vec![byte]);
assert!(matches!(
tokenizer.decode_token(*id),
Err(TokenizeError::Utf8Error)
));
}
assert_eq!(tokenizer.decode_token(1).unwrap(), "a");
}
#[test]
fn concatenated_token_bytes_equal_the_decoded_sequence() {
use crate::core::tokenize::Tokenize;
let tokenizer = byte_fallback_tokenizer();
let ids = tokenizer.encode("a𐍈c");
let joined: Vec<u8> = ids
.iter()
.flat_map(|&id| tokenizer.decode_token_bytes(id).expect("every id is known"))
.collect();
assert_eq!(joined, tokenizer.decode_lossy(&ids).into_bytes());
assert_eq!(String::from_utf8(joined).unwrap(), "a𐍈c");
}
#[test]
fn trait_decode_lossy_and_streaming_decoder_match_the_inherent_pair() {
use crate::core::tokenize::Tokenize;
let tokenizer = skipping_tokenizer();
let ids = [50256, 200, 201, 4242];
assert_eq!(Tokenize::decode_lossy(&tokenizer, &ids), "Hello world");
assert_eq!(
Tokenize::decode_lossy(&tokenizer, &ids),
Tokenizer::decode_lossy(&tokenizer, &ids)
);
let mut streamed = Tokenize::streaming_decoder(&tokenizer).expect("BPE always streams");
let mut out = streamed.add_tokens_lossy(&ids).unwrap_or_default();
out.push_str(&streamed.flush());
assert_eq!(out, "Hello world");
}
#[test]
fn the_pre_token_rung_yields_the_same_split_as_pre_tokenize() {
let mut encoder = FxHashMap::default();
for b in 0u32..256 {
encoder.insert(vec![b as u8], b);
}
let tokenizer = Tokenizer::new(encoder, FxHashMap::default(), r"\S+|\s+")
.expect("the test pattern compiles");
for text in [
"hello world",
" leading and doubled ",
"你好,世界 mixed",
"",
] {
let mut streamed = Vec::new();
tokenizer.for_each_pre_token(text, |piece| streamed.push(piece.to_owned()));
assert_eq!(
streamed,
tokenizer.pre_tokenize(text),
"the rung and `pre_tokenize` disagree on {text:?}"
);
}
}