ReplaySSM: Cache SSM Inputs, Not State

发表时间: 2026-06 · Blog post by Tri Dao (tridao.me)

原文: https://tridao.me/blog/2026/replayssm/

https://github.com/Johnny-Liou/ReplaySSM.git

速读

一句话结论 本文提出了 ReplaySSM 方法,通过在解码时缓存最近的输入而非每步写回循环状态,大幅降低了内存流量,在 vLLM 服务中实现了标准解码最高 1.48 倍、投机解码最高 1.96 倍的端到端加速。

要解决什么问题 状态空间模型(如 Mamba-2 等)将历史信息压缩为固定大小的循环状态以保持内存恒定,但这种汇总并丢弃原输入的机制带来了三个机制卡点。首先是内存 I/O 瓶颈:SSM 状态更新是轻量级的秩 1 外积加累加,算术强度极低,但每步都必须从高带宽内存读写庞大的状态矩阵,完全被内存流量限制。其次是状态不可逆导致投机解码回滚困难:SSM 的历史压缩是有损的,草稿 token 被拒绝时无法像 Transformer 那样移动指针撤销。现有 vLLM 实现只能为每个投机位置存储完整的状态快照,使内存流量暴增了投机窗口倍数,在服务批处理大小下毫无加速效果。最后是串行状态依赖导致并行性丧失:状态严格依赖前一步,验证多个草稿 token 时无法融合成大的矩阵乘法,只能退化为串行循环。

怎么做的 核心思路是改变每步写回内存的内容:不再于每步更新并存储完整的循环状态,而是维护一个小容量的环形缓冲区,缓存最近的 SSM 输入。仅当缓冲区满时才执行刷新,将缓冲的历史输入汇总到状态中并清空;在两次刷新间的绝大多数步骤里,直接从缓存输入重建状态或计算输出。该设计绕开了上述卡点:将每步主导的内存流量减半;且因显式缓存了原始输入,投机解码拒绝草稿 token 时只需移动指针即可回滚,免去了状态快照读写。关键设计由三个部件构成。第一是基于历史路线的状态重建。以 Mamba-2 为例,标准更新为:$$S_t = a_t S_{t-1} + \Delta_t (v_t k_t^\top)$$ReplaySSM 不走摘要路线,而是将新输入追加到缓冲区。第二是仅输出解码机制。因两次刷新间检查点状态不变,多数步骤只需计算当前 token 输出,无需物化庞大状态矩阵。通过结合律,将先构建状态再读取的外积路线,替换为先计算标量内积再缩放的内积路线:$$y_t = \bar a_t (S_0 q_t) + \sum_{j=1}^{t} s_{j,t} v_j (k_j^\top q_t)$$这避免了生成状态矩阵,并利用多头机制中共享键查询的特性大幅减少计算量。第三是针对投机解码的提前刷新与分块并行机制。为防投机窗口被静默截断,当缓冲区剩余空间不足两个投机窗口时会提前刷新;对于 Gated DeltaNet 这类带校正项的模型,通过缓存校正项并应用分块并行求解下三角矩阵,将串行验证转化为并行矩阵乘法,移除串行循环。

