Prhub

#46781 [Model Runner V2][Spec Decode] Implement block verification for rejection sampling

原始 PR 作者 TheEpicDolphin 合并时间 2026-06-30 23:07 文件变更 4 提交数 1 评论 3 代码增减 +641 / -85

执行摘要

实现块验证拒绝采样,提升推测解码接受率

当前 Model Runner V2 不支持块验证拒绝采样,但 V1 正在支持(#40819)。块验证被期望产生至少与标准拒绝采样一样高的接受率,且成本相当。实现此功能可以为推测解码场景带来更高的效率和吞吐量。

该 PR 值得精读,特别是对推测解码和 Triton 内核开发感兴趣的工程师。实现展示了如何基于论文在推理引擎中集成复杂采样算法,并附带了完整的基准测试和测试验证。

讨论亮点

该 PR 没有引发实质性的技术讨论。WoosukKwon 在审查后批准了 PR,表示 'LGTM. Thanks for doing this!'。此前 mergify[bot] 提示存在合并冲突(后已解决)。

实现拆解

  1. 配置扩展:在 vllm/config/speculative.py 中,RejectionSampleMethod 类型别名新增 block 字面量,并在文档字符串中说明。
  2. 采样器集成:在 vllm/v1/worker/gpu/spec_decode/rejection_sampler.pyRejectionSampler.__init__ 中根据配置设置 self.use_block_verification,并在 __call__ 中传递给 rejection_sample
  3. 核心 Triton 内核:在 vllm/v1/worker/gpu/spec_decode/rejection_sampler_utils.py 中重命名辅助函数,新增 _compute_global_residual_mass_compute_global_target_argmax_compute_global_logprobs_and_logsumexp_compute_cumulative_log_p_kernel_compute_local_residual_mass_kernel 等内核,并重写 _compute_block_stats_kernel 以支持控制流分支。
  4. 测试覆盖:在 tests/v1/spec_decode/test_rejection_sampler_utils.py 中添加两个测试,验证分布正确性和接受长度担保。
文件 模块 状态 重要度
vllm/v1/worker/gpu/spec_decode/rejection_sampler_utils.py 推测解码 modified 8.93
vllm/v1/worker/gpu/spec_decode/rejection_sampler.py 推测解码 modified 5.73
vllm/config/speculative.py 配置层 modified 5.31
tests/v1/spec_decode/test_rejection_sampler_utils.py 测试覆盖 modified 6.7

关键符号

_compute_block_max_and_sumexp _compute_max_and_sumexp _compute_global_lse _compute_global_logsumexp _compute_block_stats_kernel _compute_global_residual_mass _compute_global_target_argmax _compute_global_logprobs_and_logsumexp _compute_cumulative_log_p_kernel _compute_local_residual_mass_kernel test_block_verification_rejection_sample test_block_verification_accepts_at_least_as_many RejectionSampler.__init__ RejectionSampler.__call__

关键源码片段

vllm/v1/worker/gpu/spec_decode/rejection_sampler_utils.py core-logic

核心实现文件,包含所有块验证相关的 Triton 内核和算法逻辑。

下面是块验证核心内核 _compute_global_residual_mass 的实现。该内核计算残差归一化因子 Z_i,用于接受阈值和重采样分布。

@triton.jit
def _compute_global_residual_mass(
    local_residual_mass_ptr,
    local_residual_mass_stride,
    prefix_joint_ratio,
    target_logits_ptr,
    target_logits_stride,
    target_local_max_ptr,
    target_local_max_stride,
    target_local_sumexp_ptr,
    target_local_sumexp_stride,
    draft_sampled_ptr,
    logit_idx,
    vocab_num_blocks,
    PADDED_VOCAB_NUM_BLOCKS: tl.constexpr,
    HAS_DRAFT_LOGITS: tl.constexpr,
):
    if HAS_DRAFT_LOGITS:
        # 完整 draft 分布:累加各子块残差质量
        blocks = tl.arange(0, PADDED_VOCAB_NUM_BLOCKS)
        mask = blocks < vocab_num_blocks
        partials = tl.load(
            local_residual_mass_ptr + logit_idx * local_residual_mass_stride + blocks,
            mask=mask,
            other=0.0,
        )
        return tl.sum(partials, axis=0)
    else:
        # one-hot(贪婪)draft:M_s 是 draft_token 上的点质量
        # 残差质量 = prefix_joint_ratio * (1 - M_b(draft_token))
        draft_token = tl.load(draft_sampled_ptr + logit_idx + 1).to(tl.int64)
        target_lse = _compute_global_logsumexp(
            target_local_max_ptr,
            target_local_max_stride,
            target_local_sumexp_ptr,
            target_local_sumexp_stride,
            logit_idx,
            vocab_num_blocks,
            PADDED_VOCAB_NUM_BLOCKS,
        )
        target_logit = tl.load(
            target_logits_ptr + logit_idx * target_logits_stride + draft_token,
        ).to(tl.float32)
        m_b = tl.exp(target_logit - target_lse)
        return prefix_joint_ratio * (1.0 - m_b)
vllm/v1/worker/gpu/spec_decode/rejection_sampler.py core-logic

采样器类,将配置标志传递到核心采样函数。

下面是采样器类的接口变更,展示了 use_block_verification 标志的设置和传递。

class RejectionSampler:
    def __init__(
        self,
        sampler: Sampler,
        spec_config: SpeculativeConfig,
        device: torch.device,
    ):
        self.sampler = sampler
        self.num_speculative_steps = spec_config.num_speculative_tokens
        rejection_sample_method = spec_config.rejection_sample_method
        self.use_block_verification: bool = False
        self.synthetic_conditional_rates: torch.Tensor | None = None
        if rejection_sample_method == "synthetic":
            # 合成拒绝采样:使用预设的接受率
            assert spec_config.synthetic_acceptance_rates is not None
            self.synthetic_conditional_rates = torch.tensor(
                unconditional_to_conditional_rates(
                    spec_config.synthetic_acceptance_rates
                ),
                dtype=torch.float32,
                device=device,
            )
        elif rejection_sample_method == "block":
            # 启用块验证
            self.use_block_verification = True
​
    def __call__(
        self,
        logits: torch.Tensor,
        input_batch: InputBatch,
        draft_logits: torch.Tensor | None = None,
    ) -> SamplerOutput:
        ...
        sampled, num_sampled = rejection_sample(
            processed_logits,
            draft_logits,
            draft_sampled,
            input_batch.cu_num_logits,
            pos,
            input_batch.idx_mapping,
            input_batch.expanded_idx_mapping,
            input_batch.expanded_local_pos,
            self.sampler.sampling_states.temperature.gpu,
            self.sampler.sampling_states.seeds.gpu,
            self.num_speculative_steps,
            self.synthetic_conditional_rates,
            use_fp64=self.sampler.use_fp64_gumbel,
            use_block_verification=self.use_block_verification, # 传递块验证标志
        )
        ...

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

主要风险集中在 vllm/v1/worker/gpu/spec_decode/rejection_sampler_utils.py 中的新 Triton 内核。数值计算(指数、对数)可能引入精度问题,尤其是在低精度模式下。块验证算法新增控制流路径,可能影响标准模式和合成模式的正确性。需要确保新内核与 GPU 架构的兼容性(已在 CUDA 上测试)。

对用户:提供 block 作为 rejection_sample_method 的新值,用户可通过配置启用。对系统:新增一个或两个额外的 Triton 内核启动,但计算量相对小,且可能减少拒绝步骤,从而提升总体吞吐量。对开发团队:需要维护额外的采样方法,文档和测试已更新。

核心路径变更 数值精度 新增 API 配置

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论