Prhub

#26888 [KDA] Add target_verify support for speculative decoding

原始 PR 作者 yyq0210 合并时间 2026-07-25 19:52 文件变更 6 提交数 7 评论 39 代码增减 +701 / -4

执行摘要

为 KDA 添加 EAGLE 推测解码 target_verify 支持

GDN (a sister linear attention backend) already supports full EAGLE speculative decoding. This PR achieves feature parity for KDA with specific adaptations, including is_kda=True gating in fused sigmoid recurrent kernel, convolution state transpose, and full convolution weights applied to combined mixed_qkv.

该PR值得精读,尤其是_detect_conv_window_axis的设计体现了同一代码库支持多种Conv Layout的策略。合并后可作为KDA推测解码的基础。建议同时阅读相关PR #30113和#28197了解完整上下文。

讨论亮点
  • extra_buffer支持范围:YAMY1234指出测试文档中使用了extra_buffer,但当前分支启动时会断言失败,建议改为no_buffer。作者已修正文档和测试配置。
  • mamba跟踪开销:YAMY1234指出KDAAttnBackend.init_forward_metadata重写计算了mamba_track_mask_indices等,但KDA后端不消费,且引入设备同步(.nonzero()),建议移除。作者同意并删除了该重写。
  • 卷积窗口轴检测:YAMY1234指出原dedup分配器假定GDN的(dim, K-1)布局,KDA使用(K-1, dim)导致视图构建错误。作者引入_detect_conv_window_axis自动检测轴位置,并通过is_kda参数使KDA使用密集布局(不参与dedup)。
  • 与PR #28197统一:yuan-luo要求与并行PR #28197统一,作者rebase后解决了冲突,合并了功能。

实现拆解

  1. KDA内核target_verify入口:在Triton内核fused_sigmoid_gating_delta_rule_update中添加target_verify支持,通过disable_state_update=Trueintermediate_states_buffer参数,使单次前向即可验证多个draft token而不修改SSM状态。
  2. KDAAttnBackend适配:新增target_verify方法委托给verify_kernel.target_verify;修改forward_extend,当forward_mode.is_target_verify()时跳过gate激活(内核内部处理),处理卷积状态转置((conv_width, qkv_dim) → (qkv_dim, conv_width))并传递中间状态缓存。
  3. 内存池布局感知:新增_detect_conv_window_axis函数自动检测卷积窗口轴顺序(优先GDN尾轴布局),支持KDA的(K-1, dim)布局,并通过conv_window_dedup_enabled(..., is_kda=True)使KDA保持密集布局,避免溢出。
  4. 模型forward调整:在KimiLinearForCausalLM.forward中,当forward_modeis_target_verify时也不进行gate激活,与decode模式一致。
  5. 测试覆盖:新增test_kda_target_verify.py验证kernel等价性(fp32精确匹配,bf16误差<1e-3);test_ngram_mamba_verify_update.py测试commit_mamba_states_after_verify正确性;test_kda_spec_integration.py端到端验证正常推理、prefix caching和batch推理无回归。
文件 模块 状态 重要度
python/sglang/srt/mem_cache/memory_pool.py 内存池 modified 7.02
python/sglang/srt/layers/attention/linear/kda_backend.py 注意力后端 modified 5.5
python/sglang/srt/models/kimi_linear.py 模型加载 modified 5.92
test/registered/unit/spec/test_ngram_mamba_verify_update.py 状态验证测试 added 8.14
test/manual/test_kda_target_verify.py KDA 内核测试 added 7.44
test/manual/test_kda_spec_integration.py KDA 集成测试 added 7.54

关键符号

_detect_conv_window_axis KDAAttnBackend.target_verify KDAAttnBackend.forward_extend KimiLinearForCausalLM.forward

关键源码片段

python/sglang/srt/mem_cache/memory_pool.py core-logic

新增关键函数 `_detect_conv_window_axis` 以自动检测卷积窗口轴顺序,支持 KDA 的 (K-1, dim) 布局与 GDN 的 (dim, K-1) 布局共存,确保 deduplicated sliding-window view 正确构建。

def _detect_conv_window_axis(
    self, conv_state_shape: List[Tuple[int, int]], win_len: int
) -> int:
    """
    自动检测卷积窗口轴位置。
    GDN 的 conv_state 形状为 (dim, K-1),尾轴长度为 K-1;
    KDA 的形状为 (K-1, dim),首轴长度为 K-1。
    优先选择 GDN 的尾轴布局,若所有层一致则返回检测到的轴。
    """
    axis = None
    for conv_shape in conv_state_shape:
        # 检查尾轴是否匹配卷积大小
        if conv_shape[-1] == win_len:
            shape_axis = len(conv_shape) - 1 # GDN 布局
        elif conv_shape[0] == win_len:
            shape_axis = 0 # KDA 布局
        else:
            raise ValueError(
                f"conv_state shape {conv_shape} 没有长度为 win_len={win_len} 的轴"
            )
        if axis is None:
            axis = shape_axis
        elif axis != shape_axis:
            raise ValueError(
                f"各层卷积窗口轴不一致: {conv_state_shape},无法共享 buffer"
            )
    return axis# 在 _allocate_deduplicated_conv_window 中使用该轴构建物理形状和 as_strided view
