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: