ReplaySSM:缓存 SSM 输入而非状态,让混合模型推测解码提速近 2 倍 · 干货攻略
- 链接: https://x.com/tri_dao/status/2066518563184365953
- 分类: x-tips
- 来源: X @tri_dao
- 作者: Jay
- 更新: 2026-07-25
这是什么
ReplaySSM(Johnny-Liou / Dao AI Lab,2026 年 6 月 15 日发布)是一种让混合状态空间模型(Hybrid SSM)在 vLLM 推理引擎中同时加速自回归解码和推测解码的核内核优化技术。其核心洞察极其简洁:
不要每步都把 SSM recurrent state 写回 HBM——缓存最近的 SSM 输入,需要时再重建状态。
论文原文:"ReplaySSM caches the recent inputs instead and rebuilds the state on the fly. Same outputs, half the memory traffic."
官方基准数据: - 标准自回归解码:比 vLLM 原生 SSM kernel 提速最高 1.48x(大 MoE 模型上 1.43x) - 推测解码(Speculative Decoding):vLLM 现有实现在 serving batch size 下实际低于标准解码吞吐,而 ReplaySSM 解锁 1.87–1.96x 提速
官方资源: - 官方博客:ReplaySSM: Cache SSM Inputs, Not State(tridao.me,2026-06-15) - GitHub:Johnny-Liou/ReplaySSM(基于 vLLM,Apache-2.0) - vLLM 上游 RFC:#47572,PR:#47576
为什么值得关注
问题背景:SSM 解码的三个实际挑战
混合模型(如 Qwen3.5、Nemotron-Ultra、Kimi Linear)在生产环境中大量使用 SSM 层(Mamba-2 / Gated DeltaNet)搭配少量注意力层,SSM 层数通常是对应注意力层数的 3~6 倍。然而 SSM 的 recurrent state 机制在工程实现上引入了三个被低估的瓶颈:
1. Memory-bound:每步都要读写完整 state
Mamba-2 的 recurrent state 形状为 (nheads, d, n),通常 d 和 n 各自为 64 或 128。每步解码需要从 HBM 读取 state、更新、写回 HBM——而 state update 本身的算术强度只有 ~1 FLOP/byte(H100 上 matmul 是 ~300 ops/byte)。state 读写成为绝对瓶颈,而非计算。
2. Summarization is irreversible:状态没有 undo 标准 SSM 每步都会对历史做不可逆压缩(summarization)。一旦某步的 draft token 被拒绝、需要回滚,SSM 无法像 attention 那样简单地丢弃 KV——它已经把信息压缩进了 state,回滚代价极高。
3. 推测解码在 serving batch size 下反而更慢 vLLM 现有的推测解码实现,在实际 serving batch size(≥64)下,由于每次都要 materialize 并写回 state,总吞吐量反而低于什么都不做的标准自回归解码。
Tri Dao 的 X 帖为什么重要
Tri Dao(原 FlashAttention/FlashAttention-2/3 作者,Dao AI Lab 负责人)在 X 上指出了一个反直觉但工程上极有效的 insight:对于长上下文 agent 场景,SSM state 不是「缓存加速」而是「瓶颈本身」——"load the states, compute, but don't store them"这个简单改动,解锁了 SSM 的推测解码工程可行性,让原本无法落地的 spec decode 路径变得值得使用。
核验过程
官方来源
| 来源 | 读取内容 | 关键结论 |
|---|---|---|
| 官方博客 | 全文(含方法、公式、benchmark 数据) | 确认 1.43x(标准 AR)/ 1.87-1.96x(spec decode)速度提升数字;确认三个核心挑战;确认 ring buffer 机制 |
| GitHub README | README、benchmark 脚本路径、kernel 文件列表 | 确认基于 vLLM commit 37ce34922;确认支持 Mamba-2 和 Gated DeltaNet;确认 benchmark 命令;确认 upstreaming 到 vLLM |
| vLLM RFC #47572 | GitHub Issue 讨论 | 确认 upstream 状态;确认该 PR 面向 vLLM 集成 |
| vLLM PR #47576 | PR 详情 | 确认 productized 版本正在合入主线 |
交叉验证
原帖说法 vs 官方数据:
| 原帖主张 | 官方数据 | 结论 |
|---|---|---|
| "make this 2x faster" | spec decode 1.87–1.96x(spec window=4,B300,NVFP4) | ✅ 官方数据完全吻合,"2x"是保守表述 |
| "hybrid models (Qwen 3.5 / Nemotron Ultra)" | 博客测试了 Qwen3.5-4B、Qwen3.5-122B-A10B-NVFP4、Nemotron-3-Nano-4B、Nemotron-3-Super-120B | ✅ 官方数据完全吻合 |
| "Gated-DeltaNet / Mamba states become a bottleneck" | 博客 Section 2 明确列出"Three challenges in SSM decoding",第一条就是 memory-bound | ✅ 吻合 |
| "recompute trick finally unlocks spec decoding for SSMs" | 博客 Section 6.2:vLLM 现有 spec decode 在 batch≥64 时低于标准解码,而 ReplaySSM spec 解码实现 1.87-1.96x | ✅ 官方数据完全吻合 |
benchmark 命令核验(GitHub README):
- python benchmarks/replayssm/e2e_decode_speedup.py --model-id nvidia/NVIDIA-Nemotron-3-Nano-4B-BF16 → 标准解码对比
- python benchmarks/replayssm/e2e_spec_decode_throughput.py --batch-size 512 --model-id nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-NVFP4 → 推测解码对比
- 两套 benchmark 均可复现,命令行参数与博客数据一致
上手步骤
核心原理速览
ReplaySSM 把 SSM decode 分成两类:
普通步(大多数):不写回 state,只把输入 (v_t, k_t, g_t) 追加到一个小的 ring buffer。然后直接用 ring buffer 内的最近 L 条记录,重建 state,再读出输出——完全不需要每次都访问 HBM 中的完整 state matrix。
Flush 步(周期性):当 ring buffer 满时,执行一次完整的 state write-back 到 HBM,作为后续重建的 checkpoint。
# 伪代码(基于博客 Algorithm 描述)
# 标准 SSM:每步都读 S、写 S
S = load_state_from_HBM()
S = a_t * S + Δ_t * outer(v_t, k_t) # update
y_t = S @ q_t # read
store_state_to_HBM(S)
# ReplaySSM:普通步只追加到 ring buffer,按需重建
# 绝大多数步(non-flush):
ring_buffer.append((v_t, k_t, g_t)) # 写 HBM:small input tuple
S = reconstruct_from_buffer(ring_buffer) # 纯计算,无 HBM 读
y_t = S @ q_t
# Flush 步才做完整的 state write-back
推测解码中的优势:当 draft token 被拒绝时,标准 SSM 无法回滚(state 已不可逆更新);ReplaySSM 只需要从 ring buffer 中移除对应条目即可——rollback 变成了一个纯粹的 buffer 操作,代价极低。
安装(基于 Johnny-Liou/ReplaySSM)
# 克隆研究参考实现(基于 vLLM Apache-2.0)
git clone https://github.com/Johnny-Liou/ReplaySSM.git
cd ReplaySSM
# 依赖已通过 vLLM 满足,参考官方安装
# 正式使用推荐等 vLLM upstream PR 合并后从主线安装
标准解码基准对比
# 在单卡 H100 上,对比 ReplaySSM AR decode vs vLLM 标准 SSM kernel
python benchmarks/replayssm/e2e_decode_speedup.py \
--model-id nvidia/NVIDIA-Nemotron-3-Nano-4B-BF16
# Qwen3.5-4B(推荐 buffer-len 16)
python benchmarks/replayssm/e2e_decode_speedup.py \
--model-id Qwen/Qwen3.5-4B \
--buffer-len 16
推测解码吞吐量对比
# 单卡 B300,batch=512,对比:AR vs vLLM spec decode vs ReplaySSM spec decode
# Nemotron-3-Super-120B(MoE,NVFP4)
python benchmarks/replayssm/e2e_spec_decode_throughput.py \
--batch-size 512 \
--model-id nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-NVFP4
# Qwen3.5-122B-A10B(NVFP4),使用 qwen3_next_mtp spec method
python benchmarks/replayssm/e2e_spec_decode_throughput.py \
--batch-size 512 \
--model-id nvidia/Qwen3.5-122B-A10B-NVFP4 \
--spec-method qwen3_next_mtp
vLLM 命令行使用(等 upstream PR 合并后)
# 安装主线 vLLM(预计 PR #47576 合并后)
pip install vllm
# 使用 ReplaySSM kernel(upstream 提供 --use-replayssm flag)
vllm serve Qwen/Qwen3.5-4B \
--enforce-eager \
--use-replayssm \
--replayssm-buffer-len 16
Triton Kernel 文件(开发者参考)
| Kernel 文件 | 功能 | Entry 函数 |
|---|---|---|
selective_state_update_replayssm_output_only.py |
Mamba-2 AR decode,output_only 路由(默认) | selective_state_update_replayssm_output_only |
selective_state_update_replayssm_state_and_output.py |
Mamba-2 AR decode,state_and_output 路由 | selective_state_update_replayssm_state_and_output |
selective_state_update_replayssm_spec.py |
Mamba-2 推测解码(circular buffer) | selective_state_update_replayssm_spec |
fused_recurrent_replayssm.py |
Gated DeltaNet AR decode | fused_recurrent_gated_delta_rule_replayssm |
gdn_replayssm_spec_decode.py |
Gated DeltaNet 推测解码 | gdn_replayssm_spec_decode |
坑与适用边界
已知的坑
1. Blackwell / NVFP4 的 FlashInfer autotuner 在 CUDA graph capture 下可能不稳定
官方 benchmark 脚本默认禁用 --disable-flashinfer-autotune,可手动 override。预发布 Blackwell 驱动环境下需注意。
2. buffer-len 需要调参 ring buffer 长度直接影响重建精度与内存开销。官方博客测试了不同 buffer-len 值,最优值因模型规模而异(Qwen3.5-4B 推荐 16)。
3. NVFP4 量化目前依赖特定硬件 官方 benchmark 使用 B300 + NVFP4 精度,H100 不支持原生 NVFP4,实际部署需确认硬件。
4. 代码尚未进入 vLLM 主线 当前可用的仍是 Johnny-Liou 的研究分支,基于 vLLM commit 37ce34922。生产使用建议等 upstream PR #47576 合并。
适用边界
- 适合:使用 Mamba-2 / Gated DeltaNet 混合模型的推理服务,尤其在长上下文、agent 场景下;需要推测解码加速的场景;batch size ≥ 64 的高并发 serving
- 不适合:纯 Transformer 模型(非 SSM);纯 CPU 推理;极小 batch size(此时原版 spec decode 差距不大);对延迟而非吞吐优化的场景
一句话结论
ReplaySSM 把 SSM 的 recurrent state 从「每步读写 HBM」变成「按需从 ring buffer 重建」——实测标准解码最高 1.48x、推测解码 1.87-1.96x 提速,让混合 SSM 模型在 agent 长上下文场景下终于能高效跑 spec decode;代码已 upstream 到 vLLM,生产可用前建议跟踪 PR #47576 合并状态。