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 合并状态。