效果如何 实验在 vLLM 框架下搭建并启用 CUDA Graph,使用 GSM8K 提示词评估端到端吞吐量。测试模型涵盖代表 Mamba-2 路线的 Nemotron-3 家族(4B 密集至 550B 混合专家模型)和代表 Gated DeltaNet 路线的 Qwen3.5 家族(4B 密集至 122B 混合专家模型),硬件使用 H100 和 B300 GPU。投机解码采用各自的 MTP 头作为起草者,窗口设为 4。对比基线为传统的 SSM 标准解码(每步读写状态)和现有 vLLM 投机解码基线(串行更新并为每个草稿存快照)。量化结果显示,在批处理大小 256 的标准解码测试中,ReplaySSM 在 Nemotron-3 上实现 1.20 至 1.48 倍端到端加速,在 Qwen3.5 上实现 1.20 至 1.27 倍端到端加速。投机解码场景下,扫描批处理大小至 512 时,ReplaySSM 在最大批处理下达到标准解码吞吐量的 1.87 至 1.96 倍,最高达基线投机路径的 2.14 倍。因免去基线为每个草稿预分配状态的开销,在固定高带宽内存预算下,ReplaySSM 支持的最大并发请求数提升了 3.0 至 3.3 倍。该方法的局限在于,端到端加速幅度低于内核层面的加速(内核加速可达 1.84 倍),因优化仅局限于 SSM 内核;此外,缓冲区容量需精心权衡,过小会导致频繁刷新成本过高,过大则引发计算瓶颈,通常中等大小的缓冲区(如 8 或 16)才能达到最佳性能。

主要贡献

在强化学习(RL)后训练和推理服务中,由于长思维链追踪,解码成为了主要的瓶颈。尽管Transformer仍然是行业标准,但其KV缓存使得处理长上下文的成本高昂,序列增长会导致内存流量和容量需求线性增加。状态空间模型(State Space Models, SSMs,如Mamba-2、Gated DeltaNet等)旨在解决这一瓶颈,它们将历史信息压缩为固定大小的循环状态,使内存需求保持恒定。然而,SSM将历史记录汇总为单一状态并丢弃输入的机制带来了三大挑战:受限于内存I/O、状态不可逆导致投机解码回滚困难、以及串行状态依赖导致并行性丧失。

为了解决上述问题,本文提出了ReplaySSM。该方法的核心创新在于:不再于每个解码步骤中更新并将循环状态写回高带宽内存(HBM),而是缓存最近的SSM输入,仅在需要时才重建状态,其余时间直接从缓存中计算输出。

这一设计在不改变输出结果的前提下,显著降低了内存流量,使得标准解码速度提升高达1.48倍(在大型MoE模型上为1.43倍)。更重要的突破在于投机解码(Speculative Decoding)领域:现有的vLLM实现在服务批处理大小下,其吞吐量甚至低于标准解码,而ReplaySSM成功解锁了1.87–1.96倍的加速比。

图1a ReplaySSM缓存最近的SSM输入代替每步存储循环状态,仅在缓冲区满时重建状态并写回
图1a ReplaySSM缓存最近的SSM输入代替每步存储循环状态,仅在缓冲区满时重建状态并写回
图1b vLLM中的端到端解码吞吐量对比。左图:ReplaySSM在标准解码中加速达1.43倍。右图:ReplaySSM在投机解码中提供1.87-1.96倍加速
图1b vLLM中的端到端解码吞吐量对比。左图:ReplaySSM在标准解码中加速达1.43倍。右图:ReplaySSM在投机解码中提供1.87-1.96倍加速

背景知识与设计原则

内存受限:I/O主导而计算匮乏

在每个解码步骤中,SSM重复执行两个操作:用新输入更新状态(摘要),并从中读取输出。在Mamba-2中,步骤$t$的循环形式为:

$$S_t = a_t S_{t-1} + \Delta_t (v_t k_t^\top) \quad \text{(更新)}, \qquad y_t = S_t q_t \quad \text{(读取)}$$


其中$S$为状态,$y$为输出。$v_t$和$k_t$是时间$t$的新输入,$q_t$是用于读取输出的输入。$a_t = e^{A\Delta_t}$和$\Delta_t$是每步的标量。
Delta-rule家族(如Gated DeltaNet, GDN)通过更复杂的计算获得了擦除状态的能力,其循环形式为:

$$S_t = e^{g_t} S_{t-1} (I - \beta_t k_t k_t^\top) + \beta_t (v_t k_t^\top) \quad \text{(更新)}, \qquad y_t = S_t q_t \quad \text{(读取)}$$
其中$I - \beta_t k_t k_t^\top$被称为Householder项,用于擦除信息。

