Prhub

#41162 [Model Runner V2] Rebuild attn metadata between draft decode steps

原始 PR 作者 TheEpicDolphin 合并时间 2026-05-05 08:44 文件变更 3 提交数 2 评论 5 代码增减 +198 / -121

执行摘要

修复 MTP>2 时 DSV4 因缺少注意力元数据重建导致的崩溃

修复 DeepSeek V4 在 MTP > 2 时出现的 'invalid memory access' 崩溃。根因是 draft 解码步骤间未更新注意力元数据中的位置依赖状态。任何注意力后端都可能受此影响。

建议精读。该 PR 巧妙解决了 CUDA Graph 静态性与动态 metadata 重建的矛盾,通过标量张量传递步骤信息,并拆分单步生成逻辑,设计思路值得借鉴。同时包含性能基准数据,对理解推测解码性能特性有帮助。

讨论亮点
  • gemini-code-assist[bot] 指出 multi_step_decodeattn_metadata 在 FULL CUDA 图模式下被忽略,因为图捕获的是静态输入。PR 作者未回应,但 WoosukKwon 批准,说明此风险可接受。
  • gemini-code-assist[bot] 指出隐藏状态从融合内核移至 torch.copy_ 会引入额外内核启动开销。PR 未对此回应,但 WoosukKwon 批准,表示可接受。
  • WoosukKwon 询问是否应跳过最后一个解码步骤的 update_eagle_draft_inputs。TheEpicDolphin 解释由于 CUDA Graph 限制无法跳过,但额外调用因 Kernel Launch 被图重放覆盖而不会增加延迟。WoosukKwon 批准。

实现拆解

  1. EagleSpeculator.__init__ 中新增 self.current_draft_step 标量张量和 self.arange 辅助张量。
  2. 新增 multi_step_decode 方法,在循环 1..num_speculative_steps-1 中:重建 attn_metadata(若未跳过),更新 current_draft_step,然后调用捕获的单步生成或 generate_draft
  3. 重构 _sample_draft,使其接受 draft_step 张量和 draft_logits 张量,通过 gumbel_sampleoutput_processed_logits_col 写入正确列。
  4. 修改 gumbel_sample 及其 Triton 内核,新增 processed_logits_col_ptr 参数,支持在 CUDA Graph 内动态指定输出列。
  5. 更新 probabilistic_rejection_sampler_utils.py 中的 _resample_kernel,适应新的参数签名。
文件 模块 状态 重要度
vllm/v1/worker/gpu/spec_decode/eagle/speculator.py 推测解码 modified 8.65
vllm/v1/worker/gpu/sample/gumbel.py 采样器 modified 6.17
vllm/v1/worker/gpu/spec_decode/probabilistic_rejection_sampler_utils.py 拒绝采样 modified 5.02

关键符号

multi_step_decode _sample_draft generate_draft update_eagle_inputs update_eagle_draft_inputs gumbel_sample _gumbel_sample_kernel gumbel_block_argmax _resample_kernel

关键源码片段

vllm/v1/worker/gpu/spec_decode/eagle/speculator.py core-logic

核心变更文件:重构生成草稿的流程,新增 multi_step_decode 方法以实现步骤间注意力元数据重建,引入 current_draft_step 标量张量支持 CUDA Graph 的动态写入,并重构 generate_draft 为单步逻辑。

# vllm/v1/worker/gpu/spec_decode/eagle/speculator.py
# multi_step_decode 方法在 draft 多步解码循环中运行def multi_step_decode(
    self,
    num_reqs: int,
    skip_attn: bool,
    batch_desc: BatchExecutionDescriptor,
    num_tokens_across_dp: torch.Tensor | None,
) -> None:
    positions = self.input_buffers.positions[:num_reqs]
    query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1]
    idx_mapping = self.idx_mapping[:num_reqs]
​
    for step in range(1, self.num_speculative_steps):
        # 为每个 draft 解码步骤重建注意力元数据
        attn_metadata = None
        slot_mappings_by_layer = None
        if not skip_attn:
            # 即使在全图重放时也必需重建,以确保注意力元数据构建器的状态被更新
            slot_mappings = self.block_tables.compute_slot_mappings(
                idx_mapping, query_start_loc, positions, batch_desc.num_tokens,
            )
            slot_mappings_by_layer = build_slot_mappings_by_layer(
                slot_mappings, self.kv_cache_config
            )
            attn_metadata = self._build_draft_attn_metadata(
                num_reqs=num_reqs,
                num_reqs_padded=batch_desc.num_reqs or num_reqs,
                num_tokens_padded=batch_desc.num_tokens,
            )
