[Bugfix][Model Runner V2] Preserve all allowed_token_ids in the logit bias kernel (#46245)

Signed-off-by: Ting Sun <suntcrick@gmail.com>
This commit is contained in:
Ting SUN
2026-06-23 07:01:13 +00:00
committed by GitHub
parent 6c427dd401
commit a46f3eb232
2 changed files with 12 additions and 0 deletions
@@ -152,6 +152,14 @@ def test_allowed_token_ids(llm):
output = llm.generate(PROMPT, SamplingParams(allowed_token_ids=allowed_token_ids))
assert output[0].outputs[0].token_ids[-1] == TOKEN_ID
# Each single-token allowlist must force that token (kernel used to drop some).
for token_id in (1, 5, 100, 500, 2518, 9834, 31999):
output = llm.generate(
PROMPT,
SamplingParams(temperature=0, max_tokens=1, allowed_token_ids=[token_id]),
)
assert output[0].outputs[0].token_ids[-1] == token_id
# Reject empty allowed_token_ids.
with pytest.raises(ValueError):
_ = llm.generate(PROMPT, SamplingParams(allowed_token_ids=[]))
+4
View File
@@ -189,6 +189,8 @@ def _bias_kernel(
logits_ptr + token_idx * logits_stride + allowed_token_ids, mask=mask
)
tl.debug_barrier() # save must read original logits before the -inf overwrite
# Set logits to -inf for all tokens.
for i in range(0, vocab_size, LOGITS_BLOCK_SIZE):
offset = i + tl.arange(0, LOGITS_BLOCK_SIZE)
@@ -198,6 +200,8 @@ def _bias_kernel(
mask=offset < vocab_size,
)
tl.debug_barrier() # -inf overwrite must finish before restoring saved logits
# Restore logits for allowed token IDs.
tl.store(
logits_ptr + token_idx * logits_stride + allowed_token_ids,