Skip to content

vllm.v1.worker.gpu.sample.gumbel

Functions:

gumbel_noised_argmax(logits, keys, mask, seed, pos, temp, IS_DRAFTING, USE_FP64, APPLY_TEMPERATURE=True)

Argmax of logits under Gumbel-max sampling, or plain argmax at temp 0.

keys indexes the noise, so the same token draws the same noise wherever it appears; pos and seed place the draw in the request's stream, which is what lets a draft and its verification agree.

Source code in vllm/v1/worker/gpu/sample/gumbel.py
@triton.jit
def gumbel_noised_argmax(
    logits,
    keys,
    mask,
    seed,
    pos,
    temp,
    IS_DRAFTING: tl.constexpr,
    USE_FP64: tl.constexpr,
    APPLY_TEMPERATURE: tl.constexpr = True,
):
    """Argmax of logits under Gumbel-max sampling, or plain argmax at temp 0.

    `keys` indexes the noise, so the same token draws the same noise wherever it
    appears; `pos` and `seed` place the draw in the request's stream, which is
    what lets a draft and its verification agree.
    """
    if temp != 0.0 and APPLY_TEMPERATURE:
        # Match the behavior of _temperature_kernel: if that kernel uses
        # tl.div_rn, this must too.
        logits = logits / temp

    # fp32 is the default reduction dtype; fp64 is ~1/32-1/64x the throughput
    # on H100/Ada/Blackwell and empirically indistinguishable for Gumbel-max.
    if USE_FP64:
        logits = logits.to(tl.float64)
    if temp != 0.0:
        if IS_DRAFTING:
            pos = pos + _DRAFT_NOISE_SALT
        if USE_FP64:
            u = murmur3_uniform64(seed, pos, keys)
            gumbel_noise = -tl.log(-tl.log(u))
        else:
            u = murmur3_uniform32(seed, pos, keys)
            # Draw the large-noise tail (which decides the argmax winner) from
            # u -> 0, where fp32 has fine resolution. Avoid backend-specific
            # log1p while preserving precision in the winning tail.
            gumbel_noise = -tl.log(-_log1p_neg_stable(u))
        logits = tl.where(mask, logits + gumbel_noise, float("-inf"))

    return tl.max(logits, axis=0, return_indices=True)