在形状方面,$v_t$是头部维度为$d$的向量,$k_t$和$q_t$在一组头之间共享(形状维度为$n$)。每头状态是形状为$(d, n)$的矩阵。状态更新是轻量级的秩1外积更新加累加,无法高效映射到Tensor Cores等现代矩阵乘法加速器上。每个解码步骤都必须加载和存储$(nheads, d, n)$大小的状态,导致算术强度极低(约为1),完全被内存I/O主导。在混合模型(如Nemotron-3、Qwen3.5)中,尽管注意力层的成本随上下文增长,但由于中短上下文中注意力成本适中、SSM层数量通常是注意力层的3到6倍,且SSM内核缺乏优化,SSM状态更新内核在高达100K token的情况下仍然是主要的延迟瓶颈。

图2 Nemotron-3-Super-120B-A12B-NVFP4的延迟细分。尽管注意力随上下文长度缩放,但固定的SSM内核延迟在高达100K tokens时仍是最大组成部分
图2 Nemotron-3-Super-120B-A12B-NVFP4的延迟细分。尽管注意力随上下文长度缩放,但固定的SSM内核延迟在高达100K tokens时仍是最大组成部分

摘要的反噬:状态无法撤销

Transformer在KV缓存中显式存储整个历史,而SSM将历史压缩到固定大小的状态中。这种摘要是有损且不可逆的,一旦状态更新,模型就无法恢复产生它的确切token。这在投机解码中是一个严重问题,因为当草稿token被拒绝时需要回滚。注意力机制只需向后移动KV缓存中的指针,而SSM没有显式的token历史。vLLM中常见的解决方法是为每个投机token存储一个单独的SSM状态快照,在拒绝时恢复。这为本已受内存限制的路径增加了$T$倍($T$为投机窗口)的内存流量,导致投机解码在服务批处理大小下几乎无法带来加速。

图3 注意力机制通过移动KV缓存指针回滚,而SSM不可逆地将输入汇总到状态中,无法撤销
图3 注意力机制通过移动KV缓存指针回滚,而SSM不可逆地将输入汇总到状态中,无法撤销

并行性丧失:串行状态依赖

循环SSM状态是顺序依赖的,每个状态依赖于前一个状态。对于Transformer,投机解码可以增加并行性,因为对$T$个草稿token的验证可以批处理为一个大的GEMM。但SSM无法获得这种批处理,因为验证$T$个草稿token需要每个投机位置的状态和输出,计算仍然是一个长度为$T$的循环,无法用一次组合状态转换替换整个窗口。

方法细节

核心思想:缓存最近的输入而非存储状态

SSM解码步骤包括加载状态、更新状态、生成输出和将更新后的状态存储回内存。但写回状态的唯一消费者是下一步。根据Mamba-2的定义:

$$S_t = a_t S_{t-1} + \Delta_t (v_t k_t^\top) = \sum_{i \le t} \Big( \textstyle\prod_{i < j \le t} a_j \Big) \Delta_i (v_i k_i^\top)$$


这两个表达式暗示了两种计算状态的方式:
1. 摘要路线(左):加载摘要$S_{t-1}$并用新输入$v_t, k_t$更新。
2. 历史路线(右):从最近的输入窗口$(v_i, k_i)$重建$S_t$。

ReplaySSM改变了每步存储到内存的内容。它不存储循环状态,而是在小缓冲区中缓存最近的输入(对于Mamba-2,存储$(v, k)$对和衰减因子)。当需要状态时,ReplaySSM选择历史路线从缓冲的输入重建状态。当缓冲区足够大,加载它的成本超过写回状态时,ReplaySSM会刷新(flush)缓冲区:将缓冲的历史输入汇总到状态中,清空缓冲区,并重新开始缓存。状态写回仅在刷新步骤发生,大多数步骤只缓存小输入。

一项改变带来的三大优势

