[model-gateway] fix tokenizer to match transformers special token handling (#16087)

This commit is contained in:
Simo Lin
2025-12-29 08:13:03 -08:00
committed by GitHub
parent b5d9fc873b
commit c31f62722c
18 changed files with 265 additions and 132 deletions
@@ -174,21 +174,25 @@ async fn test_cache_produces_identical_tokens() {
// Tokenize with base (no cache)
let base_encoding = base_tokenizer
.encode(turn)
.encode(turn, false)
.expect("Base tokenization failed");
let base_tokens = base_encoding.token_ids();
// Tokenize with L0-only
let l0_encoding = l0_tokenizer.encode(turn).expect("L0 tokenization failed");
let l0_encoding = l0_tokenizer
.encode(turn, false)
.expect("L0 tokenization failed");
let l0_tokens = l0_encoding.token_ids();
// Tokenize with L1-only
let l1_encoding = l1_tokenizer.encode(turn).expect("L1 tokenization failed");
let l1_encoding = l1_tokenizer
.encode(turn, false)
.expect("L1 tokenization failed");
let l1_tokens = l1_encoding.token_ids();
// Tokenize with L0+L1
let l0_l1_encoding = l0_l1_tokenizer
.encode(turn)
.encode(turn, false)
.expect("L0+L1 tokenization failed");
let l0_l1_tokens = l0_l1_encoding.token_ids();
@@ -397,13 +401,13 @@ async fn test_cache_correctness_with_edge_cases() {
test_count += 1;
let base_tokens = base_tokenizer
.encode(query)
.encode(query, false)
.expect("Base encoding failed")
.token_ids()
.to_vec();
let cached_tokens = cached_tokenizer
.encode(query)
.encode(query, false)
.expect("Cached encoding failed")
.token_ids()
.to_vec();
@@ -33,7 +33,7 @@ fn compute_hashes_for_tokenizer<E: Encoder>(tokenizer: &E, prompts: &[&str]) ->
.iter()
.map(|&prompt| {
tokenizer
.encode(prompt)
.encode(prompt, false)
.expect("Failed to encode prompt")
.get_hash()
})
@@ -63,7 +63,9 @@ fn test_tokenizer_encode_decode_lifecycle() {
.expect("Failed to load HuggingFace tokenizer");
for prompt in TEST_PROMPTS.iter() {
let encoding = tokenizer.encode(prompt).expect("Failed to encode prompt");
let encoding = tokenizer
.encode(prompt, false)
.expect("Failed to encode prompt");
let decoded = tokenizer
.decode(encoding.token_ids(), false)
@@ -82,10 +84,14 @@ fn test_sequence_operations() {
);
for prompt in TEST_PROMPTS.iter() {
let encoding = tokenizer.encode(prompt).expect("Failed to encode prompt");
let encoding = tokenizer
.encode(prompt, false)
.expect("Failed to encode prompt");
let mut sequence = Sequence::new(tokenizer.clone());
sequence.append_text(prompt).expect("Failed to append text");
sequence
.append_text(prompt, false)
.expect("Failed to append text");
assert_eq!(
sequence.len(),
@@ -123,7 +129,9 @@ fn test_decode_stream() {
);
for prompt in TEST_PROMPTS.iter() {
let encoding = tokenizer.encode(prompt).expect("Failed to encode prompt");
let encoding = tokenizer
.encode(prompt, false)
.expect("Failed to encode prompt");
let mut decoder = DecodeStream::new(tokenizer.clone(), &[], false);
let mut output = String::new();
@@ -148,11 +156,11 @@ fn test_long_sequence_incremental_decode_with_prefill() {
for (input_text, output_text) in LONG_TEST_PROMPTS.iter() {
let input_encoding = tokenizer
.encode(input_text)
.encode(input_text, false)
.expect("Failed to encode input");
let output_encoding = tokenizer
.encode(output_text)
.encode(output_text, false)
.expect("Failed to encode output");
let mut decoder = DecodeStream::new(tokenizer.clone(), input_encoding.token_ids(), false);
@@ -191,7 +199,7 @@ fn test_stop_sequence_decoder() {
let mut decoder = StopSequenceDecoder::new(tokenizer.clone(), config, false);
let encoding = tokenizer.encode(input).expect("Failed to encode");
let encoding = tokenizer.encode(input, false).expect("Failed to encode");
let mut output = String::new();
let mut stopped = false;
@@ -238,7 +246,9 @@ fn test_factory_creation() {
let tokenizer = factory::create_tokenizer(tokenizer_path.to_str().unwrap())
.expect("Failed to create tokenizer via factory");
let encoding = tokenizer.encode(TEST_PROMPTS[0]).expect("Failed to encode");
let encoding = tokenizer
.encode(TEST_PROMPTS[0], false)
.expect("Failed to encode");
let decoded = tokenizer
.decode(encoding.token_ids(), false)
@@ -254,7 +264,7 @@ fn test_batch_encoding() {
.expect("Failed to load tokenizer");
let encodings = tokenizer
.encode_batch(&TEST_PROMPTS)
.encode_batch(&TEST_PROMPTS, false)
.expect("Failed to batch encode");
assert_eq!(encodings.len(), TEST_PROMPTS.len());
@@ -300,7 +310,7 @@ fn test_thread_safety() {
let tokenizer_clone = tokenizer.clone();
thread::spawn(move || {
let encoding = tokenizer_clone
.encode(prompt)
.encode(prompt, false)
.expect("Failed to encode in thread");
let decoded = tokenizer_clone
.decode(encoding.token_ids(), false)