python/sglang/srt/layers/attention/linear/kda_backend.py core-logic

核心后端变更:新增 `target_verify` 方法并将验证委托给 verify_kernel;修改 `forward_extend` 添加 is_target_verify 分支,处理卷积状态转置和中间 SSM 状态缓存,是 speculative decoding 的关键执行路径。

class KDAAttnBackend(...):
​
    def target_verify(
        self,
        A_log, dt_bias, q, k, v, a, b,
        *, ssm_states, cache_indices, query_start_loc, **kwargs,
    ) -> torch.Tensor:
        """验证多个 draft token 的 SSM 状态,不修改原始状态。"""
        return self.verify_kernel.target_verify(
            A_log, dt_bias, q, k, v, a, b,
            initial_state_source=ssm_states,
            initial_state_indices=cache_indices,
            cu_seqlens=query_start_loc,
            use_qk_l2norm_in_kernel=True,
            softplus_beta=1.0,
            softplus_threshold=20.0,
            is_kda=True,
            disable_state_update=True,
            intermediate_states_buffer=..., # 外部传入 buffer
            intermediate_state_indices=...,
            cache_steps=...,
            retrieve_parent_token=None,
        )
​
    def forward_extend(self, ...):
        # ... 其他逻辑
        if forward_batch.forward_mode.is_target_verify():
            # 跳过 gate 激活,target_verify 内核内部处理
            # KDA 的 conv_state 形状为 (K-1, dim),需要转置为 (dim, K-1) 以匹配
            # causal_conv1d_update 的期望布局
            conv_state = conv_state.transpose(-2, -1).contiguous()
            output = self.verify_kernel.target_verify(
                ...,
                intermediate_states_buffer=intermediate_ssm_buffer,
                cache_steps=spec_info.draft_token_num,
            )
        else:
            # 正常 extend 路径
            # ...

评论区精华

extra_buffer 支持范围与文档修正 正确性

YAMY1234 指出测试文档中 extra_buffer 模式在当前分支不支持(启动断言失败),建议改为 no_buffer。作者已修正文档。

结论:文档已修正,测试配置改为 no_buffer。 · 已解决

KDAAttnBackend 中 mamba 跟踪代码开销与可维护性 性能

YAMY1234 指出 init_forward_metadata 重写计算了 mamba_track_mask_indices 等,但 KDA 后端不消费,且引入设备同步(.nonzero()),建议移除。作者同意并删除。

结论:已移除重写函数,留下注释说明。 · 已解决

卷积窗口轴检测与 KDA 布局兼容性 设计

YAMY1234 指出内存池中 dedup 滑动窗口分配器原假定 GDN 的 (dim, K-1) 布局,KDA 使用 (K-1, dim) 导致 view 构建错误。引入 _detect_conv_window_axis 自动检测轴,并通过 is_kda 参数使 KDA 保持密集布局(不参与 dedup)。

结论:已实现 _detect_conv_window_axis,KDA 密集布局,GDN 使用 dedup。 · 已解决

与重复 PR #28197 的统一 other

yuan-luo 要求与并行 PR #28197 统一,作者完成 rebase 并解决冲突,合并了功能。

结论:已统一为一个版本。 · 已解决

风险与影响

  1. 核心路径变更:修改了KDAAttnBackend.forward_extendKimiLinearForCausalLM.forward,可能影响非推测解码路径的正常推理,但通过e2e测试验证无回归。
  2. 功能重叠风险:本PR功能与已合并的PR #30113存在部分重叠,rebase时已消除重复的target_verify定义和分发逻辑,但可能仍存在隐式依赖。
  3. 性能风险:移除了extend路径中不必要的.nonzero()同步,当前代码无额外同步开销;target_verify路径批处理draft token,预期提升kernel效率。
  4. 测试覆盖不足:kernel级测试充分,但缺少CUDA graph和radix cache结合的手动测试(CI中未覆盖),可能遗漏rollback相关bug。

对用户:KimiLinearForCausalLM现在可以启用推测解码(如--speculative-algorithm NGRAM),提升推理吞吐。对系统:新增约700行代码(测试570+源码130),改动集中在KDA后端和内存池,影响范围有限。对团队:引入KDA专用的conv布局检测逻辑,需持续维护与GDN的差异。

核心路径变更 功能重叠风险 CUDA graph 测试未覆盖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论