减少内存流量
在基线解码中,需加载和存储状态,每头的内存流量为$8dn + 2(d + 2n + 1)$(主导项为状态流量$8dn$)。在ReplaySSM中,假设缓冲区已缓存$h$个输入,需加载状态、缓存的$(v, \Delta, k)$及当前输入,并存储当前步输入。总流量为$4dn + 2h(d + n + 1) + 2(d + 2n + 1) + 2(d + n + 1)$。ReplaySSM将主导的状态流量从$8dn$减半至$4dn$。由于SSM解码受内存限制,减少内存流量直接改善了延迟。

回滚变为缓冲区操作
因为ReplaySSM显式缓存了最近的SSM输入(如草稿token),不执行每步的不可逆汇总,因此回滚拒绝的草稿token只需要删除它们的缓冲区条目,即移动指针,无需恢复完整状态。即使在投机解码触发更频繁刷新的情况下,大多数步骤仍避免了写入完整循环状态。

图4 ReplaySSM缓存每个草稿的原始输入,回滚拒绝的草稿只是指针移动,无需状态写回
图4 ReplaySSM缓存每个草稿的原始输入,回滚拒绝的草稿只是指针移动,无需状态写回

放宽要求与仅输出解码(Output-only decode)
在两次刷新之间,检查点状态不改变,最近的历史在缓冲区中,这意味着大多数解码步骤只需要当前token的输出。这解锁了新的解码方式:直接从检查点状态和缓存输入计算输出,而不物化状态。
假设零初始状态,输出为$y = S q = (v k^\top) q$。这可以有两种结合方式:
1. $(v k^\top) q$:构建完整状态然后读取。用于刷新缓冲区时。
2. $v (k^\top q)$:先计算标量内积,再缩放$v$,从建物化状态。用于大多数解码步骤。
对于非零检查点状态$S_0$和缓冲区,输出为:

$$y_t = \bar a_t (S_0 q_t) + \sum_{j=1}^{t} s_{j,t} v_j (k_j^\top q_t)$$


该形式不仅避免了物化$d \times n$的状态,还利用了$k$和$q$在头组内共享的特性,预计算$k_j^\top q_t$。在FLOPs方面,外积路线支付$2(h+1)dn$来物化状态,而内积路线将其替换为$2(h+1)(d+n)$,计算量大幅减少。

图5a 两种输出计算路线对比。一种构建完整状态$S = v k^\top$,另一种先形成标量$k^\top q$,从不物化状态
图5a 两种输出计算路线对比。一种构建完整状态$S = v k^\top$,另一种先形成标量$k^\top q$,从不物化状态
图5b ReplaySSM的两种路线与矩阵形状。状态和输出路线物化$S_t$,而仅输出路线先计算$K^\top q_t$
图5b ReplaySSM的两种路线与矩阵形状。状态和输出路线物化$S_t$,而仅输出路线先计算$K^\top q_t$

算法实现

Mamba-2 标准解码
- 基线:每步计算衰减$a \gets e^{A\Delta}$,加载$S$,更新$S \gets a S + \Delta (v k^\top)$,计算$y \gets S q$,并将$S$存回HBM。
- ReplaySSM:维护检查点$S_0$和容量为$L$的缓冲区$\mathcal{B}$。每步将$(v, \Delta, k)$追加到缓冲区;计算衰减权重$\bar a$和$s_j$;通过$y \gets \bar a (S_0 q) + \sum_{j=1}^{h+1} s_j (k_j^\top q) v_j$计算输出。如果缓冲区满,则执行刷新:$S_0 \gets \bar a S_0 + \sum_{j=1}^{h+1} s_j (v_j k_j^\top)$,并清空缓冲区。

