replayssm(B)SpecDecode
PR #28695 详解:GDN ReplaySSM Ring Spec-Verify
本文件逐字记录了关于 SGLang PR #28695 「[GDN] Support ReplaySSM Ring Spec-Verify」的全部讲解。
它是 #28511 的 Part B(移植 Dao AI Lab 2026 的 ReplaySSM)。Part A(decode 版)见 姊妹文档
replayssm_pr28451_explained.md。配套源码(PR head
yuan-luo/sglang@fe66699):python/sglang/srt/layers/attention/fla/gdn_replayssm_spec_decode.py(741 行)。
0. TL;DR(一句话)
PR #28695 把 GDN 投机解码的 target-verify 路径——目前为每个 draft token都往 intermediate_ssm 写一份完整 [V,K] 递归状态以便回退——替换为每个 slot 一个环形 (d,k,g) 缓存 + 冻结 checkpoint:整窗 verify 输出靠chunked delta-rule (I+A)^{-1} UT-transform output-only 重构(非 flush 步永不物化状态),完整状态只在每 L 个已提交 token 才折叠回 checkpoint(flush),被拒
draft 的回退退化为环形游标的一次指针移动。短输出数学无损(GSM8K 持平),长输出非无损(见 §8)。默认关闭,--enable-gdn-replayssm-spec opt-in,仅 GDN +linear-chain(speculative-eagle-topk <= 1,即 NEXTN / MTP)。
1. 并行性分析(最初的提问)
问题:查看 replayssm 的 spec decode 代码,分析是否用了 matmul 并行 + layer并行?
结论先行:用了 matmul 并行(token 轴),没有用 layer 并行。 这和 decode 版的ReplaySSM (#28451) 完全一致,原因是结构性的。
1.1 Matmul 并行(token / draft-window 轴):✅ 用了,而且是核心
verify kernel gdn_replayssm_spec_circular_kernel 把整个 draft window(BS 个token)一次性 matmul 重构,没有逐 token 递推循环。关键 tl.dot 链(全在 tensorcore 上跑):
| 行 | 运算 | 含义 |
|---|---|---|
:267 |
kk_mat += tl.dot(k_tile, kT) |
窗口内 key-key 交互 |
:268 |
kq_mat += tl.dot(k_tile, qT) |
key-query |
:286-289 |
hw_q/hw_k += tl.dot(sc_tile, qT/kT)、scores_q/k += tl.dot(khist_tile, ...) |
checkpoint 态 + 环形历史 (d,k,g) 投影到输出 |
:315-323 |
A_mat 严格下三角 + T = (I+A)^{-1} UT-transform(Neumann 展开 I−A+A²…) |
窗口内 delta-rule 解耦 |
:337 |
D_spec = R @ Tᵀ |
窗口内贡献 |
:349 |
DF = D_spec @ F,:350 O = expG·hw_q + DF |
最终输出 |
这就是 chunked delta-rule 的闭式重构——和 #28451 decode kernel 用 tl.dot 把 做成一次 matmul 是同一套数学。d/k/g 环形缓存里存的是 d(已解耦的 rank-1 因子),所以不必串行递推,可以直接矩阵化。token 轴并行 = 有。
1.2 Layer 并行:❌ 没有,且结构上不可能
grid 在 :517:
1 | grid = (triton.cdiv(V, BV), B, HV) # (V块, batch, value-head) —— 没有 layer 轴 |
- kernel 的所有指针(
q/k/v/a/b是[total_tokens, ...]、A_log/dt_bias是[HV]、checkpoint_state是[num_slots, HV, V, K])都只对应单个 layer。 - backend (
gdn_backend.py) 在 forward 的 target-verify 路径里每层各调一次这个kernel。
原因和 decode replayssm 一样:verify 跑在 forward 内部,第 L 层的输入依赖第L−1 层的输出(链式依赖),无法把 36 层堆进一次 launch。这正是与 KV-buffer replay的根本区别——KV-buffer replay 跑在 forward 之外(verify 之后重放),各层内存不相交,所以能 layer-batched 单次 launch;而 ReplaySSM(无论 decode #28451 还是这个spec #28695)赢的是 token 轴、输不了 layer 轴。
1.3 两个辅助 kernel(非数值,纯游标)
_advance_gdn_spec_cursors_kernel(:375):每个 decode step 把环形游标write_pos/cache_base/is_flush按 accept 数推进一格(拒绝回滚 = 指针移动,不写状态)。grid(1,),一个 block 扫全 batch。_reset_gdn_replayssm_spec_cursors_kernel(:432):prefill→decode 交接时重置首步行的游标。
二者都是 device-side、CUDA-graph 安全的游标算术,不涉及并行重构。
1.4 并行轴总结表
| 并行轴 | #28695 spec-verify | 机制 | 与 KV-buffer replay 对比 |
|---|---|---|---|
| token / draft-window(matmul) | ✅ 有 | (I+A)^{-1} UT-transform + tl.dot 闭式重构整窗,tensor core |
KV-buffer 串行递推,输 |
| layer | ❌ 无 | grid (V/BV, B, HV),每层一次 launch,forward 内链式依赖 |
KV-buffer forward 外、内存不相交可单次 launch,赢 |
1.5 PR 自己披露的已知局限
这条路径仅对短输出 lossless(GSM8K 0.892 vs 0.887 持平);长输出(16k+,AIME-2024)因 (I+A)^{-1} 折叠的灾难性抵消导致状态漂移、退化成重复循环(0.767 vsrecurrent 0.933),所以默认关闭、--enable-gdn-replayssm-spec opt-in,且只 GDN +linear-chain(topk≤1 的 MTP/NEXTN),tree-verify/KDA/NPU 自动回退到递推 verify。
2. 为什么 KV-buffer 能 layer 并行、ReplaySSM 不能?——forward 内 vs forward 外
问题:KV buffer 的 verify 也是串行的,但是 verify 之后的 replay 是可以 layer并行的,这个 PR 为啥不可以 replay 的时候并行?
核心区别:两者"replay"的时机和输入完全不同。
2.1 KV-buffer replay:在 forward 之外,所以能 layer 并行
时间线:
1 | forward (verify) → 采样接受/拒绝 → replay 重建状态 |
replay 发生在整个 forward 跑完、采样结果已知之后。此时:
- 每层要重放的输入
(mixed_qkv, a, b)在 verify 阶段已经全部算出来并存好了(存在intermediate_kvbuffer 里)。 - 第 L 层 replay 的输入不依赖第 L−1 层 replay 的输出——L 层的 q/k/v 早就在forward 时算完了,replay 只是拿这些已知的输入重新跑一遍 GDN 递推去重建该层自己的
ssm_states[L]。 - 各层内存不相交 → 36 层堆进一次 launch(layer 轴上 grid 扩 36 倍)。
关键:replay 的输入是"已经物化的激活值",层间没有数据依赖。
2.2 ReplaySSM 的 reconstruction 就是 forward 本身
ReplaySSM 没有"verify 之后再 replay"这个独立阶段。它的 (I+A)^{-1} 闭式重构就是verify forward 的那一次前向计算:
1 | forward (verify) = reconstruction |
- 这个 kernel 的输出
o[:, HV, V]是该层 attention 的输出,要喂给同一层后面的MLP,再喂给下一层。 - 第 L 层的 q/k/v(kernel 的输入)来自第 L−1 层输出经过整个 transformer block 的结果——第 L−1 层不算完,第 L 层的输入根本不存在。
所以它没有"replay 阶段"可以并行,因为 reconstruction 不是事后补算,而是 forward流水线的一环。
2.3 一句话对比
| 输入何时就绪 | 层间依赖 | 能否 layer 并行 | |
|---|---|---|---|
| KV-buffer replay | forward 全跑完,所有层输入已物化 | 无(各层独立重放自己的状态) | ✅ 能 |
| ReplaySSM 重构 | 重构=forward,第L层输入需第L−1层先算完 | 有(链式) | ❌ 不能 |
2.4 ReplaySSM 用什么换 layer 并行的损失?
它放弃 layer 轴,但在 token 轴赢回来:KV-buffer replay 在每层内部还是逐 token串行递推重建状态(代价 ∝ batch × accept_len),而 ReplaySSM 因为存了已解耦的d,每层内部一次 tl.dot 把整个 draft window 并行重构出来。
两者是互补的,不是谁能用 layer 并行谁就赢:
- KV-buffer:forward 外 → 赢 layer 轴,输 token 轴(层内串行)。
- ReplaySSM:forward 内 → 输 layer 轴,赢 token 轴(窗口 matmul)。
如果硬要让 ReplaySSM “layer 并行”,就得把所有层的输入提前物化——但那需要先把整个forward 跑完,那就回到 KV-buffer 的路子了,ReplaySSM 省 forward 内带宽的意义就没了。
3. 详解"forward 内 vs forward 外"(具体例子)
设定:36 层模型,batch 里一个请求,draft 了 4 个 token。
- 模型:36 层,每层 =
GDN linear attention + MLP。 - forward 必须逐层串行:layer0 算完 → layer1 → … → layer35。因为 layer1 的输入= layer0 的输出。这是 transformer 的铁律。
- spec decode:先用 4 个 draft token 跑一遍 forward(verify forward),拿到每个位置的 logits,采样后才知道接受了几个(比如接受前 3 个,第 4 个拒绝,accept_len=3)。
3.1 什么是 “forward 内”
指:事情发生在 verify forward 这条逐层流水线的内部,是 layer 计算的一部分。
ReplaySSM 的状态重构就在这里。走一遍 layer0 的 GDN:
1 | layer0 的 GDN kernel 输入:x0 (= embedding) |
为什么不能 layer 并行:你想同时启动 layer0~layer35 的重构 kernel,但 layer5的输入 h4 还不存在——它要等 layer4 算完。layer4 又要等 layer3……这是一条锁死的链。重构和 forward 是同一件事,没有独立出来的余地。
3.2 什么是 “forward 外”
指:verify forward 整条流水线已经全部跑完、采样也做完之后,再单独做的一步。
KV-buffer replay 在这里。时间线分两段:
第一段:verify forward(逐层串行)
每层的 GDN 不更新持久状态,只做两件事:
1 | layer0 GDN: 算出 o0,并把原始输入 (mixed_qkv_0, a_0, b_0) 存进 buffer[layer0] |
forward 结束 → 采样 → 得到 accept_len=3。
此刻关键的事实:36 层每一层要重放的输入,全部已经躺在 buffer 里了。
第二段:replay(forward 外,可以 layer 并行)
现在要把状态重建到"接受 3 个 token 后"的正确值。每层做:
1 | ssm_states[layer] = 拿 buffer[layer] 里前 3 个 token 的输入,重跑 GDN 递推 |
问:layer5 的 replay 需要等 layer4 的 replay 吗?
答:不需要! layer5 重放用的输入 buffer[layer5] 在第一段 forward 时就算好存好了,跟 layer4 的 replay 结果毫无关系。每层只读写自己的 buffer[layer] 和ssm_states[layer],内存不相交。所以可以一次 launch,grid 多加一个 36 的 layer 轴,36 层同时重放。
3.3 一张图对比
1 | KV-buffer replay(两段分离): |
3.4 一句话本质
判断能否 layer 并行,只看一个问题:“各层要算的东西,它们的输入是否已经全部就绪、 互不依赖?”
- KV-buffer replay:把昂贵的状态重建推迟到 forward 之后,那时所有层的输入都物化在 buffer 里、层间无依赖 → 能并行。
- ReplaySSM:状态重构塞在 forward 里就地完成,省了"存输入 + 事后重放"那一整套,但代价是它被绑死在逐层链上 → 不能并行。
ReplaySSM 换来的好处在另一个轴:它每层内部用一次 matmul 把整个 draft window 并行重构(token 轴赢),而 KV-buffer 在每层 replay 时还是逐 token 串行递推(token 轴输)。两者各赢一个轴。
4. ReplaySSM 怎么回退?——游标移动,零重算
问题:KV buffer 的 forward 之外 replay 是因为需要回退。ReplaySSM 没有 forward之外的 replay,咋回退的?
结论:不需要 replay。ReplaySSM 的回退完全靠环形缓存的游标移动,根本不重算、不重放。
4.1 状态的两个组成部分
每个请求的状态由两部分组成:
- 冻结的 checkpoint
h0:上次 flush 时固化的完整状态[V,K],verify 期间只读,绝不动。 - 环形缓存
(d,k,g):从 checkpoint 之后每个已提交 token 的 rank-1 因子,配三个游标:write_pos:已提交了多少个 token(环里 valid 历史的长度)cache_base:环形起点is_flush:本步是否该 flush
关键:verify 时一个字都不往持久状态里写。 输出靠 h0(只读)+ 环形历史 + 当前draft 窗口,三者一次 matmul 重构出来(gdn_replayssm_spec_circular_kernel)。draft 的d/g 顺手写到环里 write_pos 之后的临时位置(phys_spec,:357-371),但
游标 write_pos 没有前进*。
4.2 走一遍:draft 4 个,接受 3 个
1 | verify 前: write_pos = 100(已提交 100 个 token 在环里) |
_advance_gdn_spec_cursors_kernel (:404-410) 的核心:
1 | total_commit = num_acc # = 3 |
- 接受 3 个 = 游标推进 3,临时槽 [100,101,102] 自动变成"已提交历史"。
- 拒绝第 4 个 = 游标不推进到 103 以上,临时槽 [103] 里那条 draft 记录留在原地不管它——下次 verify 直接覆盖写。它从没被算进
write_pos < 历史长度的有效范围(kernel 里cache_valid = o_c < write_pos,:140),所以等于不存在。
回退 = “本来要 +4,现在只 +3”,一个加法。没有任何重算、重放、回写。
4.3 为什么 KV-buffer 做不到这种"零成本回退"
| 持久状态在 verify 时被改了吗 | 回退方式 | 代价 | |
|---|---|---|---|
| KV-buffer | (它选择不存中间态,只存原始输入) | forward 外重放前 accept_len 个 token,逐 token 递推重建 | 一整段串行 replay(∝ batch×accept_len),但可 layer 并行 |
| ReplaySSM | 没改(h0 只读,全程不写持久状态) | 游标 write_pos += num_accepted |
一次加法,O(1) |
本质原因:ReplaySSM 把"提交"和"计算"解耦了。
- 计算(重构输出 + 把 draft 的 d/g 塞进环)在 verify 时做,但不改变任何被认定为"已提交"的状态。
- 提交(让这些 d/g 生效)是事后游标一推,且只推被接受的数量。
因为"已提交状态"在 verify 期间从未被污染,拒绝的 draft 天然没进入有效历史,无需回退——它压根没生效过。这正是 ReplaySSM 不需要 forward 外 replay 段的原因。
4.4 那 flush 时呢?(状态怎么真正前进)
环是有限长 L。每攒够一窗就 flush:把 checkpoint 往前折叠一次——
1 | # _advance_gdn_spec_cursors_kernel :407-410 |
flush kernel (IS_FLUSH 分支 :270-284) 把"h0 + 已提交历史"折进新的 ht,写回checkpoint。注意 flush 只折叠已 commit 的 token,所以被拒绝的 draft 永远不会进checkpoint。
顺带:PR 自己披露的长输出退化(0.767 vs 0.933)就出在这个 flush 的
(I+A)^{-1}折叠上——每次 flush 都把整个 checkpoint 经(I+A)^{-1}过一遍 bf16,~1000 次折叠后条件数误差累积、状态漂移。这是 flush(状态前进)的问题,不是回退的问题;回 退本身(游标加法)是无损的。
4.5 一句话
KV-buffer 因为没存中间态,verify 后必须事后重放才能回到正确状态(于是有了forward 外那段、于是能 layer 并行);ReplaySSM 在 verify 时对已提交状态零写入,接受/拒绝只表现为环形游标 write_pos += num_accepted 一个加法——拒绝的 draft 从未生
效,无需回退,自然也就没有 forward 外的 replay 段。
5. 接受的 token 输出 o 何时存?——verify 当场算完,用完即弃
问题:draft token=4,verify 之后 accept 的 token=3,如果没有 replay,那么这 3 个接受的 token 的输出 o 是啥时候被存到 layer cache 里的(最终返回的)?
结论:接受的 3 个 token 的输出 o 在 verify forward 那一次就已经算出来并写进out 了,根本不需要事后补存。
5.1 out 是整个 draft 窗口一次性全算出来的
verify kernel 不是只算"接受的"——它把全部 4 个 draft 位置的输出一次性重构出来,写进预分配的 out[total_tokens, HV, V]:
1 | # :350 整窗输出(BS 个位置,这里 BS≥4) |
p_o (:113) 指向 out 里 [bos, bos+1, bos+2, bos+3] 这 4 个 token 的位置。所以forward 跑完时,out 里已经有 4 个位置的 attention 输出了,包括第 4 个(后来被拒的)。
5.2 谁来"挑出"接受的 3 个?——上层调度,不是这个 kernel
这个 kernel 只负责"把 4 个位置都算对"。接受/拒绝的裁剪发生在更上层的 spec-decode框架,时间线:
1 | verify forward |
5.3 区分三种"输出",不要混
| “输出” | 是什么 | 何时产生 | 拒绝的第 4 个怎么办 |
|---|---|---|---|
attention 输出 o |
该层 GDN 的输出张量,喂给 MLP | verify kernel 一次算全 4 个,:352 写进 out |
算了但下游丢弃 |
| token(最终返回给用户的) | 采样出的 token id | lm_head + 采样后 | 拒绝采样剔除 |
| SSM 持久状态 | 环形 (d,k,g) + checkpoint |
verify 时 d/g 写临时槽;commit 时游标转正 | 临时槽留着被下次覆盖,从未转正 |
你问的"接受的 3 个 token 的输出 o"——指的是表里第一行的 attention 输出 o,它在 verify forward 那一刻就和第 4 个一起被算出来、写进 out 了。不存在"事后把这3 个的 o 补存进 layer cache"这个动作,因为:
o不入任何持久 cache。o只是层间的临时激活——layer N 的o喂给 layer N的 MLP,再传给 layer N+1,用完即弃。它本来就不存到跨 step 的 cache 里(无论ReplaySSM 还是普通 decode 都如此)。- 真正跨 step 持久化的是 SSM 状态,那部分由游标 commit(
write_pos += 3)完成,转正的是(d,k,g),不是o。
5.4 对比 KV-buffer
即使在 KV-buffer 里,事后 replay 也不是为了补存 o:
- KV-buffer 的 forward 外 replay 重建的是 SSM 持久状态(
ssm_states[layer]),不是o。o同样是 verify forward 当场算完、当场用掉的。 - replay 之所以必要,是因为 KV-buffer 在 verify 时没更新持久状态,得事后用接受的3 个 token 重跑递推把状态补到位。
5.5 一句话
接受的 3 个 token 的 attention 输出 o,和被拒的第 4 个一起,在 verify forward 那一次 matmul 重构里就全部算出来并写进了 out(:350-352);o 是层间临时激活、用完即弃,从不进跨 step 的 cache。"接受 3 个"这件事影响的只是——(a) 上层把哪 3 个token id 收进序列,(b) 游标 write_pos += 3 把这 3 个的 (d,k,g) 转正为已提交历史。没有"事后补存 o"这个步骤,也不需要。
6. 核心洞察:持久化 dkg 而非 ssm 状态
问题(用户总结):ReplaySSM 因为持久化的是 dkg 而不是 ssm 状态,所以不需要replay 计算出 ssm,在 verify 的时候已经存了 dkg,只不过 commit 的时候决定持久化哪些token。
校准:方向对,但更精确说是 两者都持久化,但分工不同:
- checkpoint
h0:完整 ssm 状态,但只在 flush 时(每 L 个 token)才更新一次,verify 期间只读。 - 环形
(d,k,g):从上次 flush 到现在、每个已提交 token 的 rank-1 因子(增量历史)。
所以"当前真实状态" = h0(陈旧的基底) + 环里的 (d,k,g)(增量)。verify 时不需要把这俩合成出完整 ssm 状态——kernel 直接拿 h0 只读 + 环 + draft 窗口,一次 matmul出输出(:286-289 用 h0,:307-308 叠环历史,:350 出结果)。
6.1 核心逻辑链(完全正确)
1 | 持久化 dkg(增量),不在 verify 时合成/写 ssm 状态 |
6.2 为什么"存 dkg"就免了 replay
KV-buffer 存的是原始输入 (mixed_qkv, a, b),这些还没"消化"成对状态的贡献——要变成状态必须串行递推走一遍 GDN,这就是它 forward 外那段 replay 的本质(把输入递推成状态)。
ReplaySSM 存的 d 是已经解耦好的 rank-1 因子(,在verify 时就算出来了),它对状态的贡献是 ——可加、可矩阵化、无串行依赖。所以"转正"只是把它纳入有效历史范围(游标一推),下次重构时它自然被 matmul 吸收,永远不需要把它"递推"成显式状态。
一句话:KV-buffer 存的是"未消化的输入"→ 必须事后递推(replay);ReplaySSM 存的是"已解耦的增量 dkg"→ 只需游标转正(commit),重构时 matmul 现取现用。 commit 决定哪些 token 生效,拒绝的从未生效即等于回退。
7. flush 是 matmul,不是 replay
问题:如果 accept 了 token 之后需要 flush 了,不就需要 replay 了吗?
结论:不需要 replay。flush 和 replay 是两件不同的事——flush 也是一次 matmul,不是逐 token 串行递推。
7.1 flush 在算什么
flush 要做的是:把"陈旧 checkpoint h0 + 环里这一窗已提交的 (d,k,g)"折叠成一个新的完整 checkpoint,写回去。数学上就是闭式那一步:
- (整窗总衰减,
b_total_decay,:170) - (replay 衰减,
b_replay_decay,:169)
看 kernel 的 flush 分支(IS_FLUSH=True,:270-284):
1 | sw_f = tl.dot( |
tl.dot(b_d_scaled, khist_tile) —— 把整窗的 d(已带好 W_j 衰减,:180)和历史k 一次矩阵乘,把所有 d_j·k_jᵀ 的贡献并行累加,再叠上衰减后的旧 h0。一次tl.dot 出整个新状态。
7.2 为什么 flush ≠ replay
| flush(ReplaySSM) | replay(KV-buffer) | |
|---|---|---|
| 输入 | 已解耦的 d/k/g(rank-1 因子) |
未消化的原始输入 mixed_qkv,a,b |
| 算法 | 一次 tl.dot 闭式累加 |
逐 token 串行 |
| token 轴 | 并行(matmul) | 串行递推 |
| 代价 | O(1) 次 matmul | O(accept_len) 步递推 |
根本原因还是那个:d 是已解耦的因子, 可加、可矩阵化、无串行依赖。 所以把它们折成完整状态只需一次带衰减权重的 matmul,不必"一个 token 一个 token 地递推"。这正是 ReplaySSM 存 d 而不存原始输入的全部意义——连 flush 都能并行。
7.3 时间线对照(accept 3 个后触发 flush)
1 | verify forward(matmul 重构输出,写 out + 把 draft 的 dkg 写临时槽) |
整条链上没有任何逐 token 串行重放。flush 只是"把增量 matmul 进基底",KV-buffer的 replay 是"把原始输入串行递推成状态"——前者赢在 token 轴并行,这也是为什么ReplaySSM 即便不能 layer 并行,仍能在带宽受限的 decode/verify 上拿到收益。
补充:长输出退化恰恰出在这个 flush matmul 的
(I+A)^{-1}折叠上(~1000 次折叠后 bf16 条件数误差累积)。但那是精度问题,不改变"flush 是 matmul、不是 replay"这 个事实。
8. flush 何时执行?——下一次 verify 时;early-flush 保证不溢出
问题:如果本次 verify 接受的 token 数量超过了 flush 的规定数量,需要 flush,是在下次 verify 的时候 matmul 吗?
结论:是的,flush 在下一次 verify 时执行。当前这一步的 commit 只负责打标记,真正的 matmul 折叠发生在下一个 forward 的 flush launch 里。
8.1 commit 设标记 / 下一步 verify 执行
看 commit kernel _advance_gdn_spec_cursors_kernel(:404-424):
1 | new_wp = write_pos + total_commit # 推进游标 |
commit 只更新 is_flush 标志位,没做任何 matmul。折叠在下一步:
1 | step t: |
注意 :115-132 的互斥:某个 slot 在一步里要么走 verify、要么走 flush。flush 分支(:270-284)不只折叠状态,还顺便用刚折好的新状态重构当前窗口的输出(:275 hw_q+= tl.dot(sw_tile, qT)、:352 写 out)。所以"flush 步"= 折叠 + 出输出,一次 kernel
pass 搞定。
8.2 “接受数超过 flush 规定数量”——其实不会发生(early-flush 预留余量)
这是设计的精髓。它提前一窗 flush,永远不让接受数把环撑爆:
:413-424 的注释和逻辑:
1 | # margin = 2 * max_spec_len, strict '>' |
- 触发条件不是"已经满了",而是"再来一窗就会满"——
new_wp + 2·max_spec_len >max_cache_len。 - config 强制
max_cache_len >= 2 * max_spec_len,保证任何一步 verify 都满足write_pos + spec_len <= max_cache_len,spec 窗口绝不溢出环。
所以不存在"这一步接受太多、超过 flush 上限、本步必须立刻 flush"的情况:环里永远预留了至少一整窗的空头,接受多少都装得下;装到逼近边界时,下一步自动转为 flush步把历史折进 checkpoint、游标归位。
为什么必须严格预留(注释
:416-423):proposer(MTP/n-gram)的 draft 数量不可控,拒绝采样要读每个窗口位置的 logits,若某位置溢出环就会喂给采样器一个未初始化的out垃圾 logit → 吐错 token + 状态失同步。所以宁可提前一窗 flush,零风险。
8.3 一句话
commit 只在游标逼近边界时打 is_flush 标记;真正的 flush matmul 折叠在下一次verify forward 的 flush launch 执行(同时产出那一步的输出)。而且因为 early-flush 预留了 ≥1 窗余量,单步接受数永远不会超出环容量——不存在"本步被迫 flush",溢出在结
构上被排除。
9. verify / flush 是互斥分支,不是"先 verify 再 flush"
问题:某个 slot 在一步里要么走 verify、要么走 flush 是啥意思?我理解 flush 是在verify 的基础上更新了 ssm,然后把 ring 清空,write_pos 置为 0?
校准:方向对,但有一个关键误解:flush 不是"在 verify 基础上额外做的一步",而是verify 的一个互斥分支。同一个 slot 在某一步里只会进其中一个分支,不会先 verify 再flush。
9.1 "要么 verify 要么 flush"的代码依据
gdn_replayssm_spec_decode 一步发两次 launch(launch_mode="both",:616-677):一次 IS_FLUSH=False(verify launch),一次 IS_FLUSH=True(flushlaunch)。但每个 slot 只在其中一次真正干活,另一次直接 return:
1 | # verify launch (IS_FLUSH=False), :130-132 |
所以由 is_flush[slot] 这个标志二选一路由:
is_flush==0的 slot → 只在 verify launch 干活,flush launch 里 return。is_flush==1的 slot → 只在 flush launch 干活,verify launch 里 return。
两次 launch 是为了让一个 batch 里不同 slot 走不同分支(有的该 flush 有的不该),不是对同一个 slot 连做两遍。
9.2 flush 步也产出输出(flush 分支自带 verify 的活)
“flush 是在 verify 基础上更新"这种感觉的来源——机制不是"verify 之后再 flush”,而是flush 分支自己把输出也算了。看 flush 分支 :270-284:
1 | sw_f = tl.dot(b_d_scaled, khist_tile, acc=b_total_decay*sc_tile) # 折叠新状态 |
flush 分支 = 先把环折进 checkpoint 得到新状态 sw_tile,再用这个新状态当基底重构当前这一窗的输出。后面 :350-352 照样写 out。所以 flush 步也产出这一步的attention 输出,不会漏掉。
区别只是基底不同:
- verify 分支:基底 = 旧
h0(:286 hw_q += tl.dot(sc_tile, qT),sc_tile是旧h0)+ 环历史。 - flush 分支:基底 = 折叠后的新
sw_tile(h0+环已经合一),不再单独叠环历史。
9.3 “清空 ring、write_pos 置 0”——校准
不是清空、不是置 0,是游标重定位。flush 后环里的折叠贡献已经进了 checkpoint,逻辑上"旧历史"作废,但物理上不擦数据,只移指针。看 commit kernel :407-410:
1 | new_base = (cache_base + write_pos) & (CACHE_BUF_LEN - 1) # 环起点前移到旧窗末尾 |
cache_base(环起点)前移write_pos格,跳过已折叠的旧历史——等效"清空",但靠移起点,不擦内存。write_pos不是置 0,而是设为total_commit(本步刚接受的 token 数)。因为flush 步当前这一窗的d/g也写进了环(新起点之后),它们成了折叠后新 checkpoint 的"新增量历史",得算进write_pos。
9.4 修正后的完整图(accept 3、本步触发 flush)
1 | flush 分支(一个 kernel pass 内全做完): |
9.5 一句话
“要么 verify 要么 flush” = 同一 slot 在这一步只走一个分支,由 is_flush 标志二选一路由,不是先 verify 再 flush。flush 分支自带输出重构(用折叠后的新状态当基底),所以它既更新 checkpoint 又产出本步输出。flush 后不是擦环/置 0,而是 cache_base前移(逻辑清空旧历史)、write_pos 设为本步提交数(留住新写入的增量)——纯指针操作,零拷贝。
10. cache_base 是什么?环形是怎么实现的
问题:这个 cache_base 是啥,ring 不是环形的吗?
结论:cache_base 就是这个环形 buffer 的逻辑起点指针——环是物理固定的一段内存,cache_base 标记"当前有效历史从环的哪个物理槽开始"。环形正是靠它 + 取模实现的。
10.1 为什么环形需要一个 base 指针
物理上 (d,k,g) cache 是一段定长内存,MAX_CACHE_LEN(2 的幂)个槽,槽号0..L-1。“环形"不是内存真的弯过来,而是用取模让索引绕回。要绕,就需要知道"起点在哪”,这就是 cache_base:
1 | # kernel :143-144 逻辑下标 → 物理槽,靠 base + 取模 |
o_c=0,1,2…是逻辑下标(“第几个历史 token”),用户视角连续。phys_c是物理槽号 =(cache_base + 逻辑下标) mod L。& (L-1)就是mod L(L 是 2 的幂),这一步实现"绕回"。
所以一条记录的物理位置 = 起点 cache_base + 偏移,超过 L 就绕到开头。环形 = base
- 取模,缺一不可。
10.2 三个量的分工
| 量 | 含义 | 类比 |
|---|---|---|
cache_base |
有效历史的物理起点槽号 | 环形队列的 head 指针 |
write_pos |
从起点算起,有效历史长度(也是下一个写入的逻辑位置) | 队列里的元素个数 |
phys = (base+offset)&(L-1) |
任意逻辑下标对应的物理槽 | 绕回后的真实地址 |
有效历史占的物理槽 = cache_base, cache_base+1, …, cache_base+write_pos-1(全部 modL)。
10.3 为什么 flush 是"移 base"而不是"清空"
flush 把 [cache_base, cache_base+write_pos) 这段旧历史折进 checkpoint 后,这段就作废了。但不擦内存,只把起点跳过去:
1 | # commit :407-410 |
旧段的物理槽被"逻辑遗弃",下次写入从新 cache_base 开始,绕一圈后自然覆盖它们。O(1) 指针移动,零拷贝——这正是环形 buffer 不需要搬数据的好处。
10.4 走个例子(L=8)
1 | 初始: cache_base=0, write_pos=0 环: [_ _ _ _ _ _ _ _] |
base 一直往前走、绕圈,有效历史是从 base 起、长度 wp 的一段滑动窗口。
10.5 一句话
环是物理定长内存,"环形"靠 (cache_base + 偏移) & (L-1) 取模绕回实现;cache_base 就是这个环形队列的 head 指针,标记有效历史的物理起点。flush 时把base 往前跳过已折叠的旧段(new_base=(base+write_pos)&(L-1)),等于 O(1) 清空旧历史、零拷贝,旧槽留待下次绕回时被覆盖。
附录 A:kernel 参数与 grid 速查
源码 python/sglang/srt/layers/attention/fla/gdn_replayssm_spec_decode.py。
A.1 gdn_replayssm_spec_circular_kernel(verify + flush,:39)
主要张量参数:
| 参数 | 形状 | 说明 |
|---|---|---|
q,k |
[total_tokens, H, K] |
post-conv,split 布局 |
v |
[total_tokens, HV, V] |
post-conv |
a,b |
[total_tokens, HV] |
gate / beta 原始输入 |
A_log, dt_bias |
[HV] fp32 |
GDN 静态权重 |
o |
[total_tokens, HV, V] |
预分配输出 |
h0 / ht |
[num_slots, HV, V, K] |
checkpoint(h0==ht,原地) |
d_cache |
[num_slots, HV, L, V] |
环形 d |
k_cache |
[num_slots, H, L, K] |
环形 k |
g_cache |
[num_slots, HV, L] fp32 |
环形 g |
query_start_loc |
[B+1] |
packed cu_seqlens |
ssm_state_indices |
[B] |
每请求物理块 |
write_pos / cache_base / is_flush_flags |
[num_slots] |
block-keyed 游标 |
constexpr:H, HV, K, V, BK, BV, BS, BC, NK, BKT, MAX_CACHE_LEN, SOFTPLUS_THRESHOLD,USE_QK_L2NORM_IN_KERNEL, IS_FLUSH, NULL_BLOCK_ID。
grid(:517):(triton.cdiv(V, BV), B, HV)——(V块, batch, value-head),无 layer轴。num_warps=1(默认)。
block 维度(:505-515):BK=npow2(K)、BKT=BK//nk(≥16 供 tl.dot)、
BV=min(npow2(V),64)、BS=max(bs_min, npow2(max_spec_len))(draft 窗)、
BC=max(16, npow2(max_cache_len))(环历史)。
A.2 游标 kernel
_advance_gdn_spec_cursors_kernel(:375):commit 一步推进write_pos +=num_accepted,计算next_is_flushearly-flush 标记,flush 时cache_base前移、write_pos归到本步提交数。grid(1,)。_reset_gdn_replayssm_spec_cursors_kernel(:432):prefill→decode 交接重置首步行游标(write_pos=0, cache_base=0, is_flush=INIT_FLUSH)。grid(1,)。
A.3 Python wrappers
gdn_replayssm_spec_decode(...)(:570):一步发两次 launch(verifyIS_FLUSH=False+ flushIS_FLUSH=True),device-side per-row 路由,CUDA-graph 可捕获。commit_gdn_replayssm_spec(...)(:681):每 decode step 调一次,推进游标。reset_gdn_replayssm_spec_cursors(...)(:714):重置首步行游标。
附录 B:PR #28695 元信息与改动文件
- 标题:
[GDN] Support ReplaySSM Ring Spec-Verify - 状态:open(截至记录时)
- head:
yuan-luo/sglang@fe66699d90295ad80aa85bae66ceae12d36a4509 - 分支:
gdn_replayssm_spec_decode - 父 RFC:issue #28511(Part B)
改动文件:
| 文件 | +/- | 作用 |
|---|---|---|
fla/gdn_replayssm_spec_decode.py |
+741 / 0 | 新 verify+flush 环形 kernel + 游标 kernel + wrappers |
linear/gdn_backend.py |
+117 / −16 | target-verify 在环已分配时 dispatch 到 replayssm kernel,否则回退递推 |
mem_cache/memory_pool.py |
+147 / −3 | 每层环 replayssm_{d,k,g} + 每 slot 游标加入 SpeculativeState(Optional,flag off 时 None),prefill 重置、COW 复制、RDMA 排除 |
model_executor/model_runner_kv_cache_mixin.py |
+2 / 0 | 把 enable_gdn_replayssm_spec 接到 pool 初始化 |
server_args.py |
+23 / 0 | --enable-gdn-replayssm-spec 开关 + config 校验(max_cache_len >= 2*spec_num_draft) |
speculative/spec_utils.py |
+57 / 0 | commit_mamba_states_after_verify:每步推进游标(含 bonus)替代拷中间态;conv-state accept-rollback 仍跑 |
B.1 正确性(PR 自述)
| 层级 | 方法 | 结果 |
|---|---|---|
| kernel 数学 | vs chunk_gated_delta_rule(GDN prefill 权威) |
~5e-4,draft_token_num 1–8 真实维度 |
| 生命周期 | 多步环 + commit + flush + reset vs chunk | bounded ~5e-3(bf16,flush 处重锚) |
| 端到端 | GSM8K(400q),Qwen3.5-35B-A3B TP4 NEXTN,CUDA-graph on | 0.892(on)vs 0.887(off)——噪声内,0 invalid |
B.2 性能(PR 自述,Qwen3.5-35B-A3B TP4 NEXTN,draft=4,256in/512out)
| concurrency | 输出吞吐 Δ | median TPOT Δ |
|---|---|---|
| 8 | +6.5% | −8.0% |
| 16 | +5.6% | −5.3% |
| 32 | +13.1% | −11.8% |
带宽利用(决策门槛,H20-3e peak ~3.85 TB/s,递推 GDN verify):batch 1→6.7%,16→60.8%,64→74.7%,256→79.4%。≥70%@batch≥64 → 全状态写摊销是真加速。
另:递推基线在 --mem-fraction-static 0.8 OOM,replayssm 可跑(~12 GBintermediate_ssm per-draft 快照 vs 小环)——固定 HBM 预算下 replayssm 撑更高并发。
B.3 长输出非无损(PR 自述)
长输出(16k+,AIME-2024,box-243 TP4,ctx 40960,max_tokens 32768,greedy):
| verify | AIME-2024 | avg out tok | truncations |
|---|---|---|---|
| recurrent(基线) | 0.933 | 16001 | 1 |
| replayssm(本 PR,bf16) | 0.767 | 18244 | 8 |
根因(结构性,非精度):GDN 递推 是收缩的( 方向特征值 ,其余 ),误差衰减、数值稳定。ReplaySSM 改用 chunked UT-transform 推进 committed 状态,其交错级数 在 intra-ring delta-rule 交互 大时灾难性抵消。每折叠的条件数误 差在长解码(16k≈1000 折)累积 → 状态漂移 → 退化重复循环。设计缺陷是同一个病态 既算输出又推进 checkpoint。bf16 放大 ~1000×(每次 flush 把整个checkpoint round-trip 过 bf16)。
11. 全文一句话总结
PR #28695 用环形 (d,k,g) 缓存 + 冻结 checkpoint 替换 GDN target-verify 的per-draft 全状态快照:token 轴用 (I+A)^{-1} UT-transform 一次 matmul 并行重构整窗输出(赢),layer 轴因重构=forward 的链式依赖无法并行(输,与 KV-buffer replay
互补);接受/拒绝是 write_pos += num_accepted 一个游标加法(零成本回退,拒绝draft 从未转正即等于回退);flush 是 S_new=A·h0+tl.dot(d,k) 一次 matmul 折叠(不是 replay)而非串行递推,在下一步 verify 的 flush 分支执行(与 verify 互斥路由);early-flush 预留 2·max_spec_len 余量保证环永不溢出;cache_base 是环形 head 指针,flush 时 O(1) 前移、零拷贝清空旧历史。短输出数学无损(GSM8K 持平),长输出因(I+A)^{-1} 折叠的灾难性抵消非无损 → 默认关闭、仅 GDN linear-chain、短输出 scoped。
12. 补充数学:corrected delta 、 求逆的来源与 closed-loop 修复
本节补记对话后段关于"为什么 verify 要求逆、decode 不用、病态 是什么、怎么修"的逐题讲解,是对 §7/§8/§B.3 的数学展开。
12.1 corrected delta 是什么
delta-rule 线性注意力的单步状态更新:
就是"这一步真正往 里写了什么"的那个 rank-1 因子。输出和状态推进都靠它。ReplaySSM 的环里存的正是 (而不是原始输入),重建公式:
12.2 为什么 decode(Part A)算 不需要求逆
decode 窗口 = 1。那一步 : 是已知的(上一步算完了),所以 是一次 matvec + 标量缩放,直接出,无未知量互相依赖,不解方程、不求逆。这就是 #28451 decode kernel 缓存 却从不出现 的原因——窗口=1 是退化情形。
12.3 为什么 verify(Part B)必须求逆——m 个 的耦合
verify 窗口 = m,窗口内顺序递推: 依赖含 的 , 依赖含 的 ……把 沿窗口展开代入,每个 都依赖前面所有 $d_{j
写成矩阵就是一个严格下三角线性系统:
严格下三角 ⇒ 单位下三角、可逆,逆 也是单位下三角(UTtransform)。于是 (代码 D_spec = R @ Tᵀ,:337)——一次 matmul 并行解出窗口全部 m 个耦合的 。这是 chunked delta-rule 能并行的核心,也是 的唯一来源。
12.4 m 个未知 是怎么一步步得到的(手算,m=3)
串行视角 = 前向代入(forward substitution):
并行视角 = 直接乘 :因 幂零(),逆是有限项级数
| 视角 | 怎么解 | 复杂度形态 | 数值性质 |
|---|---|---|---|
| 串行 forward-substitution | 逐行代入 | 串行 | 每步收缩,稳定 |
| 并行 (UT transform) | 一次 tl.dot |
张量核并行 | 可能病态 |
ReplaySSM 选并行版吃满张量核,代价是把数值稳定性押在 的条件数上。
12.5 病态的 是什么、两层误差链
是交错级数。当 intra-ring 的 大时 元素大,交错相加出现 catastrophic cancellation(大数相减丢有效位),算出的 带误差——这就是"病态 "。注意两层误差链,别混:
- 误差源:
D_spec = R @ Tᵀ(:337),即 计算 。这里才有抵消。 - 误差载体/消费者:flush 的 (
:271)。它在代数上无损(只是把 L 个 rank-1 攒成一次并行tl.dot吃带宽红利),只是把已经错了的 折进 checkpoint。
所以" 用 体系一次性算出来"要拆开: 本身是"L 步交互攒成一次并行matmul"的无损折叠; 是在算 时出现的有损那一步。二者不是一回事。
12.6 为什么长输出会塌(结构性,非纯精度)
- 真实 GDN 递推是收缩的( 方向特征值 ,其余 ),误差衰减、稳定。
- 但 折叠不收缩,且同一个病态 既算输出又推进 checkpoint。
- 每次 flush 把误差写进 ;16k token ≈ ~1000 折,误差累积 → 状态漂移 → 复读环。
- bf16 把误差再放大 ~1000×(每次 flush 整个 checkpoint round-trip 过 bf16),但全fp32 也只能把 AIME —— 证明主因是结构(不收缩的折叠),不是精度。
12.7 closed-loop exact fold 修复
核心思想:把"算输出"和"推进状态"解耦。
- 输出:仍用 chunked (bf16 容忍——输出错了只影响当步比对,不累积进checkpoint)。
- 状态推进:改用精确、收缩的真实递推。为此额外存 和 ,flush 时用原始 重新计算 (而非复用病态 ),让 checkpoint 沿真实 GDN 收缩递推前进,误差自然衰减。
效果:AIME 。
12.8 权衡:要么无损、要么 +13%
精确递推天然串行(状态推进不能并行折叠),掉 ~15% 吞吐。结论是结构性的:
在这套 ReplaySSM verify 框架下,无损与 +13% 吞吐不可兼得——并行折叠(快但病态)与串行精确递推(无损但慢)二选一。
| 方案 | AIME(16k 长输出) | 吞吐 | 适用 |
|---|---|---|---|
| 纯 chunked (现状) | 0.767 | +13% | 短输出,对长输出质量不敏感 |
| 全 fp32 chunked | 0.800 | ≈ 持平/略降 | 仅证明病态是结构性的 |
| closed-loop exact fold | 0.867 | −15%(相对现状) | 长输出,要质量 |
12.9 与 Part A(#28451 decode)的对照
| Part A decode(窗口=1) | Part B verify(窗口=m) | |
|---|---|---|
| 怎么算 | , 已知,直接出 | m 个 耦合,解 |
| 是否求逆 | 否 | 是( UT transform) |
| 数值稳定 | 沿真实收缩递推,稳 | 交错级数可能病态 |
| 病态来源 | 无 | 窗口内耦合 + 不收缩折叠 |
两者共享同一套数据结构( checkpoint + 环 + 三游标),全部差异都来自窗口长度 1 vs m。
- 标题: replayssm(B)SpecDecode
- 作者: 鱿鱼圈
- 创建于 : 2026-06-30 22:30:00
- 更新于 : 2026-06-30 22:31:19
- 链接: https://yuyanqi.com/2026/06/30/replayssm(B)SpecDecode/
- 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。