Prhub

#51855 [K3] support recoverssm for K3

原始 PR 作者 ZJY0516 合并时间 2026-08-17 19:34 文件变更 19 提交数 15 评论 11 代码增减 +2235 / -75

执行摘要

Kimi-K3 推测解码新增 RecoverSSM 恢复路径,缓存容量 +10.97%

KDA 推测解码在验证阶段会为每个推测位置实例化一份完整 recurrent state,当 DP/EP 并行且每个 rank 有多个序列时,缓存浪费会被放大。作者在讨论中给出了 GB300 上的量化评估:DP + EP 16 时 KDA state 约 400 MB,若每 rank 4 个序列且每个序列 7 个推测 token,不用 RecoverSSM 会浪费约 12 GB 显存。PR body 明确指出:Without RecoverSSM, KDA speculative decoding materializes a full recurrent state for every speculative position,而 RecoverSSM 通过保留单一 checkpoint 和紧凑 per-token records,在每次 verify 后在接受位置重建状态,实现在相同缓存预算下有效容量提升 10.97%。

讨论亮点

核心讨论有三条线索:

  1. 性能权衡(benchislett)it looks like we're expecting ~3% overhead on the decode step to free up 10% of extra KV cache capacity,作者回应 DP 场景下浪费更严重(单 rank 4 序列 × 7 spec token 可浪费 12 GB)。benchislett 认为 commit kernel seems suboptimal at a glance,将转发给内核性能专家,但不阻塞合并。

  2. draftless decode 分类正确性(Codex P1 + benchislett 确认):Codex 指出零 draft 的 spec 请求若落入普通 decode 路径,会按扩展状态长度错误移位卷积窗口,导致后续步骤输出错误;benchislett 也确认 this bot comment is valid。该问题由后续 commit 19115f4a(Fix RecoverSSM draftless decode handling)修复。

  3. align block size 配置校验缺口(Codex P2):当 mamba_block_size < num_speculative_tokens + 1 时,配置校验放行但首次接受步运行时即崩溃,建议在配置阶段拒绝或让 commit plan 支持跨多边界窗口;材料中未见明确修复证据。

另有 benchislett 对架构的点评:I have tried to focus on moving complexity away from shared codepaths and into Kimi-specific and recoverssm-specific files,以及其对 cache.py--use-replayssm 语义注释的疑问(Kimi-K3 复用该 flag 但实现完全不同)。

实现拆解

实现按五步展开:

  1. 状态布局扩展vllm/model_executor/layers/mamba/mamba_utils.py 新增 append_kda_recoverssm_record,在 KDA 原有 (conv_state, recurrent_state) 基础上追加两个 transient 缓冲区——FP32 correction 记录 (local_num_heads, spec_query_len, head_dim) 和 BF16 key/gate 记录 (local_num_heads, spec_query_len, 2 * head_dim)vllm/models/kimi_k3/nvidia/kda.pyget_state_dtype/get_state_shape 在开启 use_kda_recoverssm 时扩展返回类型。这些恢复记录不参与 prefix-cache 状态拷贝,checkpoint 与卷积状态仍是缓存边界状态。

  2. Metadata 构建改造vllm/models/kimi_k3/nvidia/kda_metadata.pyKimiK3KDAMetadataBuilder 增加 use_recoverssm 分支:spec 行只保留 1 个 checkpoint slot(spec_state_slots = 1,原生 KDA 为 num_spec + 1);将 draftless decode 也纳入 spec 分类以保留扩展卷积窗口(对应 review 中 P1 问题的修复);新增 KDARecoverSSMAlignMetadata/KDARecoverSSMCommitMetadata 两个数据结构,并通过 _get_recoverssm_context 惰性创建共享的 KDARecoverSSMCommitContext

  3. Verify 与记录阶段:新增 vllm/models/kimi_k3/nvidia/ops/recoverssm.py(1067 行),kda_recoverssm_verify 只读 checkpoint 并按推测位置逐步计算 KDA 输出,同时把 correction u(FP32)与 key/raw gate(激活 dtype)写入记录缓冲,避免提交阶段重算状态依赖的 correction。vllm/models/kimi_k3/nvidia/kda.py_forwardhas_spec_decode 且开启 RecoverSSM 时改走该 verify 路径。

  4. 采样后提交与 align 后处理:新增 vllm/v1/worker/gpu/model_states/recoverssm.pyRecoverSSMState,在 mamba_hybrid.pyprepare_attnrecord_step 记录本步 metadata,在 postprocess_state 采样后 commit_step 调用 commit_recoverssm_state。提交阶段由 _prepare_commit_plan_kernel 计算运行时接受长度与 align 边界,_commit_kda_state_kernel 逆序一次性重建状态(非递归更新),_compact_conv_state_kernel 压缩卷积历史;align 模式下 _postprocess_recoverssm_align_kernel 更新运行块索引并把下一步拷贝偏置重置为中性值。

  5. 配置校验与测试配套vllm/config/vllm.pyvalidate_mamba_cached_kernel--use-replayssm 在 Kimi-K3 + 推测解码下映射为 use_kda_recoverssm,并校验架构、随机舍入开关、mamba_cache_mode(仅 none/align)、pipeline_parallel_size=1 等约束。测试方面 tests/models/kimi_k3/test_kda_metadata.py 覆盖单 slot、draftless decode 分类与 CUDA graph 捕获,tests/models/kimi_k3/test_kda.pytest_kda_recoverssm_verify_and_group_commit 以朴素递推为参考做精度对比(含 request-indexed、align 两种布局),tests/v1/worker/test_mamba_hybrid_model_state.pytests/v1/worker/test_mamba_utils.py 覆盖 v1 worker 集成。