投机解码
- 基线:为每个草稿位置串行执行状态更新,并为每个草稿存储完整的状态快照以备回滚。
- ReplaySSM:将草稿输入追加到缓冲区。对每个草稿$s$,计算衰减权重,读取检查点$H_{:,s} \gets \bar a_s (S_0 q_s)$;计算掩码GEMM $M_{j,s} \gets k_j^\top q_s$;计算输出$Y_{:,s} \gets H_{:,s} + \sum_{j \le p_s} w_{j,s} M_{j,s} v_j$。刷新时,仅将已提交的缓存输入汇总到状态中。
- 刷新决策机制:设$h$为当前缓存数,$T$为投机窗口。ReplaySSM在$h + 2T > L$时提前一个窗口刷新,而不是自然的$h + T > L$。这保证了每步至少有$T$个空闲槽,防止因空间不足导致投机窗口被静默截断。

内核设计细节
1. 预计算共享内积:在Mamba-2中,组内所有头需要相同的内积$k_j^\top q$。ReplaySSM在一个小型的预计算内核中计算它们(每组运行一次),主SSM更新内核读取暂存缓冲区,以降低寄存器压力。
2. 环形缓冲区:为了避免在投机解码中将接受的token重新定位回缓冲区前面,ReplaySSM使用环形缓冲区并在内核中进行索引,使得回滚纯粹变成指针移动。

实验环境

  • 数据集:使用GSM8K数据集的提示词评估端到端吞吐量。
  • 模型架构

    • Nemotron-3家族(Mamba-2):Nano-4B(4B密集, BF16),Super-120B(A12B MoE, NVFP4),Ultra-550B(A55B MoE, NVFP4)。
    • Qwen3.5家族(GDN):4B(4B密集, BF16),122B(A10B MoE, NVFP4)。
    • 投机解码使用各自的MTP头作为起草者,投机窗口为4。
  • 硬件配置:4B模型使用1×H100;120B/122B模型使用1×B300;Ultra-550B使用2×B300(张量并行TP2)。

  • 软件配置:在vLLM上实现,启用CUDA Graph。SSM状态为FP32,缓冲区向量为BF16。

实验结果

标准解码

在批处理大小256下,测试1K解码步骤。
- 结果:ReplaySSM在Nemotron-3上实现了1.43x至1.84x的内核加速和1.20x至1.48x的端到端加速;在Qwen3.5上实现了1.43x至1.64x的内核加速和1.20x至1.27x的端到端加速(图6)。端到端加速较小是因为仅优化了SSM内核。
- 缓冲区容量权衡:图7显示,中等缓冲区(Nemotron-3为8,Qwen3.5为16)平衡了频繁刷新的高成本和长缓冲区带来的计算瓶颈,呈现最佳的钟形曲线性能。

图6 跨Nemotron-3和Qwen3.5家族的内核级别和端到端每步加速比
图6 跨Nemotron-3和Qwen3.5家族的内核级别和端到端每步加速比
图7 不同缓冲区大小(4, 8, 16, 32)下的内核加速比
图7 不同缓冲区大小(4, 8, 16, 32)下的内核加速比

投机解码

  • 端到端吞吐量:在GSM8K提示词上,扫描批处理大小至512。ReplaySSM在保持相同草稿接受率的情况下,在最大批处理下达到标准解码的1.87–1.96倍,以及基线投机路径的最高2.14倍(图8)。
  • 更快的解码步骤:基线验证成本随窗口几乎线性增长(在Qwen3.5-122B上,$T=6$时基线内核成本是标准解码的4.85倍)。而ReplaySSM的内核成本保持平稳(1.27x到1.72x之间,图9)。2.28–3.33倍的内核加速最终转化为完整解码步骤的1.20–1.58倍加速(图10)。
  • 更高的最大并发性:在固定HBM预算下,基线预分配的每个草稿状态使最大解码批处理减少约4倍。ReplaySSM缓存小输入向量,恢复了3.0–3.3倍的并发请求容量(图11)。
