From 50cfbbca7aed14e82e991a0d89bf1d5cced99491 Mon Sep 17 00:00:00 2001 From: Woosuk Kwon Date: Fri, 3 Jul 2026 06:12:33 +0000 Subject: [PATCH] [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 --- tests/v1/worker/test_gpu_gumbel_sample.py | 37 +++++++++++++++++- vllm/v1/worker/gpu/sample/gumbel.py | 47 +++++++++++++---------- 2 files changed, 63 insertions(+), 21 deletions(-) diff --git a/tests/v1/worker/test_gpu_gumbel_sample.py b/tests/v1/worker/test_gpu_gumbel_sample.py index 9db175113ce..6bd8fcecb2f 100644 --- a/tests/v1/worker/test_gpu_gumbel_sample.py +++ b/tests/v1/worker/test_gpu_gumbel_sample.py @@ -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 ---------------------------------- diff --git a/vllm/v1/worker/gpu/sample/gumbel.py b/vllm/v1/worker/gpu/sample/gumbel.py index 190307d5e75..a49830b28b2 100644 --- a/vllm/v1/worker/gpu/sample/gumbel.py +++ b/vllm/v1/worker/gpu/sample/gumbel.py @@ -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.