Compare commits

...
Author SHA1 Message Date
Woosuk KwonandGitHub cc934277a9 Merge branch 'main' into woosuk/gumbel-fp32-rand64 2026-07-08 09:35:06 -07:00
Woosuk KwonandGitHub d41f6eb3e2 Merge branch 'main' into woosuk/gumbel-fp32-rand64 2026-07-06 22:32:44 -07:00
Woosuk KwonandGitHub 8c751ad3f5 Merge branch 'main' into woosuk/gumbel-fp32-rand64 2026-07-05 10:03:29 -07:00
Woosuk KwonandGitHub 1bc6183344 Merge branch 'main' into woosuk/gumbel-fp32-rand64 2026-07-04 09:16:59 -07:00
Woosuk KwonandGitHub e9df379765 Merge branch 'main' into woosuk/gumbel-fp32-rand64 2026-07-03 15:03:24 -07:00
Woosuk KwonandClaude 50cfbbca7a [MRV2] Draw 64-bit uniforms for fp32 Gumbel sampling
tl.rand only resolves the uniform down to 2**-31, which hard-caps the
fp32 Gumbel noise at ~21.5 and coarsely buckets the argmax-deciding
tail. Philox already generates 128 bits per call (tl.rand discards 96),
so draw 64 bits instead and convert uint64 -> fp32 directly: the tail
now resolves down to 2**-64 (noise up to ~44.4) with no fp64 arithmetic
and no measurable cost. The fp64 path is unchanged.

