Mask over logit rows whose draft was verified against a permissive row.
Draft i sits at local position i + 1, and drafts from num_acceptable_drafts
on met _full_mask. Local positions come from the device cu_num_logits, so
this stays right when adaptive verification trims drafts on device.
Returns None when there is nothing to invalidate.
Source code in vllm/v1/worker/gpu/structured_outputs.py
| def grammar_invalid_drafts(
input_batch: InputBatch,
grammar_req_ids: list[str],
num_acceptable_drafts: list[int] | None,
) -> torch.Tensor | None:
"""Mask over logit rows whose draft was verified against a permissive row.
Draft i sits at local position i + 1, and drafts from `num_acceptable_drafts`
on met `_full_mask`. Local positions come from the device `cu_num_logits`, so
this stays right when adaptive verification trims drafts on device.
Returns None when there is nothing to invalidate.
"""
if not grammar_req_ids or input_batch.num_draft_tokens == 0:
return None
req_id_to_idx = {req_id: i for i, req_id in enumerate(input_batch.req_ids)}
limit = np.full(input_batch.num_reqs, np.iinfo(np.int32).max, dtype=np.int32)
for i, req_id in enumerate(grammar_req_ids):
req_idx = req_id_to_idx.get(req_id)
if req_idx is not None:
# None (an older scheduler, or warmup): invalidate the whole window.
limit[req_idx] = (
num_acceptable_drafts[i] if num_acceptable_drafts is not None else 0
)
# Adaptive verification only trims drafts, so this bound holds either way.
num_drafts = input_batch.num_draft_tokens_per_req
if num_drafts is not None and (limit >= num_drafts).all():
return None
local_pos = input_batch.expanded_local_pos
cu_num_logits = input_batch.cu_num_logits
row_limit = torch.repeat_interleave(
async_tensor_h2d(limit, device=local_pos.device),
cu_num_logits[1:] - cu_num_logits[:-1],
output_size=local_pos.shape[0],
)
return local_pos > row_limit
|