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)
|