Co-authored-by: Claude
Signed-off-by: Woosuk Kwon <woosuk@inferact.ai>
2026-07-03 06:22:59 +00:00
2 changed files with 63 additions and 21 deletions
+36 -1
View File
@@ -21,7 +21,8 @@ pytest.importorskip("triton")
if not torch.cuda.is_available():
pytest.skip("CUDA required for Gumbel sampler tests", allow_module_level=True)
from vllm.v1.worker.gpu.sample.gumbel import gumbel_sample
from vllm.triton_utils import tl, triton
from vllm.v1.worker.gpu.sample.gumbel import gumbel_sample, tl_rand32, tl_rand64
DEVICE = "cuda"
VOCAB_SIZE = 200_000
@@ -164,6 +165,40 @@ def test_full_vocab_distribution_fidelity():
assert chi2 < df + 10 * math.sqrt(2 * df), f"chi2={chi2:.0f}, df={df}"
# ----------------------------- RNG precision --------------------------------
@triton.jit
def _draw_uniform_kernel(offset_ptr, out32_ptr, out64_ptr, seed, N: tl.constexpr):
idx = tl.arange(0, N)
offs = tl.load(offset_ptr + idx)
u32 = tl_rand32(seed, offs, includes_zero=False)
u64 = tl_rand64(seed, offs, includes_zero=False)
tl.store(out32_ptr + idx, u32)
tl.store(out64_ptr + idx, u64)
def test_rand32_resolves_below_tl_rand_floor():
"""`tl_rand32` draws 64 random bits, so its u -> 0 tail resolves below
`tl.rand`'s 2**-31 floor, and it must agree with `tl_rand64` up to fp32
rounding. The offsets are draws from the seed=12345 Philox stream found
(by scan) to fall below 2**-31; they are impossible for a 31-bit uniform.
"""
seed = 12345
offsets = torch.tensor(
[4982566788, 5277073014, 5357046532, 12285768576],
dtype=torch.int64,
device=DEVICE,
)
u32 = torch.empty(4, dtype=torch.float32, device=DEVICE)
u64 = torch.empty(4, dtype=torch.float64, device=DEVICE)
_draw_uniform_kernel[(1,)](offsets, u32, u64, seed, N=4)
assert (u32 > 0).all()
assert (u32 < 2.0**-31).all(), f"u32={u32.tolist()}"
assert torch.equal(u32, u64.float()), f"u32={u32.tolist()}, u64={u64.tolist()}"
# ----------------------------- Edge cases ----------------------------------
+27 -20
View File
@@ -2,16 +2,7 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import torch
from vllm.triton_utils import HAS_TRITON, tl, tldevice, triton
# Smallest positive value produced by Triton's fp32 `tl.rand`. Used to clamp
# zero draws before the flipped Gumbel transform below.
#
# Triton requires globals accessed from `@triton.jit` functions to be wrapped
# in `tl.constexpr(...)`. We can only do that when Triton is actually
# available — on the CPU worker path `tl` is a placeholder whose `constexpr`
# attribute is `None`, and `tl.constexpr(...)` would crash at import time.
_TL_RAND_MIN = tl.constexpr(4.6566127342e-10) if HAS_TRITON else 4.6566127342e-10
from vllm.triton_utils import tl, tldevice, triton
@triton.jit
@@ -59,12 +50,18 @@ def apply_temperature(
@triton.jit
def tl_rand64(seed, offset, includes_zero: tl.constexpr):
def _tl_randbits64(seed, offset):
# Philox generates 128 bits per call and `tl.rand` discards 96 of them,
# so drawing 64 bits costs the same RNG work as drawing 32.
lo, hi, _, _ = tl.randint4x(seed, offset)
lo = lo.to(tl.uint32, bitcast=True).to(tl.uint64)
hi = hi.to(tl.uint32, bitcast=True).to(tl.uint64)
r = (hi << 32) | lo
return (hi << 32) | lo
@triton.jit
def tl_rand64(seed, offset, includes_zero: tl.constexpr):
r = _tl_randbits64(seed, offset)
# 1 / 2**64
scale = 5.421010862427522170037e-20
u = r.to(tl.float64) * scale
@@ -75,9 +72,17 @@ def tl_rand64(seed, offset, includes_zero: tl.constexpr):
@triton.jit
def tl_rand32(seed, offset, includes_zero: tl.constexpr):
u = tl.rand(seed, offset)
# Same 64-bit stream as `tl_rand64`, converted to fp32. The uint64 ->
# fp32 convert rounds to nearest, which is exact for small values, so the
# u -> 0 tail resolves down to 2**-64 instead of `tl.rand`'s 2**-31 floor
# (both far above fp32's 2**-126 minimum normal). Elsewhere the rounding
# only loses bits fp32 cannot hold anyway.
r = _tl_randbits64(seed, offset)
# 1 / 2**64
scale = 5.421010862427522170037e-20
u = r.to(tl.float32) * scale
if not includes_zero:
u = tl.maximum(u, _TL_RAND_MIN)
u = tl.maximum(u, 1.1754943508222875e-38) # float32 tiny
return u
@@ -143,12 +148,14 @@ def gumbel_block_argmax(
gumbel_noise = -tl.log(-tl.log(u))
else:
u = tl_rand32(gumbel_seed, block, includes_zero=False)
# Draw the large-noise tail (which decides the argmax winner) from u -> 0,
# where fp32 has fine resolution, instead of u -> 1, where fp32 spacing is
# ~2**-24. The naive `-log(-log(u))` puts the winning tail at u -> 1,
# hard-capping the noise at ~16.6 and coarsely quantizing it; using
# `log1p(-u)` == `log(1 - u)` keeps the tail in the well-resolved region.
# Note `1 - u` would lose precision for small u, so `log1p` is required.
# Draw the large-noise tail (which decides the argmax winner) from
# u -> 0, where fp32 has fine resolution, instead of u -> 1, where
# fp32 spacing is ~2**-24. The naive `-log(-log(u))` puts the
# winning tail at u -> 1, hard-capping the noise at ~16.6 and
# coarsely quantizing it; with `log1p(-u)` == `log(1 - u)` the
# tail resolves down to u = 2**-64, i.e. noise up to ~44.4. Note
# `1 - u` would lose precision for small u, so `log1p` is
# required.
gumbel_noise = -tl.log(-tldevice.log1p(-u))
# Apply gumbel noise.