文件 模块 状态 重要度
vllm/v1/worker/gpu/model_states/recoverssm.py vllm/v1 added 9.11
vllm/models/kimi_k3/nvidia/kda_metadata.py vllm/models modified 9.08
vllm/models/kimi_k3/nvidia/ops/recoverssm.py vllm/models added 8.02
vllm/v1/attention/backends/recoverssm_metadata.py vllm/v1 added 7.93
tests/models/kimi_k3/test_kda_metadata.py kda/metadata modified 7.87

关键符号

RecoverSSMState __init__ record_step commit_step _postprocess_recoverssm_align_kernel KimiK3KDAMetadata KDARecoverSSMAlignMetadata KDARecoverSSMCommitMetadata commit_recoverssm_state _get_recoverssm_context

分析完成后,这里会展示 LLM 生成的相对完整源码片段和详细注释。

评论区精华

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

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

风险与影响

  1. 正确性风险(align 边界):Codex P2 指出 mamba_block_size < num_speculative_tokens + 1 时配置校验未拦截,首个接受步会因 spec_query_len 超窗口运行时报错。vllm/config/vllm.pyvllm/models/kimi_k3/nvidia/ops/recoverssm.py 之间存在配置-运行时契约缺口,当前仅靠测试用例覆盖了等长场景。
  2. 核心路径回归风险vllm/v1/worker/gpu/model_states/mamba_hybrid.pypostprocess_state 被重构(num_reqs == 0 提前返回逻辑调整),同时 RecoverSSMState 挂接在共享的采样后处理路径上,任何 metadata 生命周期错配(如 CUDA graph 捕获与回放时 _step 状态)都可能影响非 Kimi-K3 模型。
  3. 精度风险:重建公式依赖 FP32 correction 与 BF16 key/gate,PR 声明 GSM8K 提升 0.15~0.23 pp,但置信区间 [-0.4549, +0.9098] 跨零,MRCR 分数差异(0.74879 vs 0.66791)是否来自状态重建偏差仍需更多验证。
  4. 性能风险:verify 阶段因 FP32 correction 写入增加 13.66% 耗时;commit kernel 在采样后运行、不在 model forward 内,每步额外约 1.711 ms,PR 估算 GPU 关键路径成本约 2.046 ms/step,benchislett 认为 commit kernel 仍有优化空间。
  5. 兼容性风险--use-replayssm 的语义从 Mamba2 专用扩展为 Kimi-K3 RecoverSSM 复用,vllm/config/vllm.py 中原有

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论