图8 端到端解码吞吐量与批处理大小的关系。底部显示基线和ReplaySSM接受的token数完全相同
图8 端到端解码吞吐量与批处理大小的关系。底部显示基线和ReplaySSM接受的token数完全相同
图9 Qwen3.5-122B上投机解码内核延迟与窗口大小的关系
图9 Qwen3.5-122B上投机解码内核延迟与窗口大小的关系
图10 投机解码在内核、验证前向传递和完整解码步骤级别的加速比细分
图10 投机解码在内核、验证前向传递和完整解码步骤级别的加速比细分
图11 固定HBM预算下的最大解码并发性。ReplaySSM支持的并发请求比基线多3.0-3.3倍
图11 固定HBM预算下的最大解码并发性。ReplaySSM支持的并发请求比基线多3.0-3.3倍

结论

ReplaySSM通过一个简单的改变——缓存最近的输入而不是存储状态——减少了内存流量,实现了低成本回滚,并解锁了仅输出解码。该方法不仅适用于Mamba-2,也适用于GDN等delta-rule模型。在vLLM中的实现加速了标准解码,并消除了阻碍投机解码的关键障碍。未来计划将ReplaySSM的理念引入Mamba-3和GDN2等更多架构,并探索SSM在选择何时汇总输入和何时显式缓存方面的灵活性。

附录细节

Gated DeltaNet (GDN) 算法实现

标准解码
GDN基线在更新状态时包含一个校正项:$u \gets \beta (v - S k)$(减去状态在$k$处的读出),然后执行$S \gets S + u k^\top$。因为计算$u_t$需要状态$S_{t-1}$,不能直接缓存$(v, k)$。ReplaySSM的策略是缓存$u$。一旦$u_t$已知,状态更新变为$S_t = \alpha_t S_{t-1} + u_t k_t^\top$,与Mamba-2历史路线相同。
由于GDN步骤需要状态在$k$和$q$处的读出,ReplaySSM采用状态和输出路线:从衰减的检查点和外积$u_j k_j^\top$重建状态$S_h \gets \big(\textstyle\prod_j \alpha_j\big) S_0 + \sum_j \big(\textstyle\prod_{i>j}\alpha_i\big) u_j k_j^\top$,计算校正$u \gets \beta (v - \alpha (S_h k))$,然后输出$y \gets \alpha (S_h q) + u (k\!\cdot\! q)$。

投机解码
GDN验证同样存在串行循环和每草稿状态快照问题,且每个校正$u_s$依赖于前一个草稿之后的状态。ReplaySSM应用GDN训练时的分块并行(chunk-wise parallelism)方法,将循环展开得到$u_s = R_s - \sum_{s'<s} A_{s,s'} u_{s'}$。通过一次$T \times T$的严格下三角求解矩阵求逆,ReplaySSM可以同时计算所有$T$个校正,将整个验证步骤并行化为GEMM操作,彻底移除了串行循环和状态快照。</p>

在vLLM中运行CUDA Graph的补充设计

  • 处理批处理分歧:由于连续批处理和投机解码接受率不同,各序列到达刷新步骤的时间不同。ReplaySSM将刷新决策视为每序列数据,内核在运行时读取并分支,而不是作为编译时常量,从而允许同一个捕获的图处理整个批次。
  • 设备端提交与回滚:为了避免将接受的token计数发回主机导致流水线停滞,ReplaySSM使用一个小的提交内核直接在设备上更新缓冲区指针。这使得每种批处理大小只需一个捕获的图即可覆盖投机路径,无需主机同步。

与先前SSM投机解码方法的比较

先前的SSM投机解码方法(如Mamba-in-the-Llama和STree)虽然避免了为每个草稿保留单独状态,但仍会在每个解码步骤物化并写回至少一个循环状态到HBM,保持了顺序状态依赖。且它们主要在批处理大小为1的Mamba-2上评估。相比之下,ReplaySSM仅在刷新时写回状态,泛化到了GDN,并在服务批处理大小下在vLLM中成功运行。