Skip to content

vllm.v1.worker.gpu.structured_outputs

Functions:

grammar_invalid_drafts(input_batch, grammar_req_ids, num_acceptable_drafts)

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