​
        # 更新当前 draft 步骤标量张量,供 CUDA Graph 内使用
        self.current_draft_step.fill_(step)
​
        # 生成当前步骤的 draft tokens
        if batch_desc.cg_mode == CUDAGraphMode.FULL:
            assert self.decode_cudagraph_manager is not None
            self.decode_cudagraph_manager.run_fullgraph(batch_desc)
        else:
            self.generate_draft(
                num_reqs,
                batch_desc.num_tokens,
                attn_metadata,
                slot_mappings_by_layer,
                num_tokens_across_dp=num_tokens_across_dp,
                cudagraph_runtime_mode=batch_desc.cg_mode,
            )
vllm/v1/worker/gpu/sample/gumbel.py core-logic

支持通过列张量指定输出的列偏移,以便在 CUDA Graph 中动态选择写入目标。

# vllm/v1/worker/gpu/sample/gumbel.py
# gumbel_sample 函数新增 output_processed_logits_col 参数def gumbel_sample(
    logits: torch.Tensor, # [num_tokens, vocab_size]
    expanded_idx_mapping: torch.Tensor, # [num_tokens]
    temperature: torch.Tensor, # [max_num_reqs]
    seed: torch.Tensor, # [max_num_reqs]
    pos: torch.Tensor, # [num_tokens]
    apply_temperature: bool,
    output_processed_logits: torch.Tensor | None = None,
    output_processed_logits_col: torch.Tensor | None = None,
) -> torch.Tensor:
    num_tokens, vocab_size = logits.shape
    BLOCK_SIZE = 1024
    num_blocks = triton.cdiv(vocab_size, BLOCK_SIZE)
    local_argmax = logits.new_empty(num_tokens, num_blocks, dtype=torch.int64)
    local_max = logits.new_empty(num_tokens, num_blocks, dtype=torch.float64)
    # 将列张量作为参数传入内核,内核内部根据其值动态计算写入偏移
    _gumbel_sample_kernel[(num_tokens, num_blocks)](
        local_argmax, local_argmax.stride(0),
        local_max, local_max.stride(0),
        output_processed_logits,
        output_processed_logits.stride(0) if output_processed_logits is not None else 0,
        output_processed_logits_col, # 新增:列索引张量
        logits, logits.stride(0),
        expanded_idx_mapping,
        seed, pos, temperature,
        vocab_size,
        BLOCK_SIZE,
        apply_temperature,
    )
    # ... 后续 argmax 合并逻辑不变 ...

评论区精华

FULL CUDA 图模式下 attn_metadata 重建不生效的风险 正确性

gemini-code-assist[bot] 指出 multi_step_decode 中重新赋值的 attn_metadata 在 FULL CUDA 图模式下被忽略,因为图捕获的是静态输入。

结论:PR 作者未直接回应,但 WoosukKwon 批准,表明当前设计在 FULL 模式下使用捕获的静态输入,无需动态重建。 · 已解决

隐藏状态分离 copy 带来的额外内核启动开销 性能

gemini-code-assist[bot] 指出之前隐藏状态更新融合在 Triton 内核中,现在独立为 torch.copy_ 会增加内核启动次数,影响延迟。

结论:未进一步讨论,但 PR 被批准,表示可接受。 · 已解决

是否应跳过最后一个草稿解码步骤的 update 设计

WoosukKwon 询问是否应该跳过最后一个解码步骤的 update_eagle_draft_inputs。TheEpicDolphin 解释由于 CUDA Graph 限制无法跳过,但额外调用不会增加延迟(Kernel Launch 已被图重放覆盖)。

结论:决定不跳过,因为开销已被分摊。 · 已解决

风险与影响

  1. 性能风险:在每次解码步骤前重建注意力元数据增加开销,但基准测试显示总吞吐量未明显退化(部分场景有提升)。分离的 torch.copy_ 可能增加延迟,但 CUDAGraph 下可分摊。
  2. 正确性风险:注意力元数据重建可能引入其他后端的回归,但只影响 V2 Model Runner 和 Eagle 推测解码。
  3. CUDA Graph 兼容性:新增的 current_draft_step 是标量张量,兼容图捕获。但 attn_metadata 在 FULL 模式下可能仍无法动态更新,需要确保路径正确。

用户:修复 DSV4 在 MTP>2 时的崩溃,使功能可用。系统:无 API 变更,内部逻辑重构影响推测解码性能(变化<6%)。团队:为后续 Model Runner V2 的推测解码改进奠定基础。

CUDA Graph 兼容性 性能退化风险 核心模块重构

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论