From c53994e1348bac3496aafb88e9e731124a00a8a7 Mon Sep 17 00:00:00 2001 From: Giancarlo Delfin <32987265+TheEpicDolphin@users.noreply.github.com> Date: Thu, 25 Jun 2026 18:46:10 -0500 Subject: [PATCH] [Model Runner V2][Spec Decode] Use log1p to compute residual during rejection sampling (#46665) Signed-off-by: Giancarlo Delfin --- .../gpu/spec_decode/rejection_sampler_utils.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/vllm/v1/worker/gpu/spec_decode/rejection_sampler_utils.py b/vllm/v1/worker/gpu/spec_decode/rejection_sampler_utils.py index bad70aa0451..7020f228046 100644 --- a/vllm/v1/worker/gpu/spec_decode/rejection_sampler_utils.py +++ b/vllm/v1/worker/gpu/spec_decode/rejection_sampler_utils.py @@ -2,7 +2,7 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import torch -from vllm.triton_utils import tl, triton +from vllm.triton_utils import tl, tldevice, triton from vllm.v1.worker.gpu.sample.gumbel import gumbel_block_argmax, tl_rand64 @@ -387,14 +387,16 @@ def _resample_kernel( draft_lse = tl.load(draft_rejected_logsumexp_ptr + req_idx) target_log_probs = target_logits - target_lse draft_log_probs = draft_logits - draft_lse - # Compute the residual: max(p(x) - q(x), 0) - # Equivalent log form: log(max(exp(log_p(x)) - exp(log_q(x)), 0)) + # Compute the residual: + # r(x) = max(p(x) - q(x), 0) + # Gumbel sampling needs logits, so we compute it in log space: + # log(r(x)) = log(max(exp(log_p(x)) - exp(log_q(x)), 0)) # The more numerically stable form is: - # log(max(exp(a) - exp(b), 0)) = a + log(max(1 - exp(b - a), 0)) + # log(max(exp(a) - exp(b), 0)) = a + log(max(1 - exp(b - a), 0)) ratio = tl.exp(draft_log_probs - target_log_probs) residual_logits = tl.where( ratio < 1.0, - target_log_probs + tl.log(1 - ratio), + target_log_probs + tldevice.log1p(-ratio), float("-inf"), ).to(tl.float32) else: