ggaaooppeenngg

为什么计算机科学是无限的但生命是有限的

线性注意力与 SSM:两条技术路线的完整推导

线性注意力和 SSM(State Space Model,状态空间模型)是序列建模的两条技术路线,也是后来 KDA/GDN/DeltaNet 一族模型的两个源头:线性注意力贡献了「外积记忆状态 + 结合律」,SSM 贡献了「衰减门 + 递推骨架」。本文把这两条路线的计算过程完整推一遍,每步都有推导、数值验算和工程上的理由。

0. 为什么要固定大小的记忆状态

标准 softmax 注意力在生成文本时,必须把每个历史 token 的 Key 和 Value 都存下来(KV cache)–因为每生成一个新 token,都要和前面所有 token 重新算一遍注意力。序列越长,这笔账越贵:

开销 softmax 注意力(KV cache) 固定状态(线性注意力一族)
解码显存 O(nd)O(n \cdot d)(全部历史的 KV 都得缓存) O(d2)O(d^2)(一个固定矩阵 S,与 n 无关)
单步解码 O(nd)O(n \cdot d)(扫全部历史) O(d2)O(d^2)(更新 S + 读出,两次矩阵乘)
预填充 O(n2d)O(n^2 d) O(nd2)O(n d^2)
上下文长度翻倍 显存、延迟一起翻倍 完全不变

(单头视角,d 为特征维度。)

几十万 token 的上下文就能把 KV cache 撑到几十 GB。想摆脱这笔账,就得把「全部历史的 KV」压缩成一个固定大小的状态,而且这个状态要会写、会改、会忘–两个源头各给出了一半答案:线性注意力先把固定状态造出来,SSM 教会它怎么遗忘。


1. 路线一:线性注意力–两个前提,缺一不可

线性注意力能成立,靠的是两件事:去掉 softmax 的非线性让 Q/K 先结合掉。更准确地说,这是同一枚硬币的两面–softmax 必须作用在耦合后的 QKQK^\top 上,它既是非线性的来源,也是耦合的来源;把它换成可分离的 ϕ(q)ϕ(k)\phi(q)^\top\phi(k),非线性和耦合一起消失,结合律才重新可用。后面 KDA 一族的全部推导都建立在这个前提上,值得把两个条件一个一个讲透。

1.1 前提一:去掉 softmax,结合律才可用

先看没有 softmax 的裸注意力 Attn(Q,K,V)=(QK)V\mathrm{Attn}(Q, K, V) = (QK^\top)V,复杂度账本:

括号化 计算路径 复杂度 主导项
(QK)V(QK^\top)V 先算出 n×nn \times n 矩阵,再乘 V O(dn2)+O(nnd)=O(dn2)O(d\cdot n^2) + O(n\cdot nd) = O(dn^2) n(平方)
Q(KV)Q(K^\top V) 先算 KVRd×dK^\top V \in \mathbb{R}^{d \times d}(与 n 无关),再左乘 Q O(nd2)+O(dnd)=O(nd2)O(n\cdot d^2) + O(d\cdot nd) = O(nd^2) d(线性)

同一个数学式,两种括号化,复杂度差一个 n 的幂–这就是结合律的诱惑:矩阵乘法满足 (AB)C=A(BC)(AB)C = A(BC),先算哪边是自由的。裸注意力(线性代数层面)本来就能线性化。

那 softmax 挡在哪? 标准注意力是 softmax(QK)V\mathrm{softmax}(QK^\top)V,softmax 在 QKQK^\top 之后、乘 V 之前介入,并且它是行内非线性:第 i 行(第 i 个 query 对所有 key 的打分)必须整体过 exi/jexje^{x_i}/\sum_j e^{x_j}–指数逐项非线性 + 行内归一化耦合。两个性质各堵死一条路:

  1. 非线性破坏结合律:softmax 不是线性算子,softmax(QK)VQsoftmax(KV)\mathrm{softmax}(QK^\top)V \ne Q\,\mathrm{softmax}'(K^\top V),中间结果没法先合并;n×nn \times n 矩阵必须先完整算出来;
  2. 归一化引入全局耦合:分母 jeqikj\sum_j e^{q_i \cdot k_j} 依赖该行所有 n 个 key,哪怕只想算一个 query 的输出,也得先把整行算完–信息在 key 维度上全局耦合,没有可以「先结合掉」的独立块。

所以平方复杂度不是矩阵乘法的错,是 softmax 的作用位置的错:它非要等 Q 和 K 耦合完才动手。要线性化,就得把这个非线性从耦合点上搬走。

1.2 前提二:把非线性提前到 Q 和 K 各自身上

搬法就是核化(kernelization)。把注意力抽象成通用相似度函数 sim(q,k)\mathrm{sim}(q, k)–它不必是 softmax 的指数,多项式相似度、RBF 核都属此类。数学依据是核技巧:只要 sim 非负(Mercer 条件),就存在特征映射 ϕ()\phi(\cdot) 使

sim(q,k)=ϕ(q)ϕ(k)\mathrm{sim}(q, k) = \phi(q)^\top \phi(k)

注意这个形式的本质:非线性被吸收进 ϕ\phi,分别作用于 q 和 k 各自身上,相似度本身变成线性的内积。softmax 的 eqke^{q \cdot k} 作用在耦合后的标量上;核化的 ϕ(q)ϕ(k)\phi(q)^\top\phi(k) 让 q 和 k 各自先做完所有非线性变换,最后只留一次线性内积。

线性注意力取 ϕ(x)=ELU(x)+1\phi(x) = \mathrm{ELU}(x) + 1(Katharopoulos et al. 2020)。先把 ELU(Exponential Linear Unit,指数线性单元)本身说清楚,它是个分段函数:

ELU(x)={xx>0ex1x0\mathrm{ELU}(x) = \begin{cases} x & x > 0 \\ e^{x} - 1 & x \le 0 \end{cases}

  • 正数区:原样输出(和 ReLU 一样);
  • 负数区:输出 ex1(1,0)e^x - 1 \in (-1, 0),一条平滑曲线,越负越接近 -1 但永远不到 -1。

所以 ELU(x)+1\mathrm{ELU}(x) + 1 的取值范围是 (0,)(0, \infty):恒正。算两个具体值感受一下:x=0x = 0ELU(0)+1=0+1=1\mathrm{ELU}(0)+1 = 0 + 1 = 1;x=3x = -3e31+1=e30.05e^{-3} - 1 + 1 = e^{-3} \approx 0.05,很小但仍是正数。这个恒正不是自选的装饰,是核分解存在的条件(Mercer 条件要求相似度非负):ϕ(q)ϕ(k)\phi(q)^\top\phi(k) 是 d 个正数乘正数再求和,结果必为正,才能扮演「打分」的角色。于是:

(ϕ(Q)ϕ(K))V=ϕ(Q)(ϕ(K)V)\big(\phi(Q)\,\phi(K)^\top\big)V = \phi(Q)\big(\phi(K)^\top V\big)

结合律在非线性世界里重新可用:先算 ϕ(K)VRd×d\phi(K)^\top V \in \mathbb{R}^{d \times d},复杂度回到 O(nd2)O(nd^2)两个前提到此汇成一句话:非线性提前,耦合消失,结合律重新可用。写成逐 token 的归一化形式(softmax 的归一化也没丢,变成显式分母):

Attn(q)=iϕ(q)ϕ(ki)viiϕ(q)ϕ(ki)=ϕ(q)Sϕ(q)z,S=iϕ(ki)vi\mathrm{Attn}(q) = \frac{\sum_i \phi(q)^\top \phi(k_i)\, v_i}{\sum_i \phi(q)^\top \phi(k_i)} = \frac{\phi(q)^\top S}{\phi(q)^\top z}, \qquad S = \sum_{i} \phi(k_i)\, v_i^\top

两个新符号 S 和 z 分别是:

  • S=iϕ(ki)viRd×dS = \sum_i \phi(k_i)v_i^\top \in \mathbb{R}^{d \times d}:分子里的 KV 外积矩阵,存「key-value 关联」的记忆本体;
  • z=iϕ(ki)Rdz = \sum_i \phi(k_i) \in \mathbb{R}^{d}:分母里的 key 累积和,就是 softmax 分母 ieqki\sum_i e^{q\cdot k_i} 的核化版–归一化因子。

先说清 z,因为后续论文里它最容易被略写。z 在公式里承担的角色就是归一化:没有它,输出会随序列变长而无界增长。但后面 DeltaNet/GDN/KDA 的论文里,递推式往往只写 SS 的更新,z 要么藏在一句「输出再过 RMSNorm」的描述里,要么干脆不提–因为实测发现把归一化换成 RMSNorm 效果更好,z 就被简化掉了。所以读者常遇到的困惑是:看 KDA 论文时公式里根本没有 z,翻早期线性注意力文献才发现 z 是这里从 softmax 分母继承下来的归一化项。一句话:z = 归一化分母的累积状态,S = 关联记忆本体;后续模型改用 RMSNorm 后 z 退场,但 S 的更新规则一路演进到 KDA。

分子分母都变成对 S 的查询,S 一遍扫过序列累积即可–复杂度对 n 线性。更关键的是 S 可以写成递推,一步一读出(z 同样逐步累积):

St=St1+ϕ(kt)vt,ot=ϕ(qt)Stϕ(qt)ztS_t = S_{t-1} + \phi(k_t)\, v_t^\top, \qquad o_t = \frac{\phi(q_t)^\top S_t}{\phi(q_t)^\top z_t}

这个递推还有一个更形象的理解,叫 fast weights 视角(Schlag et al., Linear Transformers Are Secretly Fast Weight Programmers):把 S 看作一块快速权重–模型真正的参数(普通权重)训练完就固定了,而 S 每来一个 token 就被改写一次,相当于一个「随输入不断更新的小型权重」。写入靠 ϕ(kt)vt\phi(k_t)v_t^\top(每个 token 都在给这块权重编程),读出靠查询 ϕ(qt)St\phi(q_t)^\top S_t(拿当前问题去这块权重里查答案)。所以这个视角下,序列处理 = 一边更新权重、一边用权重答题。

回过头总结一下,这两个前提各自换来了什么:

前提 换来的直接好处 对后续模型的影响
抛弃 softmax 非线性 结合律可用,O(n2)O(n)O(n^2) \to O(n) 注意力矩阵 n×nn\times n 消失,换成固定大小状态 SRd×dS \in \mathbb{R}^{d \times d}
非线性提前到 Q/K 各自身上 ϕ(q)ϕ(k)\phi(q)^\top\phi(k) 可分离 S 可以递推累积 -> RNN 化 -> fast weights -> 一切后续演化的载体

注意代价也在表里:换来的 S 只会加法。这就引出了线性注意力最大的问题。

1.3 纯加性状态的问题:记忆碰撞

纯加性状态的核心问题:state 是纯加性的(St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top),写入只有「叠加」,没有「删除」。序列长度远超 state 有效容量(d×dd \times d 矩阵只能存这么多关联)时,不同 kvk \to v 关联互相干扰,旧信息永远赖在状态里,新信息无法覆盖–记忆像一个只进不出的仓库。

数值演示。设 dk=dv=2d_k = d_v = 2,S 从零开始,依次写入四个 token:

token key value 说明
t=1 k1=[1,0]k_1 = [1, 0] v1=[1,0]v_1 = [1, 0]
t=2 k2=[0.6,0.8]k_2 = [0.6, 0.8] v2=[0,1]v_2 = [0, 1]
t=3 k3=[1,0]k_3 = [1, 0] v3=[2,0]v_3 = [2, 0] k1k_1 同一个 key,写入新值
t=4 k4=[0,1]k_4 = [0, 1] v4=[3,1]v_4 = [3, 1]

线性注意力这边,把每一步都算出来。写入规则是 St=St1+vtktS_t = S_{t-1} + v_t k_t^\top,每个 vkv k^\top 是一个 2×2 外积。

为什么是 vkv k^\top 而不是 kvk v^\top? 外积的顺序决定「谁进谁出」。vkv k^\top 是一个从 key 空间到 value 空间的线性映射:key 从右边进,value 从左边出–直接验证:(vk)k=v(kk)(v k^\top)\,k = v\,(k^\top k),查询 key 与存储 key 的内积落在系数上,value 原样输出。如果反过来存 kvk v^\top 还从右边乘,得到 k(vk)k\,(v^\top k)–变成「拿 key 查、按 value 相似度返回 key」,角色整个反了,查不出想要的东西。要让 kvk v^\top 布局工作,查询必须从左边进:行向量 qS=i(qki)viq^\top S = \sum_i (q^\top k_i)\,v_i^\top,结论完全一样。所以两种写法都存在:§1.2 的 S=iϕ(ki)viS = \sum_i \phi(k_i)v_i^\top 配左乘读出 ϕ(q)S\phi(q)^\top S、KDA 论文的 SS 配读出 SqS^\top q,是 kvk v^\top 约定;本文手算例子和 DeltaNet 一节用 vkv k^\top 约定、右乘读出 SqS\,q。两个约定下的 S 互为转置,全部结论不变–只是读论文时看到外积顺序翻转不要慌,先看它查询是从哪边乘的。

t=1:v1k1=[10][10]=[1000]v_1 k_1^\top = \begin{bmatrix}1\\0\end{bmatrix}\begin{bmatrix}1&0\end{bmatrix} = \begin{bmatrix}1&0\\0&0\end{bmatrix},所以

S1=[1000]S_1 = \begin{bmatrix}1&0\\0&0\end{bmatrix}

t=2:v2k2=[01][0.60.8]=[000.60.8]v_2 k_2^\top = \begin{bmatrix}0\\1\end{bmatrix}\begin{bmatrix}0.6&0.8\end{bmatrix} = \begin{bmatrix}0&0\\0.6&0.8\end{bmatrix},叠加后

S2=[100.60.8]S_2 = \begin{bmatrix}1&0\\0.6&0.8\end{bmatrix}

t=3(关键一步,同一个 key 写入新值):v3k3=[20][10]=[2000]v_3 k_3^\top = \begin{bmatrix}2\\0\end{bmatrix}\begin{bmatrix}1&0\end{bmatrix} = \begin{bmatrix}2&0\\0&0\end{bmatrix}注意它不清除 S2S_2 里已有的 [1000]\begin{bmatrix}1&0\\0&0\end{bmatrix},直接往上叠:

S3=[300.60.8]S_3 = \begin{bmatrix}3&0\\0.6&0.8\end{bmatrix}

第一行变成了 [3, 0]–旧值 1 和新值 2 加在一起,这就是碰撞发生的瞬间。

t=4:v4k4=[31][01]=[0301]v_4 k_4^\top = \begin{bmatrix}3\\1\end{bmatrix}\begin{bmatrix}0&1\end{bmatrix} = \begin{bmatrix}0&3\\0&1\end{bmatrix},叠加后

S4=[330.61.8]S_4 = \begin{bmatrix}3&3\\0.6&1.8\end{bmatrix}

接下来演示「查询」–先说这个操作是什么。写入是每个 token 做一次 S+=vkS \mathrel{+}= v k^\top,把四步叠起来,S4S_4 其实就是四个外积之和:

S4=v1k1+v2k2+v3k3+v4k4S_4 = v_1k_1^\top + v_2k_2^\top + v_3k_3^\top + v_4k_4^\top

查询的定义:拿一个向量 kk右乘 S,即 SkS\,k。为什么这样就是「查询」?把上式代入 S4k1S_4 k_1:

S4k1=v1(k1k1)+v2(k2k1)+v3(k3k1)+v4(k4k1)S_4 k_1 = v_1(k_1^\top k_1) + v_2(k_2^\top k_1) + v_3(k_3^\top k_1) + v_4(k_4^\top k_1)

每个 token 的 value 前面多了一个系数–它存入时的 key 与查询 key 的内积。key 完全匹配(内积=1)value 完整取回;key 正交(内积=0)完全不干扰;部分对齐就按比例混进来一份。这就是「按 key 相似度加权取回 value」,也正是注意力打分的雏形(softmax 注意力把内积换成 eqke^{q\cdot k},思路相同)。

代入数字k1k1=1k_1^\top k_1 = 1,k2k1=0.6k_2^\top k_1 = 0.6,k3k1=1k_3^\top k_1 = 1,k4k1=0k_4^\top k_1 = 0:

S4k1=1×v1+0.6×v2+1×v3+0×v4=[1,0]+[0,0.6]+[2,0]=[30.6]S_4 k_1 = 1{\times}v_1 + 0.6{\times}v_2 + 1{\times}v_3 + 0{\times}v_4 = [1,0] + [0, 0.6] + [2,0] = \begin{bmatrix}3\\0.6\end{bmatrix}

(直接用矩阵乘验证:取 S4S_4 第一列,同样是 [3,0.6][3, 0.6]^\top。)

拆开看这个结果:[3.0, 0.6] = 旧值 v1=[1,0]v_1=[1,0] + 新值 v3=[2,0]v_3=[2,0] 各自完整叠加(两个「系数 1」都命中),再加 0.6 份 v2v_2(k2k_2k1k_1 内积 0.6,部分对齐,混进来 0.6 个 [0,1])。想读「key=[1,0] 对应的最新值」,读出来的却是一锅大杂烩。

delta rule 这边,同样四步。写入规则换成「先减旧值再加新值」:St=St1+(vtSt1kt)ktS_t = S_{t-1} + (v_t - S_{t-1}k_t)k_t^\top(取 β=1\beta = 1):

  • t=1:S0k1=0S_0 k_1 = 0,写入 v1k1v_1 k_1^\top,得 S1=[1000]S_1 = \begin{bmatrix}1&0\\0&0\end{bmatrix}(和线性注意力相同);
  • t=2:先查旧值 S1k2=[0.6,0]S_1 k_2 = [0.6, 0]^\top(k2=[0.6,0.8]k_2 = [0.6, 0.8] 在第一维有 0.6 的分量,所以部分命中第一列)。写入差值 u2=v2S1k2=[0.6,1]u_2 = v_2 - S_1k_2 = [-0.6, 1],外积 u2k2=[0.360.480.60.8]u_2 k_2^\top = \begin{bmatrix}-0.36&-0.48\\0.6&0.8\end{bmatrix}。注意左上角是负数–它在把 k2k_2 方向上已存的 v1v_1 残余往回抠,不是盲目叠加。得 S2=[0.640.480.60.8]S_2 = \begin{bmatrix}0.64&-0.48\\0.6&0.8\end{bmatrix};
  • t=3(同一个 key 写新值):先查旧值 S2k3=[0.64,0.6]S_2 k_3 = [0.64, 0.6]^\top。理想情况这里应该读出 [1,0][1, 0](t=1 写入的 v1v_1),但读出的被 t=2 的写入带着偏了。写入差值 u3=v3S2k3=[1.36,0.6]u_3 = v_3 - S_2k_3 = [1.36, -0.6],外积只动第一列方向,把 k3k_3 方向修正到 v3v_3。得 S3=[20.4800.8]S_3 = \begin{bmatrix}2&-0.48\\0&0.8\end{bmatrix}–第一行回到 2,新值到位;
  • t=4:先查 S3k4=[0.48,0.8]S_3 k_4 = [-0.48, 0.8]^\top,写入差值 u4=v4S3k4=[3.48,0.2]u_4 = v_4 - S_3k_4 = [3.48, 0.2],只修正 k4k_4 方向。得 S4=[2301]S_4 = \begin{bmatrix}2&3\\0&1\end{bmatrix}

最后查询 k1=[1,0]k_1 = [1, 0]:线性注意力给出 [3.0,0.6][3.0, 0.6](新旧叠加 + 残余污染),delta rule 给出 [2.0,0.0][2.0, 0.0]精确返回最新的 v3v_3

同一个 key 写两次,线性注意力读出的是两个 value 的叠加,delta rule 读出的是新的那个–这就是「记忆碰撞」(memory collision)。修复它(delta 写入、衰减门)就是后续模型的故事,见《KDA 的来龙去脉》


2. 路线二:SSM–从微分方程到遗忘门

线性注意力造出了固定状态,但这个状态只会加法。SSM 是从另一个领域来的–控制理论的状态空间模型–它恰好能精确回答「状态该怎么衰减」,最终和线性注意力在 Mamba-2 处汇合。

2.1 连续 SSM 的定义

单输入单输出(SISO)线性时不变(LTI)状态空间模型:

{h(t)=Ah(t)+Bx(t)(状态方程)y(t)=Ch(t)+Dx(t)(输出方程)\begin{cases} h'(t) = A\,h(t) + B\,x(t) & \text{(状态方程)}\\[4pt] y(t) = C\,h(t) + D\,x(t) & \text{(输出方程)} \end{cases}

符号 维度 含义
x(t)Rx(t) \in \mathbb{R} 标量 输入信号
h(t)RNh(t) \in \mathbb{R}^N N 维向量 状态:到 t 为止历史的压缩
y(t)Ry(t) \in \mathbb{R} 标量 输出
ARN×NA \in \mathbb{R}^{N\times N} 矩阵 状态转移:旧记忆如何演化/衰减
BRN×1B \in \mathbb{R}^{N\times 1} 向量 输入到状态的写入强度
CR1×NC \in \mathbb{R}^{1\times N} 向量 状态到输出的读出权重
DRD \in \mathbb{R} 标量 直通项(深度学习中常设 0,以下略去)

要解决的问题:神经网络处理的是离散序列 x1,x2,x_1, x_2, \dots,需要把微分方程改写为递推式 ht=Aˉht1+Bˉxth_t = \bar A h_{t-1} + \bar B x_t,并求出 Aˉ,Bˉ\bar A, \bar BA,BA, B 的精确关系。

2.2 预备知识:矩阵指数

定义(对标量指数的泰勒级数的直接推广):

eM:=k=0Mkk!=I+M+M22!+M33!+e^{M} := \sum_{k=0}^{\infty} \frac{M^k}{k!} = I + M + \frac{M^2}{2!} + \frac{M^3}{3!} + \cdots

本推导用到的三条性质:

  1. 微分:ddteAt=AeAt=eAtA\dfrac{d}{dt}e^{At} = A\,e^{At} = e^{At}A(与标量 eate^{at} 求导完全平行);
  2. 交换性:AAeAte^{At} 可交换(因为 eAte^{At}AA 的幂级数);
  3. :(eAt)1=eAt(e^{At})^{-1} = e^{-At}(由 eAteAt=eA(tt)=Ie^{At}e^{-At} = e^{A(t-t)} = I)。

直觉:eAΔe^{A\Delta} 是「让系统按自身动力学自由演化 Δ 时间」的算子。A 的特征值实部为负时,它就是各种速度衰减的混合。

2.3 齐次方程的解(无输入情形)

命题:若 x(t)0x(t) \equiv 0,则 h(t)=Ah(t)h'(t) = Ah(t) 的解为

h(t)=eA(tt0)h(t0)h(t) = e^{A(t - t_0)}\, h(t_0)

证明:直接验证满足方程与初值。令 h(t)=eA(tt0)h(t0)h(t) = e^{A(t-t_0)}h(t_0),则

h(t)=ddt[eA(tt0)]h(t0)=AeA(tt0)h(t0)=Ah(t).h'(t) = \frac{d}{dt}\Big[e^{A(t-t_0)}\Big]h(t_0) = A\,e^{A(t-t_0)}h(t_0) = A\,h(t). \quad\blacksquare

含义:没有输入时,旧记忆按 eAΔe^{A\Delta} 自然演化–这是通解第一项的来源。

2.4 通解推导:积分因子法

定理:h(t)=Ah(t)+Bx(t)h'(t) = A h(t) + B x(t) 满足初值 h(t0)h(t_0) 的解为

h(t)=eA(tt0)h(t0)+t0teA(ts)Bx(s)dsh(t) = e^{A(t-t_0)}\,h(t_0) + \int_{t_0}^{t} e^{A(t-s)}\,B\,x(s)\,ds

证明(积分因子法,四步):

第 1 步:移项。 把含 h 的项移到左边:

h(t)Ah(t)=Bx(t)h'(t) - A\,h(t) = B\,x(t)

第 2 步:乘积分因子 eAte^{-At} 两边左乘 eAte^{-At}:

eAth(t)eAtAh(t)=eAtBx(t)e^{-At}h'(t) - e^{-At}A\,h(t) = e^{-At}B\,x(t)

第 3 步:识别乘积导数。 由矩阵指数的微分性质,

ddt[eAth(t)]=eAtAh(t)+eAth(t)\frac{d}{dt}\Big[e^{-At}h(t)\Big] = -e^{-At}A\,h(t) + e^{-At}h'(t)

恰好等于左边。于是方程变成:

ddt[eAth(t)]=eAtBx(t)\frac{d}{dt}\Big[e^{-At}h(t)\Big] = e^{-At}B\,x(t)

这一步是整个推导的「机关」:乘以 eAte^{-At} 后,左端塌缩成一个全导数,方程立刻可积。

第 4 步:两边积分并整理。t0t_0tt 积分:

eAth(t)eAt0h(t0)=t0teAsBx(s)dse^{-At}h(t) - e^{-At_0}h(t_0) = \int_{t_0}^{t} e^{-As}B\,x(s)\,ds

两边左乘 eAte^{At}(注意 eAteAt0=eA(tt0)e^{At}e^{-At_0} = e^{A(t-t_0)},且 eAteAs=eA(ts)e^{At}e^{-As} = e^{A(t-s)}):

h(t)=eA(tt0)h(t0)+t0teA(ts)Bx(s)dsh(t) = e^{A(t-t_0)}h(t_0) + \int_{t_0}^{t} e^{A(t-s)}B\,x(s)\,ds \quad\blacksquare

两项的物理含义:旧记忆自由演化 + 区间内每一瞬输入贡献的叠加(叠加原理)。

2.5 零阶保持(ZOH)离散化

零阶保持假设:在采样区间 [tk, tk+Δ)[t_k,\ t_k + \Delta) 内,输入保持为采样值:

x(tk+τ)=xk,τ[0,Δ)x(t_k + \tau) = x_k, \qquad \forall\,\tau \in [0, \Delta)

推导:在通解中取 t0=tkt_0 = t_k,t=tk+Δt = t_k + \Delta:

hk+1=eAΔhk+0ΔeA(Δτ)Bx(tk+τ)dτh_{k+1} = e^{A\Delta}h_k + \int_{0}^{\Delta} e^{A(\Delta-\tau)}B\,x(t_k + \tau)\,d\tau

由 ZOH 假设,x(tk+τ)=xkx(t_k+\tau) = x_k 是常数,提出积分号:

hk+1=eAΔAˉhk+(0ΔeA(Δτ)Bdτ)Bˉxkh_{k+1} = \underbrace{e^{A\Delta}}_{\bar A}\,h_k + \underbrace{\left(\int_0^{\Delta} e^{A(\Delta-\tau)}B\,d\tau\right)}_{\bar B}\,x_k

计算 Bˉ\bar B 的积分(标量情形最直观;矩阵情形在 A 可逆时同样成立)。换元 s=Δτs = \Delta - \tau:

Bˉ=0ΔeAsBds\bar B = \int_0^{\Delta} e^{As}\,B\,ds

对级数逐项积分:

0ΔeAsds=0Δk=0(As)kk!ds=k=0AkΔk+1(k+1)!=A1(eAΔI)\int_0^{\Delta} e^{As}\,ds = \int_0^{\Delta}\sum_{k=0}^{\infty}\frac{(As)^k}{k!}ds = \sum_{k=0}^{\infty}\frac{A^k\Delta^{k+1}}{(k+1)!} = A^{-1}\big(e^{A\Delta} - I\big)

最后一步验证:A1(eAΔI)=A1k1(AΔ)kk!=k1Ak1Δkk!=j0AjΔj+1(j+1)!A^{-1}(e^{A\Delta} - I) = A^{-1}\sum_{k\ge1}\frac{(A\Delta)^k}{k!} = \sum_{k\ge1}\frac{A^{k-1}\Delta^k}{k!} = \sum_{j\ge0}\frac{A^j\Delta^{j+1}}{(j+1)!}

结论(ZOH 离散化公式):

 Aˉ=eΔA,Bˉ=A1(eΔAI)B \boxed{\ \bar{A} = e^{\Delta A}, \qquad \bar{B} = A^{-1}\big(e^{\Delta A} - I\big)\,B\ }

2.6 小步长近似

当 Δ 很小时,eΔA=I+ΔA+O(Δ2)e^{\Delta A} = I + \Delta A + O(\Delta^2),于是

Bˉ=A1(ΔA+O(Δ2))B=ΔB+O(Δ2)\bar B = A^{-1}\big(\Delta A + O(\Delta^2)\big)B = \Delta B + O(\Delta^2)

BˉΔB\bar B \approx \Delta B。同理 AˉI+ΔA\bar A \approx I + \Delta A(这等价于欧拉法)。

注意:近似只在 Δ -> 0 时可靠。S4/Mamba 的实际实现都用精确的 ZOH 公式;且 Mamba 中 Δ 是逐 token 变化的,每步现算 eΔtAe^{\Delta_t A}

2.7 数值验证:完整算一遍

取标量系统 A=1A = -1,B=2B = 2,C=1C = 1,Δ=0.5\Delta = 0.5

第 1 步:离散化。

Aˉ=e0.50.6065,Bˉ=e0.511×2=(10.6065)×20.7869\bar A = e^{-0.5} \approx 0.6065, \qquad \bar B = \frac{e^{-0.5}-1}{-1}\times 2 = (1 - 0.6065)\times 2 \approx 0.7869

第 2 步:递归。 h0=0h_0 = 0,输入 x=[1,1]x = [1, 1]:

k xkx_k hk=0.6065hk1+0.7869xkh_k = 0.6065\,h_{k-1} + 0.7869\,x_k
1 1 0.7869
2 1 0.6065×0.7869 + 0.7869 ≈ 1.2642

第 3 步:与连续精确解对拍。x1x \equiv 1 的常数输入,连续方程 h=h+2h' = -h + 2 的解为 h(t)=2+(h02)eth(t) = 2 + (h_0 - 2)e^{-t}。在 t=1t = 1(即两步后):

h(1)=22e120.7358=1.2642h(1) = 2 - 2e^{-1} \approx 2 - 0.7358 = 1.2642 \quad✓

离散递归与连续精确解逐步完全一致–这正是 ZOH 离散化的性质:在 ZOH 假设成立(输入确为分段常数)时,它不是近似,而是精确等价

2.8 卷积形式:递归展开 = 一维卷积

先把「卷积」本身说清楚(离散因果时序卷积)。输入时序序列 u1,u2,,utu_1, u_2, \dots, u_t,卷积核(权重模板)K0,K1,K2,K_0, K_1, K_2, \dots,其中 KkK_k与当前时刻相隔 k 步的历史信息的权重因果(causal):只能看过去,看不到未来。形式化:

yt=τ=1tKtτuτy_t = \sum_{\tau=1}^{t} K_{t-\tau} \cdot u_\tau

最简单的数字例子。输入 u=[10,20,30]u = [10, 20, 30],衰减卷积核 K=[1, 0.5, 0.25]K = [1,\ 0.5,\ 0.25](越远权重越小):

y1=K0u1=1×10=10y2=K1u1+K0u2=0.5×10+1×20=25y3=K2u1+K1u2+K0u3=0.25×10+0.5×20+1×30=42.5\begin{aligned} y_1 &= K_0 u_1 = 1\times10 = 10 \\ y_2 &= K_1 u_1 + K_0 u_2 = 0.5\times10 + 1\times20 = 25 \\ y_3 &= K_2 u_1 + K_1 u_2 + K_0 u_3 = 0.25\times10 + 0.5\times20 + 1\times30 = 42.5 \end{aligned}

每一步都把全部历史按距离加权求和,权重模板固定、随距离翻牌–这就是卷积。y1y_1 只用了 u1u_1(因果),三个输出用的是同一套 K(时不变)。带着这两个属性看 SSM 的卷积形式,就是把 KkK_k 具体化为 CAˉkBˉC\bar A^k \bar B

命题:离散 SSM ht=Aˉht1+Bˉxth_t = \bar A h_{t-1} + \bar B x_t,yt=Chty_t = C h_t(设 h0=0h_0 = 0)等价于

y=Kˉx,Kˉ=(CBˉ, CAˉBˉ, CAˉ2Bˉ, , CAˉL1Bˉ)y = \bar K * x, \qquad \bar K = \big(C\bar B,\ C\bar A\bar B,\ C\bar A^2\bar B,\ \dots,\ C\bar A^{L-1}\bar B\big)

证明(直接展开):

ht=i=1tAˉtiBˉxih_t = \sum_{i=1}^{t} \bar A^{\,t-i}\bar B\, x_i

(归纳:h1=Bˉx1h_1 = \bar B x_1;设 ht1=it1Aˉt1iBˉxih_{t-1} = \sum_{i\le t-1}\bar A^{t-1-i}\bar B x_i,则 ht=Aˉht1+Bˉxt=it1AˉtiBˉxi+Bˉxth_t = \bar A h_{t-1} + \bar B x_t = \sum_{i\le t-1}\bar A^{t-i}\bar B x_i + \bar B x_t ✓)

代入输出方程:

yt=i=1tCAˉtiBˉxi=j=0t1Kˉjxtjy_t = \sum_{i=1}^{t} C\bar A^{\,t-i}\bar B\, x_i = \sum_{j=0}^{t-1} \bar K_j\, x_{t-j} \quad\blacksquare

数值验证(§2.7 的系统,输入 x=[2,4,0,8]x=[2,4,0,8]):

卷积核:Kˉj=0.6065j×0.7869\bar K_j = 0.6065^j \times 0.7869,即 [0.7869, 0.4772, 0.2894, 0.1755][0.7869,\ 0.4772,\ 0.2894,\ 0.1755]

y4=0.7869×8+0.4772×0+0.2894×4+0.1755×2=6.295+1.158+0.3517.804y_4 = 0.7869{\times}8 + 0.4772{\times}0 + 0.2894{\times}4 + 0.1755{\times}2 = 6.295 + 1.158 + 0.351 \approx 7.804

递归验证:h1=1.574, h2=4.102, h3=2.488, h4=0.6065×2.488+0.7869×81.509+6.295=7.804h_1=1.574,\ h_2=4.102,\ h_3=2.488,\ h_4=0.6065{\times}2.488+0.7869{\times}8 \approx 1.509+6.295 = 7.804

推论:LTI 系统拥有双形式–训练用卷积(FFT,并行),推理用递归(O(1)/token)。Mamba 使 B,C,ΔB, C, \Delta 依赖输入后,Kˉ\bar K 不再固定,卷积形式失效,改由并行扫描承担训练侧并行。

2.9 从 Δ 到遗忘门:通往 Mamba 的最后一环

回看 Aˉ=eΔA\bar A = e^{\Delta A}:若 A 为负定(如 Mamba-2 取 A=aIA = -a\cdot I,a>0a>0),则

Aˉt=eΔta(0,1)\bar A_t = e^{-\Delta_t a} \in (0, 1)

就是一个衰减门:

  • Δt\Delta_t \to \infty:Aˉt0\bar A_t \to 0 – 清空旧记忆,只看当前输入;
  • Δt0\Delta_t \to 0:Aˉt1\bar A_t \to 1 – 冻结状态,忽略当前输入。

Mamba 让 Δt=softplus(Linear(xt))\Delta_t = \mathrm{softplus}(\mathrm{Linear}(x_t)),离散化公式每步现算,于是「步长」变成了看内容的标量遗忘门

2.10 附录:双线性变换离散化对比

除 ZOH 外,S4 论文实际使用的是双线性变换(Tustin 方法):用梯形法则近似积分,

Aˉ=(IΔ2A)1(I+Δ2A),Bˉ=(IΔ2A)1ΔB\bar A = \Big(I - \frac{\Delta}{2}A\Big)^{-1}\Big(I + \frac{\Delta}{2}A\Big), \qquad \bar B = \Big(I - \frac{\Delta}{2}A\Big)^{-1}\Delta B

方法 性质 适用
ZOH ZOH 假设下精确;保稳定性;公式含 eΔAe^{\Delta A} Mamba 系列:Δt\Delta_t 逐 token 变化,每步现算矩阵指数(S4D/Mamba 的对角 A 使 eΔtAe^{\Delta_t A} 就是逐元素指数)
双线性 保稳定性;避免矩阵指数;把左半平面解析映射到单位圆内 S4 原论文:A 为 HiPPO 矩阵(非对角),双线性把它变成有理函数,配合 Cauchy 核计算

一句话收束:离散化方法的选择是被 A 的结构决定的–对角 A 配 ZOH(指数按元素算),非对角 HiPPO A 配双线性(变成多项式比值,可借 Cauchy 核求解)。


3. Mamba 一族:从 S4 到 Mamba-2

上面这条 SSM 线,从深度学习视角走过了三个关键站点:

S4(2021):把连续 SSM 正经地搬进序列建模。Aˉ=eΔA\bar A = e^{\Delta A}Bˉ=A1(eΔAI)B\bar B = A^{-1}(e^{\Delta A} - I)B 都是固定矩阵(与输入无关),于是 LTI 成立、递归 = 卷积,训练用 FFT 卷积、推理用 O(1)O(1) 递归(双形式,§2.8)。A 用 HiPPO 矩阵初始化(记住长历史的专门构造),离散化用双线性变换。限制也很明显:参数不随输入变,「记忆怎么衰减」在推理前就定死了。

Mamba / S6(2023):选择性改造。让 B,C,ΔB, C, \Delta 都依赖输入(Δt=softplus(Linear(xt))\Delta_t = \mathrm{softplus}(\mathrm{Linear}(x_t))),每步现算 eΔtAe^{\Delta_t A}–「步长」变成了看内容的遗忘门(§2.9)。代价是 LTI 被打破、固定卷积核 Kˉ\bar K 失效,训练侧改用并行扫描(associative scan);A 限制为对角结构使 eΔtAe^{\Delta_t A} 逐元素可算。这一步是「SSM 学会看内容」的关键。

Mamba-2 / SSD(2024):与线性注意力形式统一。状态改写成外积形式 St=αtSt1+vtktS_t = \alpha_t S_{t-1} + v_tk_t^\top,矩阵 AA 退化为标量衰减 αt=eΔta\alpha_t = e^{-\Delta_t a}从这一步起,SSM 与线性注意力在数学上就是同一个东西的两个记法(SSD 框架)。训练用 chunkwise:chunk 内 QKQK^\top 小注意力 + chunk 间递归,比纯 scan 的 GEMM 利用率高得多。

对照第 1 节的线性注意力递推 St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top:Mamba-2 只多了一个乘在旧状态上的标量 αt\alpha_t两条路线在 Mamba-2 处汇合:线性注意力给出了外积状态的样子,SSM 给出了衰减门。


4. 全链路一图流

h=Ah+Bx连续 SSM积分因子法h(t+Δ)=eAΔh(t)+0ΔeA(Δτ)Bxdτ通解:旧记忆衰减 + 输入累积ZOHht=Aˉht1+Bˉxt离散递归(推理用)展开y=Kˉx卷积(训练用)\underbrace{h' = Ah + Bx}_{\text{连续 SSM}} \xrightarrow{\text{积分因子法}} \underbrace{h(t+\Delta) = e^{A\Delta}h(t) + \int_0^\Delta e^{A(\Delta-\tau)}Bx\,d\tau}_{\text{通解:旧记忆衰减 + 输入累积}} \xrightarrow{\text{ZOH}} \underbrace{h_t = \bar A h_{t-1} + \bar B x_t}_{\text{离散递归(推理用)}} \xrightarrow{\text{展开}} \underbrace{y = \bar K * x}_{\text{卷积(训练用)}}

线性注意力侧:softmax(QK)V\mathrm{softmax}(QK^\top)V --核化–> ϕ(Q)(ϕ(K)V)\phi(Q)(\phi(K)^\top V) --逐 token–> St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top

两线在 Mamba-2 汇合:St=αtSt1+vtktS_t = \alpha_t S_{t-1} + v_tk_t^\top

此后:GDN 在写入侧加 delta rule(可删除的写入);KDA 把标量门打开成逐通道门 Diag(αt)\mathrm{Diag}(\alpha_t) 并给衰减加下界。骨架始终是同一条:状态 ×(衰减/删除算子)+(写入项)。后续演化见《KDA 的来龙去脉》

总结

  • 线性注意力:去掉 softmax 的耦合非线性 + 核化把非线性提前,Q(KV)Q(K^\top V) 结合律重新可用,O(n2)O(n)O(n^2) \to O(n);代价是纯加性状态,只会叠加不会删除(记忆碰撞);
  • SSM 一条链推完:齐次解 eAth0e^{At}h_0 -> 积分因子法得通解 -> ZOH 假设下提出 xkx_k 得离散递归,Aˉ=eΔA\bar A = e^{\Delta A}Bˉ=A1(eΔAI)B\bar B = A^{-1}(e^{\Delta A} - I)B;
  • ZOH 不是近似:输入分段常数时离散递归与连续解逐步相等(§2.7 数值对拍 1.2642 = 1.2642);
  • 递归 = 卷积:LTI 的双形式是 S4 训练(FFT 卷积)与推理(O(1) 递归)各取所长的根基;
  • Aˉ=eΔA\bar A = e^{\Delta A} 就是遗忘门:这是从控制论到 Mamba/KDA 的思想连续性;
  • 两条路线在 Mamba-2 汇合:外积状态(线性注意力)+ 标量衰减门(SSM)= St=αtSt1+vtktS_t = \alpha_t S_{t-1} + v_tk_t^\top,此后 delta rule 和逐通道门的演化都在这条骨架上进行。

参考:

KDA 的来龙去脉:从线性注意力到 Kimi Delta Attention

KDA(Kimi Delta Attention)是 Kimi K3 里负责长序列的那 69 层(总共 93 层)。本文不铺开 K3 的全部架构,只讲一条演化线:KDA 之前的几个模型(线性注意力、Mamba、DeltaNet/GDN)分别解决了什么问题,KDA 又在它们的基础上改了什么

故事线很直白:线性注意力先把「固定大小的记忆」造出来,但这个记忆只会往里堆、不会清理;Mamba 教会它遗忘,但遗忘只能整片擦;DeltaNet 让它能精准改写某条记忆;GDN 把「擦除」和「改写」两样拼在一起;KDA 最后一步,让每个维度各自决定忘得多快。每一步都有公式推导、手动验算和工程上的理由。

0. 先算一笔账:为什么需要「固定大小的 KV 记忆」

标准 softmax 注意力在生成文本时,必须把每个历史 token 的 Key 和 Value 都存下来(KV cache)–因为每生成一个新 token,都要和前面所有 token 重新算一遍注意力。序列越长,这笔账越贵:

开销 softmax 注意力(KV cache) 固定状态(线性注意力一族)
解码显存 O(nd)O(n \cdot d)(全部历史的 KV 都得缓存) O(d2)O(d^2)(一个固定矩阵 S,与 n 无关)
单步解码 O(nd)O(n \cdot d)(扫全部历史) O(d2)O(d^2)(更新 S + 读出,两次矩阵乘)
预填充 O(n2d)O(n^2 d) O(nd2)O(n d^2)
上下文长度翻倍 显存、延迟一起翻倍 完全不变

(单头视角,d 为特征维度。)

这笔账具体有多痛:几十万 token 的上下文就能把 KV cache 撑到几十 GB。长上下文场景(长文档、agent 记忆、仓库级代码理解)里,第一瓶颈往往不是模型权重,而是 KV cache 的显存和带宽。K3 把「序列维度缩放」列为头号工程目标,93 层里 69 层换成 KDA,买的就是右列这张表。

但固定大小不是白来的。一个 dv×dkd_v \times d_k 的矩阵要装下任意长的序列,记忆就必须会写、会改、会忘–只进不出的仓库,序列一长就糊成一团。怎么设计这块「会自我管理的固定记忆」,就是本文主线,按出场顺序:

  • 线性注意力:先把固定状态本身造出来(用外积累积出一个矩阵 S),但只会做加法,只进不出;
  • Mamba/SSM:加上一个全局衰减系数 α–会忘了,但只能整片擦;
  • DeltaNet:加上定点改写(delta rule)–会改了,写入的是新旧值之差;
  • GDN:把「板擦」(α)和「铅笔」(delta)拼在一起,先删后写;
  • KDA:把一个总的衰减系数拆成每个维度一个,各维度自己决定记多久,再加固数值下界,进 K3。

线性注意力和 SSM 这两条技术路线的完整推导–softmax 为什么挡路、核化怎么搬、积分因子法、ZOH 离散化、卷积形式、S4 到 Mamba-2 的演化–单独成了一篇:《线性注意力与 SSM:两条技术路线的完整推导》。本文只要求读者记得结论,主线是:每个模型解决前一个的什么问题、留下什么新问题,讲完之后再回头看总表。


1. KDA 是什么:先看递推式

KDA 是 K3 序列维度缩放的核心,基于 delta rule 循环,并加上了通道维度的遗忘门

路线一:线性注意力–固定状态是怎么来的(概述)

线性注意力的完整推导(softmax 为什么挡路、核函数怎么把非线性搬走、结合律怎么把 O(n2)O(n^2) 降到 O(n)O(n))在前置篇《线性注意力与 SSM:两条技术路线的完整推导》里一步步展开,这里只留骨架:

  • 去掉 softmax,(QK)VQ(KV)(QK^\top)V \to Q(K^\top V):换一种括号方式,先算 KVK^\top V,复杂度 O(dn2)O(nd2)O(dn^2) \to O(nd^2);
  • 核化 sim(q,k)=ϕ(q)ϕ(k)\mathrm{sim}(q,k) = \phi(q)^\top\phi(k):非线性提前到 Q/K 各自身上,n×nn\times n 的注意力矩阵就消失了,换成一个固定大小的状态 SRd×dS \in \mathbb{R}^{d\times d};
  • 状态可以逐步累积:St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top,每来一个 token 更新一次、读一次(fast weights 视角:S 是被逐 token 编程的快速权重,查询是读出)。

代价也在这三行里:S 只会加法。写入只有「往上堆」,没有「删掉」–序列长度远超状态容量时,不同 kvk \to v 的关联互相干扰,旧信息永远赖在状态里,新信息覆盖不掉。前置篇里有个数值演示:同一个 key 写两次,线性注意力读出来的是两个 value 的叠加,而 delta rule 读出来的是新的那个。这个「记忆碰撞」就是线性注意力最大的问题,修复它是 DeltaNet 的全部动机。

DeltaNet:从加法到差值写入

线性注意力留下的 S 是个只会「贴便利贴」的状态,修正它就是 DeltaNet 的全部动机。解法:写入差值,不写全值。写入前,先用当前 key 把状态里已存的内容读一遍:

vold=St1kt(当前 key 指向的旧内容),ut=βt(vtvold)v_{\mathrm{old}} = S_{t-1}^\top k_t \quad (\text{当前 key 指向的旧内容}), \qquad u_t = \beta_t (v_t - v_{\mathrm{old}})

St=St1+ktut=(Iβtktkt)St1+βtktvtS_t = S_{t-1} + k_t u_t^\top = (I - \beta_t k_t k_t^\top)\, S_{t-1} + \beta_t k_t v_t^\top

符号 形状 含义
voldv_{\mathrm{old}} [dv][d_v] 当前 key 从状态里读出的已有内容
βt\beta_t 标量 (0,1)\in (0,1) 内容替换强度
utu_t [dv][d_v] 实际写入的差值

展开后两项分解:删除项 (Iβtktkt)(I - \beta_t k_t k_t^\top) 沿 ktk_t 方向抹掉旧内容,再加回新值项 βtktvt\beta_t k_t v_t^\top–先删旧、再写新,一次外积同时完成两件事。预测误差 vtvoldv_t - v_{\mathrm{old}} 就是 delta,Delta Rule 由此得名:每次写入的从来不是新值本身,而是新值与旧值之差。

为什么恰好是差值?因为它就是对回归损失做一步 SGD 的结果。注意 DeltaNet 没有动线性注意力的两个前置条件–ϕ\phi、外积状态、递推形式全部保留(实现里常取 $\phi = $ L2Norm,即下文单位球约束),它改的只是写入规则。而且这个规则不是拍脑袋的设计,它可以从回归损失严格推出来:

L(S)=12Sktvt2,SL=(Sktvt)kt\mathcal{L}(S) = \tfrac{1}{2}\|S k_t - v_t\|^2, \qquad \nabla_S\mathcal{L} = (Sk_t - v_t)k_t^\top

St=St1βtSL=St1(Iβtktkt)先删:擦掉 St1kt 方向的旧值+βtvtkt后写S_t = S_{t-1} - \beta_t\nabla_S\mathcal{L} = \underbrace{S_{t-1}(I - \beta_tk_tk_t^\top)}_{\text{先删:擦掉 } S_{t-1}k_t \text{ 方向的旧值}} + \underbrace{\beta_tv_tk_t^\top}_{\text{后写}}

βt\beta_t 就是学习率。论文原版写法还提供了一个插值视角:St=St1vtoldkt+vtnewktS_t = S_{t-1} - v_t^{\mathrm{old}}k_t^\top + v_t^{\mathrm{new}}k_t^\top,其中 vtnew=βtvt+(1βt)St1ktv_t^{\mathrm{new}} = \beta_tv_t + (1-\beta_t)S_{t-1}k_t–δ-rule 是软覆写,β 控制新旧比例(β=1 且 k 为单位向量时完全覆写)。q/k 做 L2 归一化(钉在单位球上)保证 kkkk^\top 特征值 1\le 1,更新稳定。

直观类比:state 是白板,key 是指针。纯加性写入是往白板上不停贴便利贴,贴满了就糊;DeltaNet 是先擦掉指针 ktk_t 指向的那块,再贴新的。βt=0\beta_t = 0:完全不写入,状态不动;βt=1\beta_t = 1:完全替换,指针指向的内容整个换掉。这种「先擦后写」的机制让状态能覆盖错误记忆,这是 GDN/DeltaNet 一线的核心改进。

关键整理:写成转移矩阵形式。定义 Ht=IβtktktH_t = I - \beta_t k_t k_t^\top,状态更新变成

St=HtSt1+βtktvtS_t = H_t\, S_{t-1} + \beta_t k_t v_t^\top

这个形式之所以非常关键,是因为它暴露了 DeltaNet 和 SSM 的同构:HtH_t 正是随输入变化的状态转移矩阵(SSM 里是 Aˉ=eΔA\bar A = e^{\Delta A}),βtktvt\beta_t k_t v_t^\top 正是写入项(SSM 里是 Bˉxt\bar B x_t)。

把递归展开,消掉时间依赖。记 Bt=βtktvtB_t = \beta_t k_t v_t^\top,递推 St=St1Ht+BtS_t = S_{t-1}H_t + B_t 逐层代入:

S1=S0H1+B1S_1 = S_0 H_1 + B_1

S2=S0H1H2+B1H2+B2S_2 = S_0 H_1 H_2 + B_1 H_2 + B_2

S3=S0H1H2H3+B1H2H3+B2H3+B3S_3 = S_0 H_1 H_2 H_3 + B_1 H_2 H_3 + B_2 H_3 + B_3

S4=S0H1H2H3H4+B1H2H3H4+B2H3H4+B3H4+B4S_4 = S_0 H_1 H_2 H_3 H_4 + B_1 H_2 H_3 H_4 + B_2 H_3 H_4 + B_3 H_4 + B_4

规律一目了然:St=S0itHi+itBii<jtHjS_t = S_0 \prod_{i \le t} H_i + \sum_{i \le t} B_i \prod_{i < j \le t} H_j。展开后递归被彻底消掉了:每个 StS_t 都是初始状态、各步写入项与转移矩阵连乘的线性组合,只剩矩阵乘法和求和。而矩阵乘法满足结合律,「从左往右扫」只是众多括号化方案之一–换个括号方式(比如二叉树式两两合并),HH 的连乘与写入项的累积可以在 O(logn)O(\log n) 深度内并行完成。这就是 parallel scan / associative scan(并行扫描/结合扫描) 类方法的核心思想,也是 chunkwise 并行化和 SSM 训练并行(如 Mamba 的 selective scan)共同的理论根基。

但问题还没解决:HtH_t 是什么? 回看定义 Ht=IβtktktH_t = I - \beta_t k_t k_t^\top。设 ktRdk_t \in \mathbb{R}^d,则 ktktk_t k_t^\topd×dd \times d秩 1(rank-1)矩阵–所有列都是 ktk_t 的标量倍,只张出一维。所以 HtH_t 本质是单位矩阵减去一个秩 1 矩阵:ktk_t 方向被缩放 1βtkt21 - \beta_t\|k_t\|^2,正交补方向原封不动。这类「I 减秩 1」的结构与**豪斯霍尔德变换(Householder transformation)**同族–QR 分解里用它做反射/消元。区别在于:Householder 取 β=2/k2\beta = 2/\|k\|^2 且要求 k=1\|k\|=1 时是精确反射(范数保长),DeltaNet 的 βt(0,1)\beta_t \in (0,1) 是「部分反射」,软化为可学习的写入强度。秩 1 结构也解释了删除项为什么便宜:抹掉一个方向不需要满秩运算,一次外积就够。

到这里两条路线可以拼在一起了:DeltaNet 的 St=HtSt1+βtktvtS_t = H_t S_{t-1} + \beta_t k_t v_t^\top 与 SSM 的 ht=Aˉht1+Bˉxth_t = \bar A h_{t-1} + \bar B x_t 结构完全同构,差的只有一件事–HtH_t 的遗忘是「沿 ktk_t 方向删一块」,没有 SSM 那种全通道的指数衰减。

拼在一起:Gated DeltaNet = DeltaNet 的精确写入 + Mamba-2 的全局遗忘。论文的核心洞察:gating 和 delta rule 是互补的两种记忆管理机制:

机制 比喻 能力 缺陷
Gating(Mamba-2) 板擦 一键大面积擦除(全局衰减) 无法定点修改
Delta rule(DeltaNet) 铅笔 定点覆写某个 key 的关联 无法快速清空

序列长度超过状态容量时,记忆碰撞必然发生;GDN 用板擦 + 铅笔的组合管理这块固定大小的白板:

St=St1(αt(Iβtktkt))+βtvtktS_t = S_{t-1}\big(\alpha_t(I - \beta_tk_tk_t^\top)\big) + \beta_tv_tk_t^\top

其中 αt(0,1)\alpha_t \in (0,1) 是数据相关的标量门(Mamba-2 的参数化:α=exp(\alpha = \exp(负 Softplus(Linear(xt)))(x_t))),在 log 空间计算以保证数值稳定)。三个极限情形读懂这个式子:

极限 行为 对应模型
αt1\alpha_t \to 1 纯 delta rule,只定点改写 DeltaNet
βt1\beta_t \to 1,k ⊥ 已有记忆 退化为 St=αtSt1+vtktS_t = \alpha_tS_{t-1} + v_tk_t^\top Mamba-2
αt0\alpha_t \to 0 整表清零再写入(硬重置) 新能力:两者都做不到

几何直觉:(Iβkk)(I - \beta kk^\top) 是广义 Householder 反射–把状态往 k 方向「压扁擦除」;再乘标量 α\alpha 把整张表均匀缩小。一个是定向操作,一个是全局操作,作用在不同维度上,所以可以叠加。

注意作用顺序:擦除量是 βt(St1kt)\beta_t(S_{t-1}k_t),用的是未衰减的旧读出;而 delta 对照的是衰减后的旧值(vtαtSt1ktv_t - \alpha_tS_{t-1}k_t)。α\alpha 乘的是整个 (Iβkk)(I-\beta kk^\top)–这一点写代码时极易弄错(把 α 只乘到擦除项上,chunkwise 形式会与递归形式对不上)。读取侧同样随时间累积衰减:token 在时间步 xx 写入,在 x+tx+t 读取时已经被 αxαx+1αx+t\alpha_x\alpha_{x+1}\dots\alpha_{x+t} 衰减过。实现中通过 γr/γi\gamma^r/\gamma^i 项修正–分子分母都是 α\alpha 连乘,相除就是区间衰减,本质是乘法形式的前缀和(prefix-sum),与 SSM 的 Aˉ\bar A 连乘完全同源。

回头看:一张图 + 一张表把全家统一(论文 Table 1)。到这里四个模型都出场了,回头看,它们其实是同一个在线优化问题的闭式解,差别只在目标函数:

Linear Attn(纯加性)Mamba-2(全局遗忘门)DeltaNet(定点覆写)    GDN(板擦+铅笔)KDA(逐通道门+下界)\text{Linear Attn}(纯加性) \to \text{Mamba-2}(全局遗忘门) \searrow \\ \text{DeltaNet}(定点覆写) \nearrow \;\; \text{GDN}(板擦+铅笔) \to \text{KDA}(逐通道门+下界)

模型 在线学习目标 闭式解(状态更新)
Linear Attn StSt1F22Stkt,vt|S_t - S_{t-1}|_F^2 - 2\langle S_tk_t, v_t\rangle St=St1+vtktS_t = S_{t-1} + v_tk_t^\top
Mamba-2 StαtSt1F22Stkt,vt|S_t - \alpha_tS_{t-1}|_F^2 - 2\langle S_tk_t, v_t\rangle St=αtSt1+vtktS_t = \alpha_tS_{t-1} + v_tk_t^\top
DeltaNet StSt1F22Stkt,βt(vtSt1kt)|S_t - S_{t-1}|_F^2 - 2\langle S_tk_t, \beta_t(v_t - S_{t-1}k_t)\rangle St=St1(Iβtktkt)+βtvtktS_t = S_{t-1}(I - \beta_tk_tk_t^\top) + \beta_tv_tk_t^\top
GDN StαtSt1F22Stkt,βt(vtαtSt1kt)|S_t - \alpha_tS_{t-1}|_F^2 - 2\langle S_tk_t, \beta_t(v_t - \alpha_tS_{t-1}k_t)\rangle St=St1(αt(Iβtktkt))+βtvtktS_t = S_{t-1}(\alpha_t(I-\beta_tk_tk_t^\top)) + \beta_tv_tk_t^\top
KDA (GDN 的逐通道化) (Iβtktkt)Diag(αt)St1+βtktvt(I-\beta_tk_tk_t^\top)\,\mathrm{Diag}(\alpha_t)S_{t-1} + \beta_tk_tv_t^\top

读法:第一项是正则(别离上一刻状态太远 = 记忆保留),第二项是拟合(SkvSk \approx v = 关联学习)。Mamba-2 与 DeltaNet 是两条独立的进化路线–前者把正则锚点换成可收缩的 αtSt1\alpha_tS_{t-1}(选择性遗忘),后者把拟合项换成回归目标(修正误差);GDN 把两条路线拼在了一起,KDA 再把拼完之后的标量门拆成逐通道门。沿着两个轴看这张表的演化:

  • 拟合轴:线性注意力/Mamba-2 的拟合项 2Sk,v-2\langle Sk, v\rangle 只要求内积大,太弱,所以只能「越叠越厚」;DeltaNet/GDN 换成回归目标 Skv2\|Sk - v\|^2 类的构造,才获得「修正误差」的闭式解–delta 写入是回归目标的直接推论;
  • 正则轴:Mamba-2/GDN 的正则锚点是 αtSt1\alpha_tS_{t-1}(可收缩的锚):正则放松 \Rightarrow 允许状态偏离旧值 \Rightarrow 选择性遗忘

所以这棵「进化树」可以反过来读:不是「谁加了什么模块」,而是目标函数从弱拟合 + 硬正则,逐级加强到强拟合 + 软正则,每一级都是上一级的严格泛化(GDN 令 βt\beta_t 退化为 0/边界值就退回全部下游模型)。再往前推一步就是**测试时训练(Test-Time Training, TTT)**一族:delta rule 是对 L(S)=12Sktvt2\mathcal{L}(S)=\frac12\|Sk_t-v_t\|^2 逐步做 SGD(β=学习率),GDN 等价于 SGD + 自适应权重衰减(α)–把「序列处理」看成「在线优化」,线性 RNN 全家都是这个视角的特例。

KDA 要做的最后一步就清楚了。Kimi Linear(后演化为 KDA)在 Gated DeltaNet 基础上的核心改进是细粒度门控:不再是每个注意力头一个标量 α\alpha,而是每个通道一个独立衰减值:

αt(0,1)dk(channel-wise 遗忘)\alpha_t \in (0,1)^{d_k} \quad (\text{channel-wise 遗忘})

作用是模型可以对不同维度做不同程度的记忆衰减:部分通道 α\alpha 接近 1,保留长期信息;部分通道 α\alpha 接近 0,快速遗忘。打个比方,Gated DeltaNet 的标量门是「全屋一个总闸」,KDA 的逐通道门是「每盏灯一个旋钮」–同一个状态里,慢通道记长程依赖,快通道只管局部上下文,记忆容量按维度重新分配。再加上衰减下界(log\log-decay 限制在 (gmin,0)(g_{\min}, 0))保证数值稳定,就得到 KDA 的完整递推式–正是下一小节的内容。

下面先把 GDN(Gated DeltaNet)的机制吃透:手算一遍、推导 chunkwise 并行形式,最后对照 KDA 看它继承了什么、改了什么。

GDN 数值验算:手算一遍

设定 dk=dv=2d_k = d_v = 2,C=3C = 3 个 token,S0=0S_0 = 0。精心设计的数据–k1=k2=e1k_1 = k_2 = e_1 制造 key 碰撞:

K=[101001], V=[102031], α=[0.8,0.5,0.9], β=[1,1,0.6]K = \begin{bmatrix}1&0\\1&0\\0&1\end{bmatrix},\ V = \begin{bmatrix}1&0\\2&0\\3&1\end{bmatrix},\ \alpha = [0.8,\,0.5,\,0.9],\ \beta = [1,\,1,\,0.6]

顺序递归(ground truth)。累积衰减 γj=ijαi=[0.8,0.4,0.36]\gamma_j = \prod_{i\le j}\alpha_i = [0.8, 0.4, 0.36]

t=1(k=e1k=e_1, v=[1,0]v=[1,0], α=0.8\alpha=0.8, β=1\beta=1):S0k=0S_0k = 0,无旧记忆可删:

S1=0.80+1([1,0]0)e1=[1000],o1=S1q1=[1,0]S_1 = 0.8\cdot 0 + 1\cdot([1,0]-0)\,e_1^\top = \begin{bmatrix}1&0\\0&0\end{bmatrix}, \qquad o_1 = S_1q_1 = [1,0]

t=2(k=e1k=e_1, v=[2,0]v=[2,0], α=0.5\alpha=0.5, β=1\beta=1):旧读出 S1k2=[1,0]=v1S_1k_2 = [1,0] = v_1碰撞发生。擦除 1[1,0]e1-1\cdot[1,0]e_1^\topv1v_1 完全擦掉;写入 delta =[2,0]0.5[1,0]=[1.5,0]= [2,0] - 0.5\cdot[1,0] = [1.5, 0]:

S2=0.5S1[1,0]e1+[1.5,0]e1=[2000]S_2 = 0.5S_1 - [1,0]e_1^\top + [1.5,0]e_1^\top = \begin{bmatrix}2&0\\0&0\end{bmatrix}

现在查 e1e_1[2,0] = v2v_2,v1v_1 已被覆写(对照线性注意力:纯加性会得到 [3,0] 的污染叠加)。

t=3(k=e2k=e_2, v=[3,1]v=[3,1], α=0.9\alpha=0.9, β=0.6\beta=0.6):旧读出 S2k3=0S_2k_3 = 0,正交无碰撞。写入 0.6[3,1]e20.6\cdot[3,1]e_2^\top,同时全表再乘 α\alpha(此处也乘了已写入的 e1e_1 行):

S3=[1.81.800.6],o3=S3q3=[1.8,0.6]S_3 = \begin{bmatrix}1.8&1.8\\0&0.6\end{bmatrix}, \qquad o_3 = S_3q_3 = [1.8,\,0.6]

检查点:为什么 S3[0,0]=1.8S_3[0,0] = 1.8 不是 2.0? t=3 的 α3=0.9\alpha_3=0.9 作用在整张表上:第 1 行 2.0×0.9=1.82.0\times0.9 = 1.8。这就是 gating 与 delta 的交互:即使 token 3 的 key 与 e1e_1 正交,它的遗忘门仍然衰减了 e1e_1 通道上的记忆。逐通道门(KDA)与下界衰减都是围绕这一约束做文章。

GDN 的 Chunkwise 并行形式:递归怎么塞进 GPU

推理用递归(O(1)O(1)/token),但训练/prefill 必须并行。GDN 的贡献是把 gating 并入 DeltaNet 的 WY 表示 chunkwise 框架。设 chunk 大小 C,chunk 入口状态 S0S_0,目标:一次矩阵乘算出整个 chunk 的 O 和 chunk 出口状态 SCS_C

第一步:部分展开递归:

Sr=γrS0PrFr+i=1rγrγiu~ikiGrS_r = \underbrace{\gamma_r S_0 P_r}_{F_r} + \underbrace{\sum_{i=1}^r\frac{\gamma_r}{\gamma_i}\,\tilde u_i k_i^\top}_{G_r}

  • γr=jrαj\gamma_r = \prod_{j\le r}\alpha_j(α\alpha 是标量,可提到矩阵连乘外面);
  • Pr=ir(Iβikiki)P_r = \prod_{i\le r}(I - \beta_ik_ik_i^\top):纯 Householder 连乘,与 gating 无关;
  • u~i\tilde u_i:吸收了 β\beta 和衰减修正的「伪 value」。

第二步:WY 表示–秩 1 连乘压缩成两个小矩阵。关键观察:Pr=IirwikiP_r = I - \sum_{i\le r}w_ik_i^\top(乘积不膨胀,始终是 I 减一个秩 r\le r 的项)。wrw_ru~r\tilde u_r 的递归(W 不带 γ\gammaU~\tilde Uγ\gamma,这是 gating 并入的位置):

wr=βr(kri<rwi(kikr)),u~r=βr(vri<ru~iγrγi(kikr))w_r = \beta_r\Big(k_r - \sum_{i<r}w_i\,(k_i^\top k_r)\Big), \qquad \tilde u_r = \beta_r\Big(v_r - \sum_{i<r}\tilde u_i\,\tfrac{\gamma_r}{\gamma_i}(k_i^\top k_r)\Big)

直觉:wrw_r 是「擦除向量」–前面的写入若与 krk_r 重叠,先从自己的擦除量里扣掉(避免重复擦除);u~r\tilde u_r 是「修正量 value 版」。两者形式相同,差别只在 u~\tilde u 的内积多了一个衰减比 γr/γi\gamma_r/\gamma_i:i 越久远、衰减越多,它对当前擦除的干扰越小

第三步:UT 变换–递归变成一次下三角方程求解。w 和 u~\tilde u 的递归是「自回归」的(依赖前面的 w/u~\tilde u),但系数矩阵是单位下三角,可整体解方程。

辨析:因果掩码 ≠ UT 变换。两者容易混,因为碰巧都是下三角,但要管的完全是两件事:

  • 因果掩码(tril):把 score 矩阵的严格上三角直接置零,一行掩码操作,没有任何「变换」可言。它管的是「第 i 个输出不许看未来的 token」;
  • UT 变换(求 (I+L)1(I+L)^{-1}):管的是已经看过的历史之间怎么算账。delta rule 每次写入是「先读、再改」,块内第 2 次写入读到了第 1 次的结果,第 3 次读到前两次–历史写入之间互相纠缠,UT 把这团纠缠解开。

两者都是下三角不是巧合,是同一个原因:因果性。位置 i 的写入只能影响 i 之后的位置,所以「干扰系数矩阵」天然下三角;对角线天然是 1(自己的写入自己完整可见),于是要逆的矩阵恰好是单位下三角–这就是「UT」(Unit Triangular)名字的由来。一句话:因果性决定了它是三角的,但做它的目的是解耦,不是掩码

为什么必须做(最短版本):并行化要求把 C 次顺序写入折成一锤子外积累加 iv~iki\sum_i \tilde v_i k_i^\top。如果直接用原始 viv_i 累加,重叠部分会被重复计算–第 1 次写入的内容会透过后续写入的「先读」环节被间接再写一遍。UT 变换算出每个 viv_i 该扣除多少,使等式精确成立。不做的代价:要么结果错,要么退回逐 token 循环。

旁注:这类下三角变换是个大家族。「顺序递推 ↔ 三角矩阵求逆」是个通用模式,KDA 的 UT 只是其中一员:

家族成员 三角结构 与 UT 的关系
三角方程组求解(前代/回代) LU、Cholesky 分解之后的三角系统 UT 变换的计算过程就是一次前代法–同一个算法
Householder QR 的紧凑 WY 表示 反射连乘 (Iβkk)=I+YTY\prod(I - \beta kk^\top) = I + YTY^\top,TT 三角 「WY」「UT」两个名字的学术出处(Schreiber-Van Loan 1989);DeltaNet 把同样的打包思想借到 delta rule
因果卷积的逆(去卷积) 因果线性系统 = 下三角 Toeplitz,其逆也是下三角 Toeplitz 信号处理经典:「用三角逆矩阵解顺序依赖」
Mamba-2/SSD 的半可分矩阵 块内注意力 (CB)L(CB^\top)\odot L,LL 下三角衰减 同一枚硬币另一面:顺序 SSM 递推等价于带结构下三角矩阵,KDA 的 AqkA^{qk} 衰减注意力项完全是这个结构
幂零矩阵 Neumann 级数 (I+L)1=IL+L2(I+L)^{-1} = I - L + L^2 - \cdots 一般矩阵是无穷级数,但严格下三角矩阵幂零(LC=0L^C = 0),有限项精确截断–UT 能精确、便宜算出的数学原因

放进大图景里记:凡是「顺序执行的因果更新」,分块并行化时都会变成一个下三角矩阵的求逆/求解问题–SSM、delta rule、因果卷积、QR 分解,全是这个模式的化身。KDA chunkwise 里的 WY 表示(第二步)与 UT 变换(第三步),分别是这个模式在「写入打包」与「依赖解耦」上的具体化。

回到公式。定义衰减感知掩码 Γij=γi/γj (i>j)\Gamma_{ij} = \gamma_i/\gamma_j\ (i>j),则:

Tplain=[I+strictLower(diag(β)KK)]1diag(β),W=TplainKT_{\text{plain}} = \big[I + \mathrm{strictLower}(\mathrm{diag}(\beta)\,KK^\top)\big]^{-1}\mathrm{diag}(\beta), \qquad W = T_{\text{plain}}K

Tgated=[I+strictLower(diag(β)(ΓKK))]1diag(β),U~=TgatedVT_{\text{gated}} = \big[I + \mathrm{strictLower}(\mathrm{diag}(\beta)\,(\Gamma\odot KK^\top))\big]^{-1}\mathrm{diag}(\beta), \qquad \tilde U = T_{\text{gated}}V

W 和 U~\tilde U 只差内积处的一个 γ\gamma 比。下三角求逆(C×CC\times C,实践中 C=64)用前代法(forward substitution)完成。

第四步:chunk 输出与出口状态。记号(论文的箭头约定):()r=γr()r\overleftarrow{(\cdot)}_r = \gamma_r(\cdot)_r(衰减到 chunk 首端),()r=γCγr()r\overrightarrow{(\cdot)}_r = \frac{\gamma_C}{\gamma_r}(\cdot)_r(衰减到 chunk 末端)。

输出(两项:读 chunk 外旧状态 + chunk 内交互):

O=QS0+(QKΓcausal)(U~WS0)O = \overleftarrow{Q}\,S_0^\top + \big(QK^\top\odot\Gamma_{\text{causal}}\big)\big(\tilde U - \overleftarrow W S_0^\top\big)

出口状态:

SC=γCS0+(U~WS0)K(U~r=γCγru~r)S_C = \gamma_C S_0 + \big(\overrightarrow{\tilde U} - \overrightarrow W S_0^\top\big)^\top K \qquad(\overrightarrow{\tilde U}_r = \tfrac{\gamma_C}{\gamma_r}\tilde u_r)

结构读法:

  • 输出的第二项是「chunk 内小注意力」:QKΓcausalQK^\top\odot\Gamma_{\text{causal}} 就是带衰减的因果注意力矩阵,attend 的对象不是 V 而是修正后的伪 value U~WS0\tilde U - \overleftarrow WS_0^\top(后者是「旧状态在这个 chunk 里该被擦掉的部分」);
  • 出口状态 = 旧状态整体衰减 γC\gamma_C + 修正量加权写入(权重 γC/γi\gamma_C/\gamma_i:越早写入衰减越多);
  • 全部计算都是 C×CC\times CC×dkC\times d_kC×dvC\times d_v 的稠密矩阵乘–Tensor Core 友好;chunk 间只传一个 dv×dkd_v\times d_k 矩阵。

数值验算:拿顺序递归的例子走一遍 chunkwise

Step 1:KK=[110110001]KK^\top = \begin{bmatrix}1&1&0\\1&1&0\\0&0&1\end{bmatrix},ΓstrictKK\Gamma_{\text{strict}}\odot KK^\top 只有 (2,1) 处非零 =γ2/γ1×1=0.5= \gamma_2/\gamma_1\times1 = 0.5

Step 2:解两个下三角方程(β=[1,1,0.6]\beta = [1,1,0.6]):

Tplain=[100110000.6],Tgated=[1000.510000.6]T_{\text{plain}} = \begin{bmatrix}1&0&0\\-1&1&0\\0&0&0.6\end{bmatrix}, \qquad T_{\text{gated}} = \begin{bmatrix}1&0&0\\-0.5&1&0\\0&0&0.6\end{bmatrix}

对比唯一差别 (2,1):-1 -> -0.5,正是 γ2/γ1=0.5\gamma_2/\gamma_1 = 0.5 的衰减–t=2 时 v1v_1 已被衰减一半,擦除它的需求也减半。

Step 3:W=TplainK=[100000.6]W = T_{\text{plain}}K = \begin{bmatrix}1&0\\0&0\\0&0.6\end{bmatrix},U~=TgatedV=[101.501.80.6]\tilde U = T_{\text{gated}}V = \begin{bmatrix}1&0\\1.5&0\\1.8&0.6\end{bmatrix}

u~\tilde u 的第二行:$v_2 - 0.5\cdot\tilde u_1(k_1^\top k_2) = [2,0] - 0.5[1,0] = $ [1.5, 0]–与顺序递归里手算的 delta 完全一致 ✓(W2=0W_2 = 0 则因为 w1w_1 已把 e1e_1 方向擦干净,β=1\beta=1 时无需再擦)。

Step 4(S0=0S_0=0,修正项消失):U~=γCγiu~i\overrightarrow{\tilde U} = \frac{\gamma_C}{\gamma_i}\tilde u_i 逐行乘 [0.36/0.8,0.36/0.4,1]=[0.45,0.9,1][0.36/0.8, 0.36/0.4, 1] = [0.45, 0.9, 1]:

U~=[0.4501.3501.80.6],SC=U~K=[1.81.800.6] \overrightarrow{\tilde U} = \begin{bmatrix}0.45&0\\1.35&0\\1.8&0.6\end{bmatrix}, \qquad S_C = \overrightarrow{\tilde U}^\top K = \begin{bmatrix}1.8&1.8\\0&0.6\end{bmatrix}\ \checkmark

与顺序递归的 S3S_3 完全一致

Step 5:Γcausal=[1000.5100.450.91]\Gamma_{\text{causal}} = \begin{bmatrix}1&0&0\\0.5&1&0\\0.45&0.9&1\end{bmatrix},O=(QKΓcausal)U~O = (QK^\top\odot\Gamma_{\text{causal}})\tilde U:

O=[1000.510001][101.501.80.6]=[10201.80.6] O = \begin{bmatrix}1&0&0\\0.5&1&0\\0&0&1\end{bmatrix}\begin{bmatrix}1&0\\1.5&0\\1.8&0.6\end{bmatrix} = \begin{bmatrix}1&0\\2&0\\1.8&0.6\end{bmatrix}\ \checkmark

与顺序递归的 o1=[1,0]o_1=[1,0]o2=[2,0]o_2=[2,0]o3=[1.8,0.6]o_3=[1.8,0.6] 逐步精确一致(含非零 S0S_0 的一般情形同样对拍通过)。

完整 block:GDN 在网络里长什么样

Llama 式宏架构,把 attention 换成 gated delta token mixer:

1
2
3
4
5
6
7
x ─┬─ W_q/W_k ─ ShortConv ─ SiLU ─ L2Norm ─-> q, k
├─ W_v ──── ShortConv ─ SiLU ───────────-> v
├─ W_α(线性投影,log 空间负 Softplus)──-> α ∈ (0,1) 标量
├─ W_β(线性投影 + Sigmoid)────────────-> β ∈ (0,1)
│ 递归/分块计算 gated delta rule ─-> õ
├─ W_g ─ SiLU ─────────────────────────-> 输出门
└─ y = W_o( Sigmoid(W_g x) ⊙ RMSNorm(õ) ),再走 SwiGLU MLP + 残差

要点:q/k 走 ShortConv+SiLU+L2Norm(局部上下文 + 单位球稳定擦除几何),α/β 只走线性投影(它们是标量,无需卷积);输出门与 Mamba 的 SiLU 门一脉相承(K3 把它升级为全秩)。

旁注:ShortConv 的数学定义

上面流程图里的 ShortConv(Short Convolution,短因果深度卷积)出自 Kimi Linear(K3 技术报告引用 [64]),位于 Q/K/V 线性投影之后、进入 KDA 递归之前–顺序是 xWq/k/vShortConvx \to W_{q/k/v} \to \mathrm{ShortConv} \to \dots,即卷积是「KDA 前置」而非「投影前置」。它只捕捉当前 token 前面少量局部上下文,无未来 token 泄露;且在 Kimi Linear / FLA 实现里卷积不是裸用的,后面紧跟 SiLU:silu(conv(x))\mathrm{silu}(\mathrm{conv}(x))(见下一个旁注)。

输入单头向量 xtRdx_t \in \mathbb{R}^d,卷积核窗口大小 W(标准取 4),深度可分离(depthwise)且因果:

ShortConv(x)t=j=0W1wjxtj\mathrm{ShortConv}(x)_t = \sum_{j=0}^{W-1} w_j \odot x_{t-j}

记号 含义
wjRdw_j \in \mathbb{R}^d 每通道独立卷积权重–depthwise,每个特征通道一套独立卷积核
\odot 逐元素相乘
因果约束 tj0t - j \ge 0;t<jt < j 时零填充,看不到未来 token

三个性质决定了它为什么放在这个位置:

  1. 时序维度滑动,只混合最近 W 个历史 token,开销远小于全局注意力;
  2. depthwise:wkw_kxtkx_{t-k} 逐通道相乘,通道间不混–局部上下文的注入不破坏各通道独立的衰减语义(与 Diag(αt)\mathrm{Diag}(\alpha_t) 逐通道门配套);
  3. 因果:与上一节因果时序卷积同理,LLM 自回归必备。

直觉上,它给 KDA 补上了线性 RNN 天生缺失的「局部听觉」–递归状态只携带压缩后的全局历史,最近几个 token 的精细局部模式(词内字符、短程搭配)由这层小卷积负责。

旁注:Swish(SiLU,自门控激活函数)

流程图里 ShortConv 之后、L2Norm 之前的 SiLU,即 Swish–严格说 Swish 带可学参数时是 xσ(βx)x \cdot \sigma(\beta x),论文和实现里固定 β=1\beta = 1,两个名字就此等价:

Swish(x)=xσ(x),σ(x)=11+ex\mathrm{Swish}(x) = x \cdot \sigma(x), \qquad \sigma(x) = \frac{1}{1 + e^{-x}}

分层计算只有两步:1对输入的每个元素算 σ(x)=1/(1+exp(x))\sigma(x) = 1/(1+\exp(-x));2原输入 x 与 sigmoid 结果逐元素相乘得到输出。Swish = 输入 × 输入自己的 sigmoid 门–自门控(self-gated):门控信号不是外部来的,是输入自身,免参数。

为什么这层用它:1平滑非单调(负区间有个小下凹,xx \to -\infty 时输出趋于 0 而非 ReLU 的硬截断),梯度处处非零,深网训练稳;2门控形式与整个 block 的「信息流过 vs 被抑制」语义一致–Q/K/V 进递归状态前先过一道自门控,相当于对局部卷积混合后的特征做一次软筛选,再交给 L2Norm 钉上单位球。K3 把输出侧的这门升级为全秩输入相关门(满秩门控一节),源头就是这里的 SiLU。

旁注:L2Norm(逐头 L2 归一化,Q/K 专用)

流程图最后一框。对单头向量 zRdk\bm{z} \in \mathbb{R}^{d_k},逐通道 L2 标准化到单位范数:

L2Norm(z)=zz2+ϵ,z2=i=1dkzi2\mathrm{L2Norm}(\bm{z}) = \frac{\bm{z}}{\|\bm{z}\|_2 + \epsilon}, \qquad \|\bm{z}\|_2 = \sqrt{\sum_{i=1}^{d_k} z_i^2}

记号 含义
ϵ\epsilon 极小防除零常数(10610^{-6} 左右)
操作维度 每个注意力头独立归一化,跨头不共享统计

为什么只有 Q/K 用、V 不用?回到 DeltaNet 一节埋下的伏笔:归一化后 kt2=1\|k_t\|_2 = 1,于是 delta rule 的擦除矩阵 IβkkI - \beta k k^\top 的特征值被夹在 [1β, 1][1-\beta,\ 1]–写入步长天然稳定,β=1 时是精确保长反射。Q 同样归一化则是让读出 o=Sqo = S^\top q 的尺度可控,查询和 key 在同一单位球上内积才有可比性,chunkwise 公式里 QKQK^\top 的 score 变成有余弦界的小量。V 不归一化:写入内容的幅值本身携带信息(value 的模长是特征,不是 bug),而且 βtktvt\beta_t k_t v_t^\top 的稳定性由 k 侧保证,与 v 的尺度无关。这与 softmax 注意力里的 QK-norm(防 logit 爆炸)动机同源,但线性 RNN 里它 additionally 承担擦除几何的稳定性–归一化不是锦上添花,是 delta rule 能正常工作的前提。

至此三个旁注凑齐流程图的 Q/K 支路:Linear 投影 -> ShortConv(局部混合) -> SiLU(软筛选) -> L2Norm(单位球),三步全部因果、逐通道、免跨头统计,正好匹配 Diag(αt)\mathrm{Diag}(\alpha_t) 的逐通道语义。更重要的是三者各管一段数值安全:ShortConv 管局部混合、L2Norm 管幅度稳定、逐通道衰减差 γiγj0\gamma_i - \gamma_j \le 0 管指数安全(下界衰减一节)–全链路数值安全正是 KDA 能在 BF16 上跑起来的主线。

混合架构(论文的另一贡献):GDN + 滑窗注意力(H1)或 Mamba-2 + GDN + SWA(H2)交错堆叠,取长补短–与 K3 的「3 KDA + 1 MLA」同款思路,GDN 论文可视为这个混合范式的先声。

从 GDN 到 KDA:K3 继承了什么、改了什么
维度 GDN KDA(Kimi Linear -> K3) 动机
更新式 St1(αt(Iβkk))+βvkS_{t-1}(\alpha_t(I-\beta kk^\top)) + \beta vk^\top (Iβkk)Diag(αt)St1+βvk(I-\beta kk^\top)\,\mathrm{Diag}(\alpha_t)S_{t-1} + \beta vk^\top α 从标量 -> 逐通道向量
遗忘门粒度 标量(整头同一衰减率) 向量 αt(0,1)dk\alpha_t\in(0,1)^{d_k} 长期记忆通道与短期工作区分离
门作用顺序 α 在外乘整个更新 Diag(α) 在内侧先衰减 S,再删写 与逐通道参数化配套(KCP 推导需要)
α 参数化 负 Softplus,值域 (,0)(-\infty,0) K3:缩放 sigmoid,下界 gmin=5g_{\min}=-5 1/Γ1/\Gamma 有界 -> 对角 tile 全走 Tensor Core
chunkwise WY + UT + 衰减箭头 同框架,Γ 从标量比变成向量累积比 表达力↑ 数值难度↑(K3 用下界解决)
输出门 低秩/简单门 K3:输入相关全秩 sigmoid 门 逐通道调节读出

两处值得注意的「暗改」:

  1. 作用顺序变了:GDN 是 S(α(Iβkk))S(\alpha(I-\beta kk^\top))(α 吸进 Householder 连乘),KDA 是 (Iβkk)Diag(α)S(I-\beta kk^\top)\mathrm{Diag}(\alpha)S(α 先作用、再删写)。KDA 的逐通道门是矩阵 Diag(α)\mathrm{Diag}(\alpha),与 Householder 不交换,顺序成为实质设计选择–它使 KCP 的「段转移分解」成为可能。
  2. 衰减在 Γ 处的复杂化:GDN 的 γ 是标量连乘,Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 只是数;KDA 的 Γ 是向量逐元素累积,chunk 公式里 K/Γ、Q⊙Γ 等运算随之复杂化,数值范围问题(1/Γ1/\Gamma 爆炸)由此而生–K3 的下界衰减正是对 GDN->KDA 这一步引入的新问题的修复。

SSM 部分速览,详细推导见前置篇

KDA 递推式的 SSM 部分来自一条完整的推导链:连续状态方程 h(t)=Ah(t)+Bx(t)h'(t) = Ah(t) + Bx(t),经积分因子法得通解,ZOH 离散化得递推 ht=Aˉht1+Bˉxth_t = \bar A h_{t-1} + \bar B x_t,其中 Aˉ=eΔA\bar A = e^{\Delta A}Bˉ=A1(eΔAI)B\bar B = A^{-1}(e^{\Delta A} - I)B;递归展开等价于卷积(S4 双形式);Mamba 让 Δt\Delta_t 依赖输入,Aˉt=eΔta(0,1)\bar A_t = e^{-\Delta_t a} \in (0,1) 变成看内容的遗忘门;Mamba-2 把状态写成外积形式 St=αtSt1+vtktS_t = \alpha_t S_{t-1} + v_tk_t^\top,和线性注意力在数学上就是一回事了(SSD 对偶)。每一步的严格推导、数值验算(ZOH 精确性 1.2642 = 1.2642)和 S4/Mamba/Mamba-2 三站的演化,见前置篇《线性注意力与 SSM:两条技术路线的完整推导》

这里只需要记住结论:SSM 这条路线最终给出的,是乘在旧状态上的标量衰减系数 αt\alpha_t–它回答了线性注意力一族「状态该怎么遗忘」的问题,但答案是整片擦除、不分通道。它和 DeltaNet(定点覆写)在 GDN 处碰头;GDN 又被 KDA 逐通道化。骨架始终是同一条:

状态×(衰减/删除算子)+(写入项)\text{状态} \times (\text{衰减/删除算子}) + (\text{写入项})

递推公式(论文 2.1.1 节,核心式 1)

St=(Iβtktkt)Diag(αt)St1+βtktvt\mathbf{S}_t = \left(\mathbf{I} - \beta_t \bm{k}_t \bm{k}_t^\top\right) \operatorname{Diag}(\bm{\alpha}_t)\, \mathbf{S}_{t-1} + \beta_t \bm{k}_t \bm{v}_t^\top

o~t=Stqt\tilde{\bm{o}}_t = \mathbf{S}_t^\top \bm{q}_t

符号逐个过一遍(论文原定义):

符号 含义
αtRdk\bm{\alpha}_t \in \mathbb{R}^{d_k} 逐通道一维保留因子向量(channel-wise one-step retention factor)
Diag(αt)\operatorname{Diag}(\bm{\alpha}_t) 向量转对角矩阵算子,把通道级衰减系数变成矩阵乘
I\mathbf{I} 同维度单位矩阵
βt(0,1)\beta_t \in (0,1) delta rule 写入强度

读这个式子的三层动作(从右往左):

  1. 通道衰减:Diag(αt)St1\operatorname{Diag}(\bm{\alpha}_t)\mathbf{S}_{t-1} 等价于对 St1\mathbf{S}_{t-1} 的每一行分别乘以对应通道的 α\alpha 系数–逐通道缩放历史记忆。普通线性注意力/SSM 大多全局统一衰减(αt\alpha_t 是标量),KDA 给每一个 key 通道分配独立衰减权重:不同语义通道可以选择更快/更慢遗忘 St1\mathbf{S}_{t-1};
  2. 方向擦除:(Iβtktkt)\left(\mathbf{I} - \beta_t \bm{k}_t \bm{k}_t^\top\right) 再对衰减后的状态做当前 token 的定点擦除(秩 1 Householder 型,详见 DeltaNet 一节);
  3. 新信息写入:叠加 βtktvt\beta_t \bm{k}_t \bm{v}_t^\top

三步合起来:先通道衰减,再方向擦除,最后写入–带通道精细遗忘的 delta 递推。作用顺序是设计选择:Diag 与 Householder 不交换,这个顺序使 KCP 的段转移分解成为可能(见下界衰减与基础设施篇)。

记号约定:论文里大写 Diag()\operatorname{Diag}(\cdot) 是「向量 -> 对角矩阵」算子;小写 diag()\operatorname{diag}(\cdot) 有时指反向操作(输入矩阵、提取对角线为向量)。本文全程大写表示向量转对角矩阵。

为什么是逐通道门:动机与在线学习视角

标量门的表达力瓶颈。GDN 的 αt(0,1)\alpha_t \in (0,1) 是标量:每一步遗忘时,所有 key 通道以同一个比例衰减–要么一起记住,要么一起忘记。但不同通道承担的角色不同:

  • 有的通道在存「长期主题」(希望 α1\alpha \approx 1,几乎不遗忘);
  • 有的通道在存「临时指针」(希望快速衰减,腾出容量)。

KDA 的核心改动只有一处:把标量换成向量 αt(0,1)dk\bm{\alpha}_t \in (0,1)^{d_k}(βt\beta_t 仍是标量,这是 KDA 的选择而非必须)。直觉:S 的第 j 列对应 key 空间的第 j 个通道,Diag(αt)\mathrm{Diag}(\alpha_t) 作用上去就是给每一列配一个独立的遗忘速度。GDN 是 KDA 在 Diag(αt)=αtI\mathrm{Diag}(\alpha_t) = \alpha_t I 时的特例。

在线学习视角:逐通道权重衰减的 delta rule。与 GDN 一节的在线学习表同构,只需把正则项换成逐通道版。每步给定新样本 (kt,vt)(k_t, v_t),希望新状态 S:1拟合新样本 SktvtSk_t \approx v_t;2不偏离「衰减后的旧记忆」太远 SSt1Diag(αt)S \approx S_{t-1}\mathrm{Diag}(\alpha_t)(以下记 Dt=Diag(αt)D_t = \mathrm{Diag}(\alpha_t),αt=exp(gt)\alpha_t = \exp(g_t),gtR<0dkg_t \in \mathbb{R}_{<0}^{d_k}):

Lt(S)=12Sktvt2拟合+12SSt1DtF2逐通道正则\mathcal{L}_t(S) = \underbrace{\tfrac12 \|S k_t - v_t\|^2}_{\text{拟合}} + \underbrace{\tfrac12 \|S - S_{t-1} D_t\|_F^2}_{\text{逐通道正则}}

St1DtS_{t-1}D_t 出发对第一项做一步梯度下降(步长 βt\beta_t):

St=St1Dtβt(St1Dtktvt)kt=St1Dt(Iβtktkt)+βtvtktS_t = S_{t-1}D_t - \beta_t\,(S_{t-1}D_t\, k_t - v_t)\,k_t^\top = S_{t-1} D_t (I - \beta_t k_t k_t^\top) + \beta_t v_t k_t^\top

正是递推式(转置约定下)。逐通道正则的含义:第 j 列的「信任区域」宽度正比于 αt(j)\alpha_t^{(j)}α\alpha 小的通道,旧记忆先被压缩,新写入覆盖几乎没有阻力(快遗忘);α1\alpha \approx 1 的通道,旧记忆原样进入下一步,delta rule 只做精细增量(慢遗忘)。这就是「细粒度记忆控制」的优化论表述:KDA = 逐通道权重衰减 + delta rule。

实现注意:作用顺序是坑DtD_t 作用于整个 St1S_{t-1}、先于擦除项–擦除时检索用的也是已衰减的状态:

St=St1Dt先衰减(Iβtktkt)再擦除/写入S_t = \underbrace{S_{t-1} D_t}_{\text{先衰减}}\underbrace{(I - \beta_t k_t k_t^\top)}_{\text{再擦除/写入}}

写成 St1(Iβkk)DtS_{t-1}(I-\beta kk^\top)D_t(顺序颠倒)或只衰减单位阵部分都是错的–DtD_t(Iβkk)(I - \beta kk^\top) 不可交换,顺序错了结果就错(与 GDN 一节「α 乘整个括号」的警告同源,但逐通道化之后错误更隐蔽)。FLA 参考实现里对应 S = S * g.exp() 之后立刻用衰减后的 S 做检索 v - k^T S

下界衰减(Lower-bounded decay)

KDA 的衰减参数化是一个关键改进。Kimi Linear 使用无界的负 Softplus 映射 g=eASoftplus(z)(,0)g = -e^A \text{Softplus}(z) \in (-\infty, 0),而 K3 改用有界缩放 sigmoid:

gth=gminSigmoid(eAhzth)(gmin,0)g_t^h = g_{\min} \cdot \text{Sigmoid}(e^{A_h} z_t^h) \in (g_{\min}, 0)

其中 gmin=5g_{\min} = -5 固定,AhA_h 是可学习的每头对数尺度。这意味着每个保留因子满足 α>e56.7×103\alpha > e^{-5} \approx 6.7 \times 10^{-3},16 token 块的累积对数衰减落在 (80,0)(-80, 0) 内,对应的重缩放因子小于 e80e^{80},在 BF16 动态范围内。

计算收益:有限范围使得因果对角块和离对角块都能用密集 Tensor Core 矩阵乘法,消除了 Kimi Linear 中需要的 position-pair 对角计算路径。

问题从哪来:1/Γ1/\Gamma 爆炸是向量门控自带的代价

为什么这个修复在 GDN 上不必要、在 KDA 上变成刚需?GDN 的标量衰减在 chunkwise 里只以比值出现(s=j+1iαs1\prod_{s=j+1}^{i}\alpha_s \le 1,iji \ge j),天然安全。KDA 把 Γ\Gamma 变成向量累积 Γt=exp(stgs)\Gamma_t = \exp(\sum_{s\le t} g_s),逐通道独立:第 4 步式的推导可以刻意只用 iji \ge j 的差(指数 0\le 0,安全);但只要换一种等价写法–把状态「反归一化」回 chunk 起点、或把衰减从 key 上整体外提(论文公式 (4) 的 K/ΓK/\Gamma 因式分解正是这种写法)–就得真的算出

Γt1=exp(stgs)(逐通道)\Gamma_t^{-1} = \exp\Big(-\sum_{s \le t} g_s\Big) \quad \text{(逐通道)}

衰减越快的通道,γt-\gamma_t 越大,1/Γ1/\Gamma 指数膨胀。量级感受(实测):

恒定 α\alpha t=100t=100 t=500t=500 t=1000t=1000 t=2000t=2000
0.9 3.8×1043.8\times10^{4} 7.6×10227.6\times10^{22} 5.7×10455.7\times10^{45} 3.3×10913.3\times10^{91}
0.5 1.3×10301.3\times10^{30} 3.3×101503.3\times10^{150} 1.1×103011.1\times10^{301} 溢出 float64
0.1 1010010^{100} 溢出 溢出 溢出

float64 上界约 e7091.8×10308e^{709} \approx 1.8\times10^{308};BF16 训练下几十步就出事。标量 -> 向量的推广放大了衰减率的动态范围,爆炸从例外变成常态–所以必须给衰减率本身夹界。

Kimi Linear 的绕行方案(补丁,不是修复):在对数空间算相对衰减(减法代替除法,不溢出),并把每个 chunk 再切成 16 token 的二级瓦片。效果:瓦片之间(非对角块)可以安全交给 Tensor Core 稠密矩阵乘;但瓦片内部(对角块)衰减可能极端,仍需逐位置对显式计算–position-pair 路径无法组织成大矩阵乘,吃不满 Tensor Core,成为块内主要瓶颈。

K3 的根治:不改公式,改参数化,让爆炸在数学上不可能发生(负 Softplus 允许 α\alpha 任意接近 0,即「一步清零」;缩放 sigmoid 不允许)。

表达力为什么不受损:快通道在瓦片内照样能把记忆衰减到 e801035e^{-80} \approx 10^{-35}(约等于零,只是不真的为零);长期遗忘靠多步连乘(每步 ×0.0067\times 0.0067,几十步后照样忘干净),不需要单步清零的能力。用一个下界换来全 Tensor Core 化,划算的买卖。

满秩门控(Full-rank gate)

K3 将 KDA 的输出门从低秩参数化改为输入相关的满秩投影。在递推输出经过 head-wise RMSNorm 后,应用数据相关的输出门控:

yt=Wo[Sigmoid(Wgxt)RMSNorm(o~t)]y_t = W_o [\text{Sigmoid}(W_g x_t) \odot \text{RMSNorm}(\tilde{o}_t)]

满秩门控允许每个 token 独立调制从循环状态读取的通道。

Chunkwise 并行形式

递推形式在推理(decode)时是优势–O(1)O(1) 状态更新;但在训练/prefill 时是灾难:每个 token 的状态依赖前一个 token,顺序循环让 GPU 的数千个核心集体围观一个 for 循环。Chunkwise 并行化把序列切成长度 C 的块:块内矩阵运算并行,块间只传状态。先看通用的复杂度框架,再看 KDA 逐通道门带来的推导细节。

通用框架Chunk 内(intra-chunk):token 交互用带衰减的因果注意力直接算,O(C2d)O(C^2 d)–C 是常数(64/128),对序列长度 N 线性。Chunk 间(inter-chunk):每个 chunk 对状态做一次递推更新,chunk 内外积先归约成固定大小的 state 增量,块间只传 d×dd \times d 的 state。总计算量:

O(NCd)chunk 内注意力+O(N/Cd2)chunk 间状态更新  =  2Nd2(固定项)+2NCd\underbrace{O\big(N \cdot C \cdot d\big)}_{\text{chunk 内注意力}} + \underbrace{O\big(N/C \cdot d^2\big)}_{\text{chunk 间状态更新}} \;=\; 2Nd^2(\text{固定项}) + 2NCd

C 形态 复杂度
C=1C = 1 每个 token 一个 chunk,chunk 内注意力消失 纯线性注意力递推,FLOPs 最少但不一定最快(GPU 对小矩阵乘利用率低)
C=NC = N 整个序列一个 chunk,递推消失 标准 O(N2)O(N^2) 注意力

实践中 C 取 64/128:小到让 chunk 间项不贵,大到让 C×CC\times C 注意力矩阵铺满 Tensor Core 的 tile。这与 S4 时代「训练用卷积、推理用递归」双形式一脉相承:选择性打破 LTI 后,chunkwise 就是卷积的继任者。

KDA 版推导:逐通道衰减下的 WY / UT

记 chunk 起点传入状态 S[0]S_{[0]}(以下用局部下标 i=1..Ci = 1..C;乘积按时间倒序)。核心记号是累积 log 衰减:

γi=s=1igsRdk,Γij=diag(exp(γiγj))(ij)\gamma_i = \sum_{s=1}^{i} g_s \in \mathbb{R}^{d_k}, \qquad \Gamma_{i \leftarrow j} = \operatorname{diag}\big(\exp(\gamma_i - \gamma_j)\big) \quad (i \ge j)

Γij\Gamma_{i\leftarrow j} 是「从第 j 步衰减到第 i 步」的逐通道算子。两条性质:iji \ge jγiγj0\gamma_i - \gamma_j \le 0 逐分量成立(gs<0g_s < 0),元素都在 (0,1](0,1],数值安全;反向 Γji1\Gamma_{j\leftarrow i}^{-1} 的元素 1\ge 1 且随 chunk 长度指数增长–这是 1/Γ1/\Gamma 爆炸的根源,先记住这个观察。

第一步:衰减 KKT 矩阵 M。类比 GDN chunkwise 里的普通 kikjk_i^\top k_j,逐通道版需要带衰减的 key-key 内积:

Mci=(kceγcγi)ki,1i<cCM_{ci} = \big(k_c \odot e^{\gamma_c - \gamma_i}\big)^\top k_i, \qquad 1 \le i < c \le C

含义:kik_i 写入的记忆衰减到第 c 步时,与 kck_c 的重叠程度。Hadamard 积 \odot 作用在 key 通道维–正是逐通道衰减出现的位置。

第二步:UT 变换。构造严格下三角 LL:Lci=βiMci (c>i)L_{ci} = \beta_i M_{ci}\ (c > i),求 T=(I+L)1T = (I+L)^{-1}(幂零,有限项截断;实践中不显式求逆,前代法逐行解),再每列乘 βj\beta_jA^=Tdiag(β)\hat{A} = T\,\operatorname{diag}(\beta)。FLA 实现里那个三行循环,数学上就是解 (I+L)T=I(I+L)T = I

与 GDN 对照:GDN 这一步的 M 是普通 kckik_c^\top k_i(标量衰减被拆成比值吸收进别的项);KDA 里衰减「长」在 M 内部,无法外提–这是向量门控带来的结构性变化,KDA chunkwise 推导的核心难点。

第三步:WY 表示:

W=A^(eγiki)i=1..CRC×dk,U=A^VRC×dv,V~=UWS[0]W = \hat{A}\,\big(e^{\gamma_i} \odot k_i\big)_{i=1..C} \in \mathbb{R}^{C \times d_k}, \qquad U = \hat{A}\,V \in \mathbb{R}^{C \times d_v}, \qquad \tilde{V} = U - W\,S_{[0]}^\top

伪值 V~\tilde V 到底是什么:一本账。UT 变换产出 U 和 W,用它们定义 V~:=UWS\tilde V := U - WS。三个符号各司其职:

符号 形状 角色
S[0]S_{[0]} dk×dvd_k \times d_v 旧账本:chunk 之前所有 token 写入的记忆,块内计算时固定不变
WW C×dkC \times d_k 翻账本的钥匙:每行是某位置的衰减 key 修正组合;WSWS 即「这个位置能从旧记忆里读到什么」
UU C×dvC \times d_v 想记的新账(已和同伴对过账):A^\hat A 作用在 V 上,块内写入重叠已扣除
V~\tilde V C×dvC \times d_v 真正上账的净额 = 新账 - (旧记忆已有的 + 同伴已写的)

为什么叫「伪」值:它们不是真实的 v(手算例子里 v~2=(2,2)v2\tilde v_2 = (-2,2) \ne v_2),而是预扣完所有重叠后可以不做任何修正直接累加的修正值。这正是 delta rule 的灵魂的单步版:单步写入 vtSt1ktv_t - S_{t-1}^\top k_t = 目标值 - 旧记忆读数;V~\tilde V 是它的 chunk 并行版,每行回答「这个位置真正新增了多少」。账在 V~\tilde V 里算清了,后面的块内注意力 AV~A\tilde V 和状态更新才敢放心一锤子矩阵乘。

第四步:输出与跨 chunk 状态:

Acjqk=(qceγcγj)kj (jc),oc=S[0](qceγc)块间:query 打折查旧状态+jcAcjqkv~j块内:衰减注意力A^{qk}_{cj} = \big(q_c \odot e^{\gamma_c - \gamma_j}\big)^\top k_j \ (j \le c), \qquad o_c = \underbrace{S_{[0]}\,(q_c \odot e^{\gamma_c})}_{\text{块间:query 打折查旧状态}} + \underbrace{\sum_{j \le c} A^{qk}_{cj}\,\tilde v_j}_{\text{块内:衰减注意力}}

S[C]=S[0]ΓC0+c=1Cv~c(kceγCγc)S_{[C]} = S_{[0]}\,\Gamma_{C\leftarrow 0} + \sum_{c=1}^{C} \tilde v_c \big(k_c \odot e^{\gamma_C - \gamma_c}\big)^\top

出口状态读法:旧状态整体按整 chunk 累积衰减缩小;每个伪值以「写入时刻衰减到块末」的 key 为地址写入。所有指数都是 0\le 0 的差,整条链路没有一个 1\ge 1 的因子–刻意保持,原因见下界衰减一节。

手算一遍(dk=dv=2d_k = d_v = 2,C = 2)

零初始状态,每步恒定衰减 α=(0.5,0.25)\alpha = (0.5, 0.25)(即 g=(ln0.5,ln0.25)g = (\ln 0.5, \ln 0.25)),β 取 1:

k1=(11), v1=(20);k2=(12), v2=(02);q1=q2=(11)k_1 = \binom{1}{1},\ v_1 = \binom{2}{0};\quad k_2 = \binom{1}{2},\ v_2 = \binom{0}{2};\quad q_1 = q_2 = \binom{1}{1}

递推式。第 1 步(S0=0S_0 = 0):S1=v1k1=(2200)S_1 = v_1k_1^\top = \begin{pmatrix}2&2\\0&0\end{pmatrix},o1=(4,0)o_1 = (4,0)。第 2 步,D=diag(0.5,0.25)D = \mathrm{diag}(0.5, 0.25),先衰减 S1D=(10.500)S_1D = \begin{pmatrix}1&0.5\\0&0\end{pmatrix}(第 1 列 ×0.5、第 2 列 ×0.25–逐通道在动),再擦除写入:

S1D(Ik2k2)=(13.500),S2=(13.524),o2=(4.5, 6)S_1D(I - k_2k_2^\top) = \begin{pmatrix}-1&-3.5\\0&0\end{pmatrix}, \qquad S_2 = \begin{pmatrix}-1&-3.5\\2&4\end{pmatrix}, \qquad o_2 = (-4.5,\ 6)

Chunkwiseeγ1=(0.5,0.25)e^{\gamma_1} = (0.5, 0.25),eγ2=(0.25,0.0625)e^{\gamma_2} = (0.25, 0.0625),eγ2γ1=(0.5,0.25)e^{\gamma_2-\gamma_1} = (0.5, 0.25)。衰减 KKT:M21=(k2eγ2γ1)k1=(0.5,0.5)(1,1)=1M_{21} = (k_2 \odot e^{\gamma_2-\gamma_1})^\top k_1 = (0.5, 0.5)\cdot(1,1) = 1。UT:L=(0010)L = \begin{pmatrix}0&0\\1&0\end{pmatrix},T=IL=(1011)T = I - L = \begin{pmatrix}1&0\\-1&1\end{pmatrix},A^=T\hat A = T。WY:

W=A^(0.50.250.250.125)=(0.50.250.250.125),U=A^(2002)=(2022)W = \hat A\begin{pmatrix}0.5&0.25\\0.25&0.125\end{pmatrix} = \begin{pmatrix}0.5&0.25\\-0.25&-0.125\end{pmatrix}, \quad U = \hat A\begin{pmatrix}2&0\\0&2\end{pmatrix} = \begin{pmatrix}2&0\\-2&2\end{pmatrix}

伪值(S[0]=0V~=US_{[0]}=0 \Rightarrow \tilde V = U):v~1=(2,0)\tilde v_1 = (2,0),v~2=(2,2)\tilde v_2 = (-2,2)。注意 v~2v2\tilde v_2 \ne v_2:因为 k2k_2 与衰减后的 k1k_1 写入重叠(M21=10M_{21} = 1 \ne 0),WY 把 v2v_2 修正为扣除重叠后真正的新增–单步 delta rule 的 chunk 版样子。衰减注意力与输出:

Aqk=(200.753),o1=2v~1=(4,0) ,o2=0.75v~1+3v~2=(4.5, 6) A^{qk} = \begin{pmatrix}2&0\\0.75&3\end{pmatrix}, \qquad o_1 = 2\tilde v_1 = (4,0)\ \checkmark, \qquad o_2 = 0.75\,\tilde v_1 + 3\,\tilde v_2 = (-4.5,\ 6)\ \checkmark

块末状态:S[2]=v~1(k1eγ2γ1)+v~2k2=(13.524)S_{[2]} = \tilde v_1(k_1 \odot e^{\gamma_2-\gamma_1})^\top + \tilde v_2 k_2^\top = \begin{pmatrix}-1&-3.5\\2&4\end{pmatrix} \checkmark 与递推式逐项一致。(随机对拍:非零初始状态、随机门控下递推 vs chunkwise 最大误差 10910^{-9} 量级。)

与论文公式 (4) 的对照:K/Γ 因式分解

K3 报告(沿用 Kimi Linear)用乘积记号:γij=r=ijαr=exp(r=ijgr)\gamma_{i\to j} = \prod_{r=i}^{j}\alpha_r = \exp(\sum_{r=i}^j g_r);ΓRC×dk\Gamma \in \mathbb{R}^{C\times d_k} 把各步 γ\gamma 按行堆叠(注意 Γ\Gamma 本身不是对角阵,每个位置 rrdiag(γr)\mathrm{diag}(\gamma_r) 才是遗忘对角阵,Γ\Gamma 是 C 个对角阵的打包)。论文状态约定 SRdk×dvS \in \mathbb{R}^{d_k \times d_v}(本文的转置),其公式 (4):

A=Tril[(QΓ)(K/Γ)],O=(ΓQ)S块间+AV~块内A = \mathrm{Tril}\big[(Q\odot\Gamma)(K/\Gamma)^\top\big], \qquad O = \underbrace{(\Gamma\odot Q)S}_{\text{块间}} + \underbrace{A\,\tilde V}_{\text{块内}}

为什么 A 能这样拆:看 (i,j)(i,j) 元素 Aij=dqi,dγi,dkj,d/γj,d=(qiγji)kjA_{ij} = \sum_d q_{i,d}\gamma_{i,d}\cdot k_{j,d}/\gamma_{j,d} = (q_i \odot \gamma_{j\to i})^\top k_j,与 AqkA^{qk} 逐元素相同–逐对位置的衰减比值被因式分解成 query 侧乘 Γ、key 侧除以 Γ 两个逐位置操作,整个 C×CC\times C 矩阵 = 一次稠密 matmul + 两次逐元素乘,完全并行。Tril 保留对角线:delta rule 里 oio_i 读的是写入当前 token 之后的状态。

用手算例子核对:QΓ=(0.50.250.250.0625)Q\odot\Gamma = \begin{pmatrix}0.5&0.25\\0.25&0.0625\end{pmatrix},K/Γ=(24432)K/\Gamma = \begin{pmatrix}2&4\\4&32\end{pmatrix},乘积 Tril 后 A=(200.753)A = \begin{pmatrix}2&0\\0.75&3\end{pmatrix},与上面 AqkA^{qk} 一致。注意 K/ΓK/\Gamma 里已出现 32 这种被放大的数–1/Γ1/\Gamma 的膨胀就藏在这一步,数值后果见下界衰减一节。

常见疑问:CCdkd_k 有倍数关系吗? 没有。C 是序列轴的切分(一个 chunk 装多少 token),dkd_k 是特征轴的宽度,两根轴独立。有整除要求的是:K3 把 chunk 再切成 16 token 二级瓦片,故 C 需是 16 的倍数;dkd_kdvd_v 按 Tensor Core 友好尺寸取是工程约束,不是数学要求。

KDA 的核心优势是:相比传统 softmax attention,它是线性复杂度的;相比纯粹的线性 RNN,它通过 delta rule 实现了更强大的信息写入和遗忘控制。

一页纸速查(KDA/GDN 全家)
维度 内容
状态 S ∈ R^{d_v×d_k}(论文记法),固定大小,每序列每头一份
更新 S = S(α(I-βkkT)) + βvkT:板擦 α + 铅笔 δ
读出 o = Sq,后接 RMSNorm + sigmoid 输出门
α GDN 标量 / KDA 逐通道,数据相关;β:sigmoid
q/k Linear -> ShortConv -> SiLU -> L2Norm
等价视角 在线最小二乘 SGD(β=学习率)+ 自适应权重衰减(α)
并行训练 chunkwise:WY 表示(W 无 γ / Ũ 有 γ)+ UT 变换(下三角求逆)+ 衰减箭头
chunk 输出 O = Q⃖S0T + (QKT⊙Γ因果)(Ũ - W⃖S0T)
chunk 状态 S_C = γ_C S0 + (Ũ-> - W->S0T)TK
记忆瓶颈修复 覆写(δ)治叠加污染;清场(α)治长期堆积
后继 KDA:α->逐通道、换作用顺序、加下界 -> 进 K3

参考:

DSpark 的实现和测评

DSpark = DFlash 的并行 backbone forward(1 次)+ N 步轻量 Markov 序列修正,全部在 CUDA Graph 内。 本文结合 vLLM 源码分析 DSpark 的实现细节,并在 Qwen3-8B 上实测 deepseek-ai 官方 draft 和社区 Dogacel draft 的效果差异。


1. 背景:投机解码与并行起草

投机解码(Speculative Decoding, SD)用一个小 draft 模型并行猜测 N 个 token,再由 target 模型一次 verify,通过 rejection sampling 保证输出分布不变。SD 的收益来自把 decode 阶段的 memory-bound 转为 compute-bound–bs=1 时 GPU 利用率极低,draft 的轻量 GEMM + target 的 batched verify 填充了 GPU 空闲。

vLLM v1 的 SD 框架支持多种 method:eagleeagle3dflashdsparkmedusangrammtp 等。DSpark 继承自 DFlash,核心改进是序列马尔可夫采样

继承链:

1
2
3
4
BaseSpeculator (ABC)
└─ DraftModelSpeculator
└─ DFlashSpeculator
└─ DSparkSpeculator ← 本文主角

模型类继承链:

1
2
3
Qwen3ForCausalLM
└─ DFlashQwen3ForCausalLM
└─ Qwen3DSparkForCausalLM

2. DSpark vs DFlash:两个核心差异

DSpark 的 docstring 写得非常清楚,和 DFlash 的差异只有两点。

2.1 Anchor-as-first-prediction(锚位即首预测)

DFlash:每个 request 发 1 + N 个 query token(1 个 anchor/bonus + N 个 mask token)。anchor 是上一步验证通过的 token,只有 N 个 mask 位置做预测:

1
2
3
4
5
DFlash query layout (1+N=9, N=8):
[anchor] [mask] [mask] [mask] [mask] [mask] [mask] [mask] [mask]
↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑
bonus pred pred pred pred pred pred pred pred
(不采样)

DSpark:anchor 本身也是预测位置,每个 request 只发 N 个 query token:

1
2
3
4
5
DSpark query layout (N=8):
[anchor] [noise] [noise] [noise] [noise] [noise] [noise] [noise]
↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑
pred pred pred pred pred pred pred pred
(采样)

代码(DSparkSpeculator.__init__):

1
2
3
4
5
6
7
self.sample_from_anchor = getattr(
self.draft_model_config.hf_config, "sample_from_anchor", True
)
if self.sample_from_anchor:
self.num_query_per_req = self.num_speculative_steps # N
else:
self.num_query_per_req = 1 + self.num_speculative_steps # 1+N (兼容旧格式)

在 Triton kernel _prepare_dflash_inputs_kernel 中,通过 SAMPLE_FROM_ANCHOR 编译常量控制采样行为:

1
2
3
4
# DSpark: 所有 N 个位置都采样,sample_pos = query_pos + 1(标准 next-token)
sample_off = 0 if SAMPLE_FROM_ANCHOR else 1
is_sample = is_query & (query_off >= sample_off)
sample_pos = query_pos + 1 if SAMPLE_FROM_ANCHOR else query_pos

2.2 Sequential Markov Sampling(序列马尔可夫采样)

这是 DSpark 的核心创新。

DFlash:N 个 mask 位置的 hidden states 一次性并行采样,各位置之间无依赖。

DSpark:先并行 forward 得到所有 N 个位置的 hidden states,然后从左到右逐个采样,每步用前一个采样出的 token 注入一个 Markov bias:

1
2
3
4
5
6
7
8
9
10
11
12
并行 backbone forward → [h₀, h₁, h₂, ..., h₇]
│ │ │ │
▼ ▼ ▼ ▼
base_logits[0] base_logits[1] ... base_logits[7]
+ + +
markov_bias( markov_bias( markov_bias(
anchor) sample₀) sample₆)
│ │ │
▼ ▼ ▼
sample₀ sample₁ ... sample₇

└──────────────→ 传给下一步作为 prev

代码在 _sample_sequential

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
def _sample_sequential(self, num_reqs, head_hidden):
n_spec = self.num_speculative_steps
# 1. 一次性算出所有 N 个位置的 base logits
base_logits = self.model.compute_draft_logits(sample_hidden) # [B, N, V]

# 2. anchor token 作为初始 prev
prev = self.input_buffers.input_ids[self._anchor_idx[:num_reqs]]

# 3. 逐位置采样
for i in range(n_spec):
markov_embed = self.model.markov_embed(prev) # [B, r]
bias = self.model.markov_bias(markov_embed) # [B, V]
logits_i = base_logits[:, i] + bias # 加上 Markov 偏置
draft_sampled_i = gumbel_sample(logits_i, ...) # 采样
self.draft_tokens[:num_reqs, i] = draft_sampled_i
prev = draft_sampled_i # 传给下一步

一句话总结:并行 forward 拿到所有位置的 base prediction,再用 N 步轻量 Markov 修正注入序列依赖–把「N 个独立预测」变成「N 个有依赖的预测」。


3. Markov Head 结构

DSparkMarkovHead 是一个 low-rank 转移偏置头:

1
2
3
4
5
6
7
8
9
prev_token_id

│ markov_w1: Embedding(V, r) ← V 是 vocab_size,r 是 markov_rank

markov_embed [B, r]

│ markov_w2: ParallelLMHead(r, V) ← r → V 的线性投影

markov_bias [B, V] ← 加到 base_logits 上

代码(qwen3_dspark.py):

1
2
3
4
5
class DSparkMarkovHead(nn.Module):
def __init__(self, vocab_size, draft_vocab_size, markov_rank, ...):
self.markov_w1 = nn.Embedding(vocab_size, markov_rank) # V×r
self.markov_w2 = ParallelLMHead(
draft_vocab_size, markov_rank, bias=False, disable_tp=True) # r×V

两个权重都是 replicateddisable_tp=True),因为 Markov head 每步都跑,分片会引入 all-reduce 和 full-vocab gather。

参数量 = 2×V×r2 \times V \times r。当 V=151936V=151936(Qwen3 词表)、r=64r=64 时约 19.4M 参数,相比 8B backbone 可以忽略。


4. 完整的 Draft 一步流程

DSparkSpeculator._generate_draft 只有两行:

1
2
3
def _generate_draft(self, num_reqs, num_tokens_padded, ...):
head_hidden = self._run_model(...) # 1. 并行 backbone forward
self._sample_sequential(num_reqs, head_hidden) # 2. 序列 Markov 采样

Step 1:并行 Backbone Forward(继承自 DFlash)

  • 输入:N 个 query token(anchor + mask/noise),position 已对齐
  • 上下文 KV 已在 precompute_and_store_context_kv 中预填充
  • 非因果 attention:N 个 query 位置可以互相 attend
  • 整个 forward 被 CUDA Graph 捕获

Step 2:Sequential Markov Sampling(DSpark 独有)

  • 取出 N 个位置的 hidden states
  • 一次性算出 base logits(compute_draft_logits
  • 逐位置:base_logits[i] + markov_bias(prev) -> gumbel_sample
  • 这个循环也被 CUDA Graph 捕获(所有 buffer 预分配固定地址)

Context KV 预计算(DFlash 的关键优化)

避免逐层跑 target 的 forward 来填充 draft KV cache,而是用 target 的中间层 hidden states 一次性投影:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
target aux hidden states [num_ctx, H_target]

│ fc 层投影到 draft hidden size

context_states [num_ctx, H_draft]

│ ① Fused GEMM(所有层的 KV projection 合成一个矩阵乘法)

all_kv_flat [num_ctx, L×2×kv_size]

│ ② Grouped RMSNorm(所有层的 K-norm 一次算完)

all_k_normed [L, num_ctx, nkv, hd]

│ ③ Fused RoPE(所有层一次应用)

all_k_final [L, num_ctx, nkv, hd] → per-layer 写入 KV cache

代码核心(DFlashQwen3Model.precompute_and_store_context_kv):

1
2
3
4
5
# 融合所有层的 KV 权重做一次大 GEMM
all_kv_flat = F.linear(normed_context_states, self._fused_kv_weight, self._fused_kv_bias)
# 分离 K/V,per-layer 写入 cache
all_kv = all_kv_flat.view(num_ctx, L, 2, nkv, hd).permute(2, 1, 0, 3, 4).contiguous()
all_k, all_v = all_kv[0], all_kv[1]

5. Probabilistic Rejection Sampling 与 Reduced Vocab

DSpark 支持 draft_sample_method="probabilistic"(Gumbel-based rejection sampling)。Draft 采样时把 logits 通过 Gumbel max trick 得到 draft_logits,Target verify 时用相同 Gumbel seed 验证,保证输出分布不变。

支持 reduced draft vocab:draft 在小词表上算 logits,然后 scatter 到 target vocab 位置:

1
2
3
4
if self._d2t_scatter_index is not None:
buf = self._draft_scatter_buf[:num_reqs] # [-inf, -inf, ...]
buf.index_copy_(1, self._d2t_scatter_index, logits_i) # 只填 draft vocab 列
logits_i = buf # 变成 target vocab 大小

6. CUDA Graph 覆盖

DFlash/DSpark 的 CUDA Graph 是 FULL mode,覆盖整个 draft step:

1
2
3
# DFlashSpeculator.init_cudagraph_manager
if wants_full and supports_full:
cudagraph_mode = CUDAGraphMode.FULL_DECODE_ONLY

为了让 Markov 循环能被 CG 捕获,所有 buffer 都是预分配的固定地址:

Buffer 用途 CG 兼容性
draft_tokens 输出 token ✅ 固定地址
draft_logits probabilistic 模式的 processed logits ✅ 固定地址
_draft_scatter_buf reduced vocab scatter buffer ✅ 固定地址
_anchor_idx 每个 request 的 anchor 位置索引 ✅ 固定地址
input_buffers.input_ids anchor token 读取 ✅ 固定地址

7. 模型加载与权重共享

load_dspark_modeldspark/utils.py)做了几件事:

  1. 创建 draft config,设置非因果注意力
  2. 加载 draft 模型
  3. Embed tokens 共享:如果 draft 没有自己的 embedding,用 target 的
  4. LM head 共享:同理
1
2
3
if _should_share(draft_model, "has_own_embed_tokens", draft_embed, target_embed):
del draft_inner.embed_tokens
draft_inner.embed_tokens = target_embed

权重加载(Qwen3DSparkForCausalLM.load_weights):

  • 跳过 t2d(训练用映射,推理不需要)
  • d2t -> draft_id_to_target_id(推理用的 draft→target 映射)
  • 跳过 mask_embedding(DSpark 通过 vocab row 做 mask,不用单独参数)和 confidence_head(未接入推理)
  • 调用 _build_fused_kv_buffers() 构建 fused KV 权重

8. 实验环境

项目 配置
Target Model Qwen/Qwen3-8B
推理引擎 vLLM v0.26.0
conda 环境 dspark-vllm
nsys 版本 2026.1.3(vLLM traces)/ 2025.3.0(DeepSpec trace)
采集参数 -t cuda,nvtx,osrt,cudnn,cublas --python-backtrace=cuda --cudabacktrace=all
Benchmark SPEED-Bench(qualitative split, coding category)

Draft Model 配置

配置名 Draft Model 架构 来源
Baseline 无(纯 Qwen3-8B) Qwen3 -
DSpark(deepseek-ai) deepseek-ai/dspark_qwen3_8b_block7 Qwen3DSparkSt deepseek-ai 官方
DSpark(Dogacel) Dogacel/Qwen3-8B-DSpark EAGLE3 社区训练

Dogacel 的 vLLM 启动参数:speculative 开启、acceptance: 0.85num_spec_tokens: 7max_model_len: 2048


9. 性能对比

9.1 端到端性能(3 prompts, 各 64 tokens)

配置 耗时 vs Baseline 加速比
Baseline 0.75s - 1.00x
DSpark(deepseek-ai) 0.50s -33% 1.50x
DSpark(Dogacel) 0.77s +3% 0.97x

9.2 投机解码指标(DeepSpec evaluator trace)

指标 DSpark(deepseek-ai)
verify_steps 12
mean_accept_len 7.1
推测 每次提议 ~7 tokens,几乎全部被接受

Dogacel 的 trace 中未发现 dspark_propose / target_verify 的 NVTX range,推测 acceptance rate 极低。

mean_accept_len=7.1 意味着 N=8 时几乎全部接受–backbone 的并行预测质量极高,Markov head 的序列修正有效。verify_steps=12 表示 12 步验证共接受约 85 个 token(12×7.112 \times 7.1)。


10. Trace 分析

10.1 Trace 文件清单

文件 大小 来源 CUDA Kernel 数据
baseline_trace.nsys-rep 1.5 MB vLLM profile_baseline.py ❌ 无
dogacel_trace.nsys-rep 1.8 MB vLLM profile_dogacel.py ❌ 无
trace.nsys-rep(DeepSpec) 3.5 MB DeepSpec evaluator ✅ 有

10.2 CUDA Kernel 缺失原因

vLLM 的 EngineCore 在子进程中运行,nsys 默认只 trace 主进程。三个 vLLM trace 均无 GPU kernel 数据。

解决方案:重新采集时添加 --trace-fork 参数:

1
2
3
4
5
6
nsys profile -t cuda,nvtx,osrt,cudnn,cublas \
--python-backtrace=cuda --cudabacktrace=all \
--trace-fork \
--force-overwrite=true \
-o baseline_trace_v2 \
bash -c '...'

10.3 DeepSpec Evaluator Kernel 分布

Kernel 耗时占比 Instances 说明
CUTLASS GEMM (16×16) 66.3% 7,177 主要 matmul(Q/K/V/O + MLP)
elementwise_kernel 3.9% 8,522 RoPE、残差等
reduce_kernel (mean) 2.8% 4,298 RMSNorm
CUTLASS GEMM (32×32) 2.7% 382 大块矩阵乘法
Flash Attention 1.6% 864 Attention 计算
Softmax forward 1.0% 168 Softmax

关键观察:GEMM 占 66.3%,但 bs=1 decode 时本质是 memory-bound(M=1 瘦矩阵乘)。kernel launch 开销显著(~20000 次 launch)。vLLM 的 CUDA Graph 会消除大部分 launch 开销,fused kernel 会压缩 elementwise/reduce 占比。预期 vLLM 路径下 GEMM 占比升至 80%+。

10.4 NVTX Range 对比

NVTX Range Baseline Dogacel deepseek-ai dspark
dspark_propose
target_verify
decode_sample
warmup
VLLM::EngineCore

Dogacel 缺少 dspark_propose/target_verify 说明其 draft forward 未走标准 dspark 代码路径


11. Dogacel 无效原因:架构不匹配

维度 deepseek-ai(有效) Dogacel(无效)
Draft 架构 Qwen3DSparkSt EAGLE3
与 vLLM dspark 实现兼容 ✅ 完全对齐 ❌ 不匹配
NVTX range 存在 ✅ propose + verify ❌ 无
Mean accept len 7.1 推测极低
端到端加速 1.50x 0.97x(负优化)

根因:Dogacel 用 EAGLE3 架构训练 draft,中间层 hidden state 接口与 dspark 实现不兼容。即使 draft 能加载运行,acceptance rate 极低,draft 开销 > SD 收益。即使模型本身学得不差,接口不对也白搭。


12. SPEED-Bench 数据集

12.1 整体结构

SPEED-Bench(SPEculative Evaluation Dataset)是 NVIDIA 出的投机解码评测基准。

Split 样本数 用途
qualitative 880(11 类×80) 测 SD 质量(acceptance rate)
throughput_1k/2k/8k/16k/32k 1536×5 测系统吞吐(高并发)

12.2 Qualitative Split(质量评测)

从 18 个公开数据源聚合,分成 11 个 category:Coding、Math、Humanities、STEM、Writing、Summarization、Roleplay、RAG、Multilingual、Reasoning、QA。每类 80 个样本,用 OpenAI text-embedding-3-small 做嵌入,greedy 选择 + swap 优化最大化语义多样性(平均 pairwise cosine similarity 从 SpecBench 的 0.22 降到 0.14)。

12.3 Throughput Split(吞吐评测)

固定输入长度桶(1K/2K/8K/16K/32K),每桶 1536 条(512×3),分 3 个难度类别:low_entropy(coding 类)、high_entropy(creative writing 类)、mixed_entropy。用 tiktoken 精确 pad/truncate,不用 random token(会扭曲 MoE routing 和 acceptance behavior)。

12.4 为什么选 coding 类做 benchmark

  1. Coding 是低熵任务–token 可预测性高,SD 的 acceptance rate 天然高,是 best-case 场景
  2. 语义多样性好–80 条 prompt 覆盖 Python(27)、C++(9)、Java(10)、Go(13)、JS(11)、Rust(3) 等,来自 LiveCodeBench、Code Contests、HumanEvalPack
  3. 固定输出长度--speed-bench-output-len 2048)–隔离 prefill 影响,纯测 decode
  4. 两种并发对比--max-concurrency 32(batched,模拟生产环境)vs --max-concurrency 1(单流,测纯 decode 延迟)
  5. --disable-shuffle 保证可复现,--temperature 1.0 高温采样更反映真实使用场景

13. 接受率与训练效果的关系

Acceptance rate 的天花板由 draft 训练质量决定,工程实现决定能打到多少天花板。

训练侧决定上限

  • Draft 的 hidden state 和 target 的中间层对齐越好,token 分布越接近,accept 越高
  • deepseek-ai 的 block7 专门按 dspark 接口训练,hidden state 严格对齐 Qwen3-8B 第 7 层,所以 mean_accept_len=7.1
  • Dogacel 用 EAGLE3 方式训练,hidden state 映射方式不同,接口不对

工程侧决定下限

  • vLLM dspark 的 propose → verify pipeline 是否正确对接 draft
  • KV cache 的 layout、position ID 对齐、temperature sampling 一致性
  • CUDA Graph 是否覆盖 draft forward(没覆盖的话 launch overhead 会吃掉 SD 收益)
维度 deepseek-ai(训练+工程都对) Dogacel(工程接口不对)
Draft 架构 Qwen3DSparkSt EAGLE3
Hidden state 接口 ✅ 正确对接 ❌ 不匹配
NVTX range ✅ propose + verify ❌ 无
Mean accept len 7.1 推测极低
端到端加速 1.50x 0.97x

一句话总结:训练决定 draft 能不能猜对,工程决定猜对的部分能不能高效用上。Dogacel 的情况是工程接口就不对,猜得再准也走不进去。


14. vLLM 推理引擎优化对 Kernel 分布的影响

无 vLLM 优化的 kernel 分布(DeepSpec evaluator)

Kernel 占比 说明
CUTLASS GEMM (16×16) 66.3% bs=1 时是 memory-bound
elementwise 3.9% RoPE、残差等,未融合
reduce (mean) 2.8% RMSNorm,未融合
Flash Attention 1.6% decode 时计算量小
Softmax 1.0% 未融合
总 kernel launch ~20000 次 launch 开销显著

vLLM 优化后的预期变化

  1. CUDA Graph:20000 次 kernel launch → 1 次 graph launch
  2. Fused kernel:RMSNorm + residual + RoPE 融合为 1 个 kernel
  3. FlashInfer/FlashAttention decode-optimized:attention kernel 更高效
  4. GEMM 占比升至 80%+:其他开销被压缩后,GEMM 成为绝对瓶颈

对投机解码的启示

Baseline 的 decode 在 vLLM 下 GEMM 占 80%+,本质是 memory-bound(M=1 瘦矩阵乘,GPU 利用率低)。SD 的价值在于用 draft 的轻量 GEMM + target 的 batched verify 填充 GPU 空闲。当 batch size 增大(高并发),decode 从 memory-bound 转向 compute-bound,SD 收益下降–这也是 SPEED-Bench throughput split 存在的意义。


15. 总结

  1. DSpark = DFlash 并行 backbone forward + N 步 Markov 序列修正,全部在 CUDA Graph 内,用极小的开销把并行预测的「无依赖」缺陷补上
  2. 实测 deepseek-ai 官方 draft 在 Qwen3-8B 上实现 1.50x 加速mean_accept_len=7.1(N=8 几乎全接受)
  3. Dogacel 社区 draft 因架构不匹配(EAGLE3 vs DSpark)完全无效,0.97x 负优化
  4. 接受率天花板由训练决定,工程决定下限–hidden state 接口对齐是前提
  5. SPEED-Bench coding 类是 SD 的 best-case 场景,低熵任务下 acceptance rate 天然高

下一步

  • 重新采集 vLLM trace(加 --trace-fork),获取真实 kernel 数据
  • 检查 vLLM v0.26.0 是否支持 EAGLE3 method(Dogacel 应走 EAGLE3 而非 dspark)
  • 用 throughput split + --max-concurrency 32 测高并发下 SD 效果
  • 尝试 ngram、medusa 等 baseline 对比

参考

  • vLLM 源码:vllm/v1/worker/gpu/spec_decode/dspark/speculator.pyvllm/model_executor/models/qwen3_dspark.py
  • vLLM PR:#50138#50694#50737
  • 模型:deepseek-ai/dspark_qwen3_8b_block7Dogacel/Qwen3-8B-DSpark
  • 数据集:nvidia/SPEED-Bench,arXiv: 2604.09557

Speculative decoding 的 drafter 架构正在经历一次范式转移。DFlash 用 block diffusion 把 drafting 从串行变并行,实现了 6× 加速;DSpark 在此基础上补了两刀——半自回归解决并行生成的后缀衰减,置信度调度解决高并发下的验证浪费。本文围绕这两篇论文,结合源码逐行分析,澄清训练注意力结构中的常见困惑,并讨论其架构设计、核心 trade-off 和工程落地。


一、背景:从串行到并行的 Drafter

Speculative decoding 的加速比为 η=Ltarget/L\eta = L_{\text{target}} / L,其中每个 cycle 的 per-token 延迟为 L=(Tdraft+Tverify)/τL = (T_{\text{draft}} + T_{\text{verify}}) / \tauτ\tau 是每个 cycle 期望接受的 token 数。

Autoregressive drafter 和 diffusion drafter 的最大区别在于 drafting 的计算方式。自回归 drafter 一个 token 一个 token 地算,Tdraft=γtstepT_{\text{draft}} = \gamma \cdot t_{\text{step}} 与 block size 线性增长。为了控制延迟,只能用极浅架构(Eagle3 仅 1 层 transformer),τ\tau 很快饱和,加速比卡在 23×\sim 2{-}3\times。Diffusion drafter 一次并行算出整个 block 的 token,TdraftT_{\text{draft}}γ\gamma 基本不敏感,因此可以用更深的网络获得更高的 τ\tau

DFlash 就是这样一个并行 diffusion drafter。它的关键设计是 KV injection:从目标模型提取 hidden context features,注入到 draft 模型每一层的 Key-Value cache 中,让 draft 模型利用目标模型的深度表征来做条件预测,而不是从头猜。但纯并行生成引入了新问题:block 内 token 之间没有依赖建模


二、DFlash:用 Diffusion 做 Drafter

2.1 核心思路

DFlash 的核心 insight 很简单:目标模型知道未来

大型自回归模型的 hidden states 隐含了多个未来 token 的信息。DFlash 不让小模型从头猜,而是把目标模型的 hidden features 作为条件,让 draft 模型变成一个"扩散适配器"——利用目标模型的深度表征来并行预测未来 block。

2.2 “Diffusion” 到底在哪?

DFlash 名字里有个 D,但翻遍代码你会发现一个事实:没有多步去噪,没有噪声调度,没有连续时间 SDE。所谓的 diffusion 只体现在两件事上:

  1. Mask token 构造:待预测位置填充为 mask token,类似于 BERT 的 [MASK],作为"全噪"起点
  2. 双向注意力is_causal=False):block 内 token 互相可见,一次 forward 出所有位置

就这两点,没有迭代去噪。标准 block diffusion(Arriola et al., 2025)还有多步迭代,DFlash 把它压成了单步。传承链条是这样的:

1
2
3
连续扩散 (LLaDA)  ->  Block 级离散扩散 (BlockDiff)  ->  单步 mask-predict (DFlash)
高斯噪声 多步迭代去噪 一步出结果
连续时间 SDE 离散 mask token BERT-style

每一步都在往"更像自回归、更不像 diffusion"的方向走。DFlash 到了极致——名字叫 diffusion,实质是 parallel mask prediction。论文用 “diffusion” 这个词主要是学术传承定位,不是方法描述。

那 mask token 的作用是什么?模型需要知道哪些位置是待预测的,哪些是已知信息。如果不用 mask token,直接放随机 embedding 进去,模型会把这些位置当成已知输入去做 attention。Mask token 是一个学习到的"我不知道"信号,跟 BERT 的 [MASK] 一回事。

2.3 KV Injection:不是输入融合,是每层注入

Eagle3 也用目标模型的 hidden features,但只在输入层融合,随着 draft 模型变深,目标信息逐渐稀释。DFlash 采用了完全不同的策略:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
目标模型 hidden states (5层)

│ concat + 线性投影 W_c (5*hidden -> hidden)

H_ctx = RMSNorm(W_c [H(l1); ...; H(l5)]) ← 压缩后的上下文特征

│ 注入到 draft 模型每一层的 KV cache

Draft Layer 1: K = [W^K · H_ctx; W^K · H_d] ← 目标特征 + draft 特征拼接
V = [W^V · H_ctx; W^V · H_d]

Draft Layer 2: 同上(H_ctx 共享)

...

Draft Layer 5: 同上

关键区别:目标特征作为额外的 KV entry 直接注入每一层,而不是经过 draft 模型的 Q projection、output projection 和 FFN。这意味着目标信息在每一层都是"常驻"的,不会因深度而稀释。

从源码看(dflash/model.py),KV 注入的实现极其直接:

1
2
3
4
5
6
# 每个 draft decoder 层的 attention 中
k_ctx = self.k_proj(target_hidden) # KV 注入:target context
k_noise = self.k_proj(hidden_states) # draft 自身的 K/V
k = torch.cat([k_ctx, k_noise], dim=1) # 拼接 K
v = torch.cat([v_ctx, v_noise], dim=1) # 拼接 V
# is_causal = False ← 双向注意力,block 内 token 互相可见

所有层共享同一份 target_hidden(经 fc + RMSNorm 投影后),且 attention 设为 is_causal=False——block 内 token 双向可见,这是并行扩散生成的必要条件。

一句话总结:KV injection 把目标模型的 hidden features 变成 draft 模型每一层的"持久上下文",让深层 draft 模型也能充分利用目标模型的信息。

2.4 共享 Embedding 和 LM Head:设计意图

在深入推理和训练流程之前,需要先澄清一个贯穿两篇论文的基础设计:共享 embedding 和 LM head

DFlash 论文明确阐述了这一设计的动机:

“To improve training efficiency, the draft model shares the token embedding layer and language modeling head with the target model and keeps them frozen during training. Only the draft Transformer layers are updated. This design reduces the number of trainable parameters and encourages the draft model to function as a lightweight diffusion adapter tightly aligned with the target model’s representation space.”

这里的"target model"就是你想加速的那个大模型——最终产出正确 token 的 LLM(如 Qwen3-4B、DeepSeek-V4)。target.lm_head 不是什么特殊构造,它就是目标模型自带的最后一个线性层——把 hidden state 映射到词汇表 logits 的那一层。

以 Qwen3-4B 为例:

  • embed_tokensvocab_size(151936) × hidden_size(2560) ≈ 390M 参数
  • lm_headhidden_size(2560) × vocab_size(151936) ≈ 390M 参数
  • 两者合计占 Qwen3-4B 总参数(4B)的约 20%

关键词是 “lightweight diffusion adapter”——共享 + 冻结 embed/lm_head 的本质目的不是省参数,而是强制 draft 模型在目标模型的表征空间内工作embed_tokens 决定输入空间,lm_head 决定输出空间,两者都锁定后,draft 模型只能学习"如何把目标模型的 hidden states 转换成未来 token 的预测",而不能自己学一套独立的表征。这正是 KV injection 设计的配套——KV injection 让目标模型的信息每层注入,共享 embed/lm_head 让 draft 的输入输出空间与目标对齐,两者合在一起确保 draft 是一个纯粹的"适配器"而非独立模型。

两份源码的共享方式不同

  • DFlash(推理时直接借用)DFlashDraftModel 类本身不持有 embed_tokenslm_head 模块。在 dflash_generate() 函数中直接调用 target 对象的属性:

    1
    2
    3
    # model.py 第111-112行 - 直接调用,不存副本
    noise_embedding = target.model.embed_tokens(block_output_ids)
    draft_logits = target.lm_head(model(...))
  • DSpark(训练时复制 + 冻结):draft 模型有自己的 embed_tokenslm_head 模块(modeling.py 第 227-246 行定义),初始化时把 target 的权重逐字节复制过来然后冻结:

    1
    2
    3
    4
    5
    6
    def initialize_embeddings_and_head(self, *, embed_tokens, lm_head, freeze=True):
    with torch.no_grad():
    self.embed_tokens.weight.copy_(embed_tokens.weight.detach())
    self.lm_head.weight.copy_(lm_head.weight.detach())
    if freeze:
    self.set_embedding_head_trainable(False) # requires_grad=False

    训练时必须用独立模块供 PyTorch autograd 走完整前向传播;推理时则像 DFlash 一样直接调用 target 的 lm_head。

2.5 推理流程:极简实现

DFlash 的仓库极其精简(4 个 Python 文件,核心逻辑 ~370 行)。推理主循环 dflash_generate() 的核心步骤:

1
2
3
4
5
6
7
8
9
while not done:
① 构造 [prev_token, mask, mask, ..., mask] block
② Draft 前向:单次并行生成整个 block 的 logits
- noise_embedding = target.model.embed_tokens(block_output_ids) # 直接用 target 的 embed
- draft_logits = target.lm_head(draft_model(...)) # 直接用 target 的 lm_head
③ 采样 draft tokens
④ 目标模型单次 forward 验证整个 block
⑤ 计算 accept_length(cumprod 找到第一个 reject 的位置)
⑥ 裁剪 draft 和 target 的 KV cache 到接受位置

注意:draft 模型在推理时直接使用目标模型的 embed_tokenslm_head(通过传入的 target 对象直接访问),自己只持有 5 个 decoder 层。block_size=1 时退化为普通自回归解码(用于 baseline 对比)。

重要说明:DFlash 仓库只包含推理代码,训练 recipe 尚未开源。DSpark 在 DFlash 架构基础上增加了独立的训练 pipeline,完整实现在 DeepSpec 仓库中。两者不是共享同一套训练框架——DSpark 的训练代码是独立开发的,包含了 Markov head、confidence head、anchor sampling 等 DSpark 特有组件。

2.6 并行扩散 drafting

DFlash 用 block diffusion 一次生成 γ\gamma 个 token:

TdraftDFlash=tparallel(与 γ 无关)T_{\text{draft}}^{\text{DFlash}} = t_{\text{parallel}} \quad (\text{与 } \gamma \text{ 无关})

这意味着 draft 模型可以用更深的架构(5 层 vs Eagle3 的 1 层),而不会让 drafting 延迟失控。实验显示,5 层 DFlash 生成 16 个 token 的延迟,低于 1 层 Eagle3 生成 8 个 token 的延迟。

需要注意的是,历史 context 的 KV cache 仍然是 causal 的。目标模型的 forward pass 是标准 causal attention,产出的 hidden states 已经编码了"只能看前面"的因果历史。DFlash 通过 KV injection 把这些 hidden states 注入到 draft 模型,所以 draft block 整体的 attention pattern 是:

1
2
3
4
5
[历史 context(causal,来自目标模型 KV 注入)]  [draft block(bidirectional)]

mask tokens 互相可见
但都只能看到历史 context
看不到"未来"

2.7 训练设计

DFlash 的训练 recipe 未开源。以下分析基于 DSpark 论文和 DeepSpec 仓库源码,两者的训练设计在 backbone 层面一致(KV injection、共享 embed/lm_head、anchor sampling 等核心机制相同),DSpark 额外增加了 Markov head 和 confidence head 的训练。

2.7.1 序列布局:拼接式而非交错式

训练时的输入序列布局是 concatenated(拼接式),不是 interleaved(交错式):

1
2
[ context: p1 p2 p3 r1 r2 r3 r4 r5 ]  [ draft: B0 B1 B2 ... ]
← 拼接在 context 之后
  • Context 部分:完整的训练样本 [prompt | response],全部是 ground truth token。目标模型对这段序列做一次 forward,提取 5 个中间层的 hidden states 作为 KV injection 来源。
  • Draft 部分:512 个 block 拼接而成,每个 block 是 [anchor, mask, mask, ..., mask],block_size=7。

2.7.2 随机 Anchor 采样

不从 response 均匀分块,而是随机采样 anchor 位置作为每个 block 的起点。源码(common.py 第 164 行)显示采样后会 .sort() 排序:

1
anchors = gathered[:, :max_n].sort(dim=1).values  # 随机采样后排序

排序后 block 按位置从小到大排列,每个 block 的 anchor 是一个真实的 ground truth token(teacher-forced),紧跟 block_size - 1 个 mask token。一次训练前向传播同时覆盖 512 个位置。

2.7.3 Flex Attention:用函数描述稀疏注意力模式

DSpark 的注意力模式高度稀疏——每个 7-token block 只能看到 anchor 之前的 context + 自己 block 的 7 个 token。用稠密矩阵(Q_LEN × KV_LEN 布尔矩阵)太浪费。

PyTorch 2.5+ 的 flex_attention API 提供了解法:用函数描述注意力模式,而不是构造稠密矩阵。

流程

  1. 提供一个 mask_mod(b, h, q_idx, kv_idx) -> bool 函数,告诉它"query 位置 q 能不能看到 key 位置 k"
  2. create_block_mask() 把这个函数编译成块稀疏格式——把整个矩阵切成小块(比如 128×128),只保留含 True 的块
  3. 实际 attention 计算时跳过全 False 的块,只算有内容的块

源码核心(common.py 第 86-96 行):

1
2
3
4
5
6
7
8
9
10
11
12
13
def dspark_mask_mod(b, h, q_idx, kv_idx):
q_block_id = q_idx // block_size
anchor_pos = anchor_positions[b, q_block_id]

is_context = kv_idx < seq_len
mask_context = is_context & (kv_idx < anchor_pos) # 只看 anchor 之前的 context

is_draft = kv_idx >= seq_len
kv_block_id = (kv_idx - seq_len) // block_size
mask_draft = is_draft & (q_block_id == kv_block_id) # 只看同一个 block

is_valid_block = block_keep_mask[b, q_block_id]
return (mask_context | mask_draft) & is_valid_block

2.7.4 Attention Mask 的两条规则

整个注意力可见性只有两条规则:

  1. Block 之间互不可见q_block_id == kv_block_id)——不同 block 的 draft token 完全隔离,双向的
  2. Anchor 之前的前缀可见kv_idx < anchor_pos)——context 中 anchor 位置之前的 token 可见

这两条规则产生了注意力矩阵中的阶梯(staircase)结构。

2.7.5 Invisible Tokens:到底是什么

论文训练图中的"白色 = invisible tokens"让人困惑。Invisible tokens 分两类:

Invisible 类型 条件 原因
Context 中 anchor 位置及之后的 token kv_idx >= anchor_poskv_idx < seq_len 因果一致性:这些是 draft 要预测的答案,看了就是 data leakage
其他 block 的 draft token q_block_id != kv_block_id 块间隔离:防止不同 block 之间的梯度互相干扰

关键澄清:Context 边界是 token 级别的,不是 block 级别的

这是理解训练图最容易混淆的地方。看注意力矩阵:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
         Context (causal)                    Draft blocks (bidirectional)
c0 c1 c2 c3 c4 c5 c6 c7 B0(a m m) B1(a m m) B2(a m m)
B0 ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗ ✓ ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗
B0 ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗ ✓ ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗
B0 ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗ ✓ ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗

B1 ✓ ✓ ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗ ✗ ✓ ✓ ✓ ✗ ✗ ✗
B1 ✓ ✓ ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗ ✗ ✓ ✓ ✓ ✗ ✗ ✗
B1 ✓ ✓ ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗ ✗ ✓ ✓ ✓ ✗ ✗ ✗

B2 ✓ ✓ ✓ ✓ ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗ ✗ ✗ ✓ ✓ ✓
B2 ✓ ✓ ✓ ✓ ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗ ✗ ✗ ✓ ✓ ✓
B2 ✓ ✓ ✓ ✓ ✓ ✓ ✗ ✗ ✗ ✗ ✗ ✗ ✗ ✗ ✓ ✓ ✓
↑ ↑ ↑
anchor=2 anchor=4 anchor=6
(B0的边界) (B1的边界) (B2的边界)
  • B0(anchor=2):只看到 context [c0, c1]——2 个 token
  • B1(anchor=4):看到 context [c0, c1, c2, c3]——4 个 token(包含 B0 看到的 + 更多)
  • B2(anchor=6):看到 context [c0, c1, c2, c3, c4, c5]——6 个 token

蓝色(可见 context)形成一个阶梯。阶梯的每一级台阶在 anchor 位置(2、4、6),是单个 token 的位置。

为什么 context 看起来也按 block 切了?

这是视觉错觉。两个相邻 anchor 之间的 context 段(比如 [c2, c3])对 B0 不可见、对 B1 和 B2 可见,在图里看起来像一个"块"。但边界是随机 anchor 的 token 位置,不是固定的 block 边界。如果 anchor 随机采到位置 1、4、9,分段就完全不同。

为什么 anchor 之后的不看?

训练时虽然完整序列都在手里,但必须用 mask 模拟推理条件。推理时 draft 模型只能看到 anchor 之前的 token(因为后面的还没生成),所以训练时也必须只让它看 [0, anchor_pos)。这和标准自回归训练的 causal mask 完全同理——你有完整序列,但人为限制可见性防止作弊,只是这里"未来"的定义从"当前位置之后"变成了"anchor 位置之后"。

每个 block 内所有 token 共享同一个 anchor_pos,所以它们看到的 context 前缀完全一样。在注意力矩阵里,这表现为同一 block 的所有行在 context 区域的可见性模式完全一致——画出来就是一个矩形块,视觉上像是 context 也按 block 对齐了。但决定可见/不可见边界的是 anchor_pos 这一个整数,是 token 级别的。

2.7.6 KV Injection 在训练中的结构

训练时的 KV injection 和推理时完全一致——目标模型的 hidden states 经过 fc(5×hidden → hidden) + RMSNorm 投影后,作为额外的 K/V entry 注入到 draft 模型每一层:

1
2
3
每层 attention 的 K/V 拼接:
K = [k_proj(target_hidden) ; k_proj(draft_hidden)] ← 两部分拼接
V = [v_proj(target_hidden) ; v_proj(draft_hidden)]

目标特征绕过 Q projection、output projection、FFN,直接作为 KV entry 进入 attention。所有层共享同一份投影后的 target_hidden。

这里的"KV"不是推理时增量生成的 KV cache,而是指 KV injection 的结构——目标模型的 hidden states 作为"常驻 KV"注入每一层。

2.7.7 指数衰减位置加权

wk=exp(k1γ)w_k = \exp\left(-\frac{k-1}{\gamma}\right)

Speculative decoding 中,早期 token 的错误会级联失效整个 block 后缀。loss 加权反映了这种不对称性——前面的 token 更重要。

2.7.8 训练 vs 推理

维度 训练 推理
Anchor Ground truth token(teacher-forced) 目标模型上一步的 bonus token
Block 数量 512 个 block 一次 forward 一次一个 block
KV injection 与推理一致(每层注入) 同左
Attention mask Flex attention block mask 标准 bidirectional
串行 head(DSpark) Teacher-forced,所有位置并行计算 Autoregressive,逐 token 串行

2.8 结果与局限

DFlash 在 Qwen3-8B 上实现 6.1× 加速,比 Eagle3 快 2.5×。但存在两个结构性局限:

  1. 后缀衰减:纯并行生成无法建模 block 内依赖。当上下文有多个合理续写(如 “of course” vs “no problem”)时,各位置独立预测可能产生不一致的组合(“of problem”)
  2. 验证浪费:所有 draft token 都送去验证,高并发场景下低置信度的后缀 token 占用 batch 容量

DSpark 正是来解决这两个问题的。


三、DSpark:半自回归 + 置信度调度

3.1 整体架构

DSpark = DFlash backbone + 轻量串行 head + 置信度 head + 硬件感知调度器

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
目标模型生成 bonus token (anchor)


┌─────────────────────────────────────┐
│ 并行 backbone (DFlash) │
│ 输入: anchor + (γ-1) mask tokens │
│ 输出: hidden h_1..h_γ, base logits │
└──────────┬──────────────────────────┘

┌──────┴──────┐
│ │
▼ ▼
┌────────┐ ┌──────────────┐
│串行 head│ │置信度 head │
│B_k(·) │ │c_k = σ(w·h) │
└───┬────┘ └──────┬───────┘
│ │
▼ ▼
采样 x_k prefix survival
(条件于 概率估计
x_<k)
│ │
▼ ▼
draft tokens ┌──────────────┐
E F G H │硬件感知调度器 │
│ 截断低置信后缀 │
└──────┬───────┘


目标模型验证
E F G (H 被砍掉)

3.2 半自回归生成:解决后缀衰减

问题本质

并行 drafter 在每个位置独立预测,相当于对前缀所有可能的 token 做 marginal 平均。当上下文存在多个合理续写路径时,不同位置可能选到不同路径的 token,产生不连贯的组合。

以论文中的例子:上下文允许 “of course” 和 “no problem” 两种续写。并行 drafter 在位置 1 独立采样得到 “of”,位置 2 仍然不知道位置 1 选了什么,可能选 “problem” 而非 “course”。这就是多模态碰撞(multi-modal collision)

解法:并行 backbone + 串行 head

DSpark 把生成拆成两个阶段:

并行阶段:DFlash backbone 一次 forward 产生所有位置的 hidden states h1,,hγh_1, \ldots, h_\gamma 和 base logits U1,,UγU_1, \ldots, U_\gamma

串行阶段:在 base logits 上叠加一个 transition bias BkB_k,逐 token 左到右采样:

pk()=softmax(Uk+Bk(x0,x<k))p_k(\cdot) = \text{softmax}(U_k + B_k(x_0, x_{<k}))

关键在于 BkB_k 条件于前面已采样的 token,解决了独立预测的问题。一旦位置 1 采样了 “of”,串行 head 在位置 2 boost “course” 并 suppress “problem”。

三种串行 Head 实现

源码在 markov_head.py 中实现了三种变体,复杂度递增:

类型 参数 机制 自回归程度
VanillaMarkov markov_w1(Embed) + markov_w2(Linear) 查表 + 低秩矩阵乘 最轻,仅依赖前一 token
GatedMarkovHead + gate_proj(Linear) 门控混合 draft hidden 与 embedding 中等,依赖前一 token + hidden
RNNHead + joint_proj(Linear) GRU-like state 跨位置传播 最强,维护整个前缀历史

Markov head(默认):BkB_k 只依赖前一个 token,低秩分解 B=W1W2B = W_1 W_2W1RV×rW_1 \in \mathbb{R}^{V \times r}W2Rr×VW_2 \in \mathbb{R}^{r \times V}r=256r=256

为什么用低秩分解? 一阶马尔可夫转移矩阵 BRV×VB \in \mathbb{R}^{V \times V} 直接存储需要 V2V^2 参数(V=151936 时约 230 亿),完全不可行。低秩分解 B=W1W2B = W_1 W_2 压缩到 2×V×r=2×151936×25678M2 \times V \times r = 2 \times 151936 \times 256 \approx 78M 参数,减少 99.7%。

计算过程:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
输入: 前一个 token id x_{k-1} (整数, 如 342)

Step 1: Embedding 查表
W_1[x_{k-1}] → 取 W_1 的第 x_{k-1} 行, 得到 r 维向量 ∈ ℝ^r
等价 PyTorch: w1_row = W1[x_prev] # shape: (r,)

Step 2: 低秩矩阵乘
B(x_{k-1}, ·) = W_1[x_{k-1}] · W_2 ∈ ℝ^V
等价 PyTorch: bias = w1_row @ W2 # shape: (V,)

Step 3: 叠加到 base logits 并采样
p_k(·) = softmax(U_k + B(x_{k-1}, ·))
x_k ~ p_k

直观理解:
W_1: 马尔可夫嵌入表 (token id → r 维隐向量)
W_2: 隐向量 → 全词表 logit 空间的投影
组合效果: 给定前一个 token, 对所有 V 个候选 token 打出转移偏置分数

RNN head 比 Markov head 略好但实现更复杂,收益有限(论文 Figure 4 显示差距很小),生产默认用 Markov head。

训练 vs 推理:一个关键区别

训练时是 teacher-forced,可以并行计算。用 ground truth token ids 作为前缀输入,所有位置的 bias 一次算完,不需要串行循环:

1
2
3
4
5
训练(teacher-forced,并行):
位置1: bias = W_1[anchor] · W_2 ← 用 ground truth 的 anchor
位置2: bias = W_1[gt_token_1] · W_2 ← 用 ground truth 的 token_1
位置3: bias = W_1[gt_token_2] · W_2 ← 用 ground truth 的 token_2
所有位置一次 forward 算完

推理时必须逐 token 串行,因为位置 k 的 bias 依赖位置 k-1 实际采样出来的 token,不是 ground truth:

1
2
3
4
5
6
推理(autoregressive,串行):
位置1: bias = W_1[anchor] · W_2
采样 x_1 ~ softmax(U_1 + bias) ← 这步必须先完成
位置2: bias = W_1[x_1] · W_2 ← 用上一步采样的 x_1
采样 x_2 ~ softmax(U_2 + bias)
...

因为每步只是一个 embedding 查表 + 低秩矩阵乘,非常轻量,在 batch size 128 的生产环境下延迟开销只有 ~1%。这也是 DSpark 叫"半自回归"的原因——backbone 是并行的,串行 head 是自回归的。

一句话总结:半自回归 = 并行 backbone 出 base logits + 轻量串行 head 加条件 bias,用 1% 的延迟代价换来 16-18% 的接受长度提升。

3.3 位置分析:为什么并行能赢自回归?

DSpark 论文里最精彩的实验是位置层面的条件接受率分析(Figure 2)。这个分析揭示了一个反直觉的结论。

反直觉:并行 drafter 在位置 1 显著优于自回归 drafter

原因纯粹是架构容量:自回归 drafter 受限于 TdraftγT_{\text{draft}} \propto \gamma,只能用 1 层 transformer;并行 drafter 延迟与 γ\gamma 无关,可以用 5 层甚至 8 层。更深的网络在位置 1(只依赖目标模型上下文,不需要 block 内依赖)自然更准。

论文 Figure 2 按不同 domain(Math、Code、Chat)分别给出了位置接受率曲线。大致趋势是:

1
2
3
4
位置:    1     2     3     4     5     6     7
DFlash: 高位 ↘ ↘ ↘ ↘ ↘ ↘ ← 缓慢但持续衰减
Eagle3: 低位 ↗ ↗ ↗ ↗ ↗ ↗ ← 反而上升(利用已采样的前缀)
DSpark: 最高 ↘ ↘ ↘ ↘ ↘ ↘ ← 高位稳定衰减

Eagle3 在后续位置反超 DFlash,因为自回归模型条件于已采样 token,随着前缀确定,后续预测变得更容易。DFlash 在后续位置仍然是独立预测,持续衰减。

DSpark 两头兼得:位置 1 继承并行 backbone 的高容量优势,后续位置由串行 head 维持稳定。

注意:以上趋势基于论文 Figure 2 的定性描述。具体数值请参阅论文原图,不同 domain(Math/Code/Chat)的绝对值差异较大。

为什么位置 1 最重要?

Speculative decoding 是严格的 prefix-matching:位置 1 被 reject,整个 block 作废。位置 1 的接受率对整体 τ\tau 的杠杆最大。

这个分析也解释了为什么 DFlash 能比 Eagle3 快 2.5×——不是因为并行生成更快(虽然确实更快),而是因为更深的网络在位置 1 的优势被 prefix-matching 机制放大了。

3.4 置信度调度验证:从固定长度到自适应

问题:高并发下的验证浪费

DFlash 和 Eagle3 都用固定长度验证:draft 模型生成 γ\gamma 个 token,全部送去目标模型验证。但在高并发场景下:

  • 每个 extra verification token 都占用目标模型的 batch 容量
  • 低置信度的后缀 token 大概率被 reject,验证它们是纯浪费
  • 被浪费的 batch 容量本可以服务其他请求

解法:Confidence Head + Hardware-Aware Scheduler

Confidence Head 的源码实现极其极简——就是一个单层线性投影:

1
2
3
4
5
6
# eval/dspark/confidence_head.py
class AcceptRatePredictor(nn.Module):
def __init__(self, input_dim: int):
self.proj = nn.Linear(input_dim, 1) # 单层线性投影
def forward(self, features):
return self.proj(features).squeeze(-1)

输入特征是 [hidden_states, markov_prev_embeddings] 拼接,输出经过 sigmoid 后得到每个位置的条件生存概率:

ck=σ(w[hk;W1[xk1]])c_k = \sigma(w^\top [h_k; W_1[x_{k-1}]])

监督信号是解析的 per-step 接受率:ck=112pkdpkt1c_k^* = 1 - \frac{1}{2}\|p_k^d - p_k^t\|_1(TV distance 的补)。训练时用 BCE loss。

推理时的置信度裁剪同样简洁:

1
2
3
4
# draft_ops.py - 找到第一个低于阈值的置信度位置,截断
below_threshold = confidence_logits.sigmoid() < threshold
first_below = torch.nonzero(below_threshold[0])[0].item()
return first_below # 只验证 [0, first_below) 的 token

Sequential Temperature Scaling (STS) 校准:原始 confidence 通常过自信(ECE 3-8%)。STS 逐位置做 1D grid search,最小化累积乘积 ikci\prod_{i \leq k} c_i 的 ECE,校准后 ECE 降到 ~1%。

Hardware-Aware Prefix Scheduler 把验证长度选择形式化为全局吞吐量最大化问题:

Θ=τSPS(B),其中 τ=r=1R(1+j=1rar,j),B=r=1R(1+r)\Theta = \tau \cdot \text{SPS}(B), \quad \text{其中 } \tau = \sum_{r=1}^{R}\left(1 + \sum_{j=1}^{\ell_r} a_{r,j}\right), \quad B = \sum_{r=1}^{R}(1 + \ell_r)

  • SPS(B)\text{SPS}(B):引擎的 steps-per-second 容量曲线,初始化时 profiling 一次
  • ar,j=ijcr,ia_{r,j} = \prod_{i \leq j} c_{r,i}:request rr 在位置 jj 的 prefix survival 概率
  • 目标:选择每个 request 的验证长度 1,,R\ell_1, \ldots, \ell_R,最大化 Θ\Theta

因为 ar,ja_{r,j} 单调递减,可以贪心求解:全局排序所有 (r,j)(r, j)ar,ja_{r,j} 降序,逐个加入验证 batch,直到 Θ\Theta 不再上升。

1
2
负载低 -> SPS(B) 几乎不变 -> 多验证 token 划算 -> 验证长度大
负载高 -> SPS(B) 快速下降 -> 少验证 token 划算 -> 砍掉低置信度后缀

Verification Length vs System Load 的 tradeoff 直观对比

低负载 高负载
SPS(B) 曲线 平缓,几乎不降 陡降,batch 增大代价高
最优验证长度 长(接近 block_size) 短(只留高置信前缀)
砍掉的 token 几乎不砍 砍掉大量低置信后缀
单次验证收益 τ\tau 大,SPS 不受影响 τ\tau 小,但 SPS 保住不崩
系统级效果 每个请求多验证,τ\tau 最大化 腾出 batch 容量服务更多请求
极端情况 无损,接近离线吞吐 baseline 崩溃,DSpark 仍稳定

核心矛盾:验证更多 token 能提高单个请求的 τ\tau,但增大 batch BB 会拉低所有人的 SPS。调度器在两者之间找全局最优–低负载时 SPS “便宜”,多验证划算;高负载时 SPS “贵”,省着用 batch 容量更重要。这和日常生活中的资源分配同理:空闲时精雕细琢,繁忙时抓大放小。

一句话总结:置信度调度把"验证多少"从静态配置变成动态优化问题–根据每个请求的 draft 质量和当前系统负载,全局分配验证算力。

3.5 训练目标

DSpark 的 loss 三项加权和:

L=αceLce+αtvLtv+αconfLconf\mathcal{L} = \alpha_{\text{ce}} \mathcal{L}_{\text{ce}} + \alpha_{\text{tv}} \mathcal{L}_{\text{tv}} + \alpha_{\text{conf}} \mathcal{L}_{\text{conf}}

Loss 项 作用 权重
Lce\mathcal{L}_{\text{ce}} 交叉熵,预测正确 token 0.1
Ltv\mathcal{L}_{\text{tv}} TV distance,匹配目标分布 0.9
Lconf\mathcal{L}_{\text{conf}} BCE,校准置信度预测 1.0

Ltv\mathcal{L}_{\text{tv}} 权重最高,因为 TV distance 直接对应接受率:per-step 接受概率 =112pdpt1= 1 - \frac{1}{2}\|p^d - p^t\|_1,最小化 TV distance 就是最大化期望接受率。


四、工程落地:从论文到 DeepSeek-V4 线上

4.1 生产部署架构

DSpark 部署在 DeepSeek-V4-Flash 和 V4-Pro 上。

配置项 DeepSeek-V4 生产环境 开源 checkpoint(如 dspark_qwen3_4b_block7)
Draft backbone 3 层 MoE + mHC + sliding window attention 128 标准 dense transformer 层
Block size 5 7
串行 head Markov head (r=256) 同左
置信度 head 线性投影 + sigmoid 同左
校准 STS (held-out validation set) 同左
调度器 异步硬件感知 prefix scheduler 仅 Transformers 评估器

注意:生产环境使用 MoE + mHC 架构和 block_size=5;开源 checkpoint 使用标准 dense 层和 block_size=7,便于社区复现。两者核心算法一致,架构配置不同。

4.2 异步调度:解决 ZOS 冲突

算法 1 的同步版本与生产系统的 Zero-Overhead Scheduling (ZOS) 冲突——ZOS 需要在当前 step 完成前知道下一步的 batch size。DSpark 的解法是用两步前的 confidence 预测来确定当前步的截断长度

1
2
3
4
Step N-2:  生成 draft + confidence
Step N-1: 用 N-2 的 confidence 确定截断 → 验证
Step N: 用 N-1 的 confidence 确定截断 → 验证
↑ 同时生成新 draft + confidence(供 N+1 使用)

这引入了轻微的时间偏差,但选择机制是 rank-preserving 的——最自信的 draft token 总是优先验证。更重要的是,异步设计形成了一道"因果屏障":截断决策只依赖历史信息,不会泄露未来 token,保证了 lossless guarantee。

4.3 生产性能

指标 V4-Flash V4-Pro
每用户速度提升(matched throughput) 60%–85% 57%–78%
吞吐提升(moderate SLA) +51% +52%
极端 SLA 下吞吐优势 +661%(baseline 接近崩溃) +406%

关键结论不是倍数本身,而是DSpark 扩展了可行的交互性边界。在 120 TPS/user 的严格 SLA 下,MTP-1 baseline 几乎无法运作,DSpark 仍然稳定——这意味着原来达不到的延迟等级现在可以服务了。


五、两篇论文的对照

维度 DFlash DSpark
Drafter 类型 纯并行 block diffusion 半自回归(并行 + 串行)
Block 内依赖 无建模 Markov/RNN head
验证策略 固定长度全验证 置信度 + 硬件感知自适应
位置 1 优势 深层网络 继承 DFlash backbone
后缀稳定性 快速衰减 串行 head 维持
高并发友好 验证浪费 动态截断
生产验证 SGLang 实验 DeepSeek-V4 线上流量
vs Eagle3 2.5× 更快 τ\tau 再 +16-18%
训练代码开源 未开源 完整 pipeline(DeepSpec)
推理后端 Transformers/SGLang/vLLM/MLX 仅 Transformers

DSpark 不是对 DFlash 的替代,而是增量改进。DFlash 解决了"能不能用 diffusion 做 drafter"的问题,DSpark 解决了"用得好不好"的问题。两者共享 KV injection 条件化、共享 embedding/LM head、位置加权 loss 等核心设计。


六、可迁移的启示

1. "目标模型知道未来"是一个深刻的观察。 大模型的 hidden states 隐含了远超 next-token 的信息。DFlash 的 KV injection 和 DSpark 的串行 head 都在利用这一点——draft 模型不需要从头推理,只需要"解读"目标模型已经知道的东西。

2. 并行 vs 自回归不是二选一。 DSpark 的半自回归架构证明,用并行 backbone 做"重活"+ 串行 head 做"精修",可以在 1% 延迟代价下获得 16-18% 的质量提升。这个思路在更广泛的 LLM 加速领域也适用——不要追求纯并行或纯串行,找正确的分割点。

3. Speculative decoding 是系统问题,不只是算法问题。 DSpark 的置信度调度把验证长度从算法参数变成系统调度参数,在真实流量下实现了负载感知的自适应。这提醒我们:脱离部署环境谈 drafter 架构是不完整的。

4. 位置 1 的杠杆最大。 在 prefix-matching 机制下,位置 1 的接受率对整体 τ\tau 的影响远大于后续位置。这意味着 draft 模型的架构选择应该优先考虑位置 1 的容量,而非后续位置的依赖建模——这正好是并行 drafter 的天然优势。

5. 共享 embed/lm_head 不是省参数,是锁定表征空间。 冻结 embed 和 lm_head 后,draft 模型被强制在目标模型的表征空间内工作,成为纯粹的"适配器"。这是 KV injection 的配套设计——前者保证空间对齐,后者保证信息每层注入。


参考

  • DFlash: Chen et al., “DFlash: Block Diffusion for Flash Speculative Decoding”, ICML 2026. arXiv:2602.06036
  • DSpark: Cheng et al., “DSpark: Confidence-Scheduled Speculative Decoding with Semi-Autoregressive Generation”, 2026. arXiv:2607.05147
  • Eagle3: Li et al., “Eagle-3: Scaling up Inference Acceleration of LLMs via Training-Time Test”, 2025. arXiv:2503.01840

开源仓库

  • DeepSpec(deepseek-ai/DeepSpec):69 个 Python 文件,包含 Eagle3、DFlash backbone、DSpark 三种 drafter 的统一训练框架,支持 Qwen3 和 Gemma4 系列目标模型。完整训练 pipeline(数据准备 → 训练 → 评估)。
  • DFlash(z-lab/dflash):4 个 Python 文件,仅含推理代码(Transformers/SGLang/vLLM/MLX 四种后端),训练 recipe 尚未开源。

MiMo-V2.5 推理优化解读:Hybrid SWA 的工程落地

小米 MiMo 团队的《Full-Pipeline Inference Optimization for MiMo-V2.5 Series》(arXiv:2607.13095)核心论点很直白:Hybrid SWA 理论上能把 KVCache 和 attention 计算量压到 Full Attention 的 1/7,但如果 KVCache 系统不跟着改造,实际反而会变慢

本文不写综述,重点拆解几个关键工程问题:SWA 的 KV cache 怎么存、怎么 offload、共享前缀命中时的重算问题、numa_balancing 为什么 +10%、以及 Prefill 为什么也要开 MTP


1. 论文速览

MiMo-V2.5-Pro 的架构组合拳:

1
2
Hybrid SWA + 稀疏 MoE + 多模态编码器
70 层 = 10 Full Attention + 60 SWA(窗口 W=128)

理论上 attention FLOPs 和 KVCache 存储都降到 1/7,但论文的核心观点是:这些理论红利不会自动兑现。Hybrid SWA 在 KVCache 管理、前缀匹配、Full 和 SWA 层语义一致性上都引入了新问题。

论文涉及的优化面很广:KVCache 分层(L1 GPU / L2 Host / L3 GCache)、调度(LLM-Router)、Prefill/Decode pipeline、多模态等。其中 LLM-Router 和 GCache 属于常见的分布式缓存和调度实现,本文不展开。重点聚焦在 SWA 相关 KVCache 这条主线上。


2. SWA 的 KV Cache 怎么存

核心设计:物理分离 + 逻辑统一

物理层:两个独立池子

每层存储(L1 GPU / L2 Host / L3 GCache)都并排维护两个 KV Pool:

1
2
3
4
5
6
7
8
9
┌─────────────────────────────────────┐
│ Full KV Pool │
│ ├─ 大小: O(seq_len × 10 层) │
│ └─ 淘汰: 按整个序列做 LRU │
├─────────────────────────────────────┤
│ SWA KV Pool │
│ ├─ 大小: O(W × 60 层) ← 严格 O(W) │
│ └─ 淘汰: 独立按 window 做 eviction │
└─────────────────────────────────────┘

关键点:SWA pool 物理上就只有 W 大小,写入新 token 时按窗口 evict 旧的。不是「存完整历史然后只用尾部 128 个」,而是存进来就已经是窗口截断后的样子

直观对比(128K token 序列):

Full Attention Hybrid SWA
60 个 SWA 层总量 60 × 128K = 7.68M 60 × 128 = 7.68K
10 个 Full 层总量 10 × 128K = 1.28M 10 × 128K = 1.28M
合计 8.96M ≈ 1.29M(省 7×)

Hybrid SWA 存储量几乎等于只算 Full 层的量,SWA 层几乎不占空间。

逻辑层:单序列视图 + 双索引前缀树

上层(前缀树、调度器)只看到一个逻辑序列,底下用 Full→SWA mapping 做透明分层。每个前缀树节点存两套元数据:

1
2
3
4
5
PrefixTreeNode {
tokens: [t0, t1, ..., tN]
full_seg_idx: [Full KV pool 位置] # Full 层复用
swa_seg_map: [SWA KV pool 位置 or 空] # 判断 window 安全性
}

用途

  1. 命中判定:token 相等 + tail W 个 token 的 swa_seg_map 都非空 → 才算真命中
  2. 独立淘汰:window 外的 SWA 段可以先扔,Full 段还留着

一句话总结:SWA 存的就是「每 SWA 层永远只存 W 个 token 的 KV」,SWA pool 和 Full pool 物理分离,上层用双索引前缀树抽象成单序列视图。


3. SWA Cache 怎么 Offload 到 CPU 内存

存什么:只存「窗口内有效 slot」

对一条 100K token 的序列:

1
2
Full 层 (10 层):  每层 100K 个 slot  ← 全部下沉到 CPU
SWA 层 (60 层): 每层只有 W=128 个 slot ← 只存窗口内

SWA 下沉到 CPU 的就是窗口截断后的形态,不是完整历史。

怎么存:Host 侧独立的 SWA Pool

L2 结构完全镜像 L1:

1
2
3
Host DRAM (L2):
Full KV Pool: [10 层, seq_len, head_dim], paged + pinned memory
SWA KV Pool: [60 层, W=128, head_dim], 独立 eviction

工程要点:

  • Pin memory:Host 池子用固定内存,D2H/H2D 用异步 DMA
  • Paged 分配:SWA pool 拆成 block,滑出窗口的 block 归还池子
  • 和 Full pool 完全独立:分开的地址空间、eviction 策略

怎么搬:按 SWA mask 传输

论文 §3.1.1 的关键一句:

“Cross-tier transfers are performed based solely on the SWA mask, ensuring only valid window data is moved.”

D2H / H2D 时不搬窗口外的数据

1
2
3
D2H writeback: 只把 mask=1 的 slot 打包 DMA 到 host
H2D prefetch: 只拉窗口内的 W 个 token 的 KV
→ 数据量小,layerwise prefetch 能 overlap 计算

这里省的是带宽。按 mask 走,60 层每层就 128 个 token,总量小到可以忽略。

D2H 时机

  • Eviction 前的备份:GPU SWA pool 满了,先 D2H 到 host 再释放
  • 一致性修复(§3.1.4):前缀树节点合并 / prefill chunk 完成时,检查 device 和 host 的 SWA 占用差 → host 补齐槽位 → 异步 D2H
  • 请求结束 / 会话切换:高价值会话主动 D2H 到 host,甚至下沉到 L3

一句话总结:SWA 下沉到 CPU 的就是窗口截断后的 128 个 slot,Host 侧用独立 pool + pinned memory 存,搬运只按 SWA mask 走有效位置。


4. 共享前缀命中时,SWA 要不要重算

这是 Hybrid SWA 落地生产最难受的一个点

问题本质

典型场景:共享 system prompt(2000 tokens),下面挂了很多用户会话:

1
2
3
4
PrefixTree:
system_prompt (2000 tokens) ← 非叶子节点
├─ user_A_session (往下延展了 50K tokens)
└─ user_B_session (往下延展了 30K tokens)

关键观察:user_A 会话延展了 50K tokens,当前窗口在 [52000-128, 52000] 附近。system_prompt 段(位置 0-2000)早就滑出窗口了,SWA KV 被 evict 掉了。

这就是「伪命中」

user_D 带着同一个 system prompt 进来:

1
2
3
4
5
命中判定:
1. Token equality 检查 → 匹配到 2000 tokens ✅
2. Full KV 检查 → 都在 → Full 命中 2000 tokens ✅
3. Window-safe 检查 → tail 128 个 SWA slot 在不在?
└─► ❌ 已被 evict

按「window-safe length」规则,SWA 匹配长度被 clip。Full 层能复用 2000 tokens,但 SWA 层不行——SWA 层要看窗口内 KV,位置 2001 的 attention 需要 [1873, 2001] 的 KV,这段没了 → 只能重算 SWA 层的 prefill。

代价

  • Full 层:省了(Full pool 还留着)
  • SWA 层:60 层全部重跑一次 prefill(占 6/7)
  • 60 层重算,几乎等于白干

小米的应对:SWA 保留策略(§3.1.4 第 4 条)

论文直接点名了这个场景:

“Medium/short sequence SWA retention strategy. Based on user request patterns, we retain relatively dense SWA KV Cache at fixed length positions for medium/short sequences… particularly beneficial for long agent sessions, multi-user shared system prompts, and repeated tool calls to the same codebase.

翻译:对高频复用的前缀(比如 system prompt),SWA pool 不严格 O(W),而是在关键位置留几份 SWA 快照

推测的实现方式:

1
2
3
4
5
6
7
普通序列的 SWA pool:
只留 tail W 个 slot ← 严格 O(W)

高频前缀(如 system prompt):
在几个 anchor 位置多留 W 个 slot
比如: [0, W], [500-W, 500], [1000-W, 1000], [2000-W, 2000]
← 相当于"检查点",每个都能作为 SWA 计算的起点

这样 user_D 命中时:匹配到 system_prompt 末尾 tail W=128 的 SWA slot 确实存在(是保留下来的 anchor),Full + SWA 都能完整复用 → 不需要重算 SWA 层

Trade-off 账

严格 O(W) 高频前缀密集保留
SWA 存储占比 极低(1/7) 略高(几倍 W)
单节点并发能力 稍低
共享前缀命中率 差(伪命中) 好(真命中)

核心思路用一点点 SWA 存储换取显著提升的共享前缀命中率。因为 system prompt 被成千上万用户共享,每次省掉 60 层 SWA prefill 的收益,远远大于多存几份 128-slot 快照的代价。

一句话总结:非叶子共享前缀命中时,默认严格 O(W) 实现会因为 tail 被 evict 导致 SWA 层必须重算。小米的解法是对高频共享前缀主动保留多个 SWA 快照,让 Full + SWA 都能完整复用,真正兑现 Hybrid SWA 在多用户共享场景下的效率红利。


5. numa_balancing 为什么影响 +10%

论文里就一句话:「关掉 numa_balancing,端到端 +10%」。背后是操作系统调度和 GPU 推理的碰撞

numa_balancing 是干嘛的

现代服务器都是 NUMA 架构,CPU 访问本 NUMA 节点内存 << 访问远端节点内存(延迟差 1.5-2×)。kernel.numa_balancing=1 是 Linux 的「自动 NUMA 均衡」特性:

1
2
3
4
5
6
7
内核周期性做的事:
1. 扫描进程页表 → 把页面标记为 PROT_NONE
2. 进程下次访问 → 触发缺页异常
3. 内核在 fault handler 里统计 NUMA locality
4. 如果 CPU 和 page 不在同一 node:
- 迁移页面到 CPU 所在 node
- 或者迁移线程到 page 所在 node

但在 GPU 推理场景下变成灾难

SGLang 有自己的 --numa-node 配置,会主动把 rank 绑到指定 NUMA node(PIN 死)。冲突来了

  • SGLang 说:这个 rank 就吃 Node 0,别乱跑
  • 内核 numa_balancing 说:我来「优化」一下

内核会做几件破坏性的事:

  1. 周期性把页面标 PROT_NONE:即使你已经手动 pin 好了,下一次访问必触发 page fault
  2. Page fault 引起大 stall:GPU 计算 → 触发 page fault → 内核处理 → GPU 继续(延迟数百 μs 到 ms)
  3. 随机页面迁移:pinned memory 被迁移 → GPU 之前记录的 DMA 地址失效

为什么这个 bug 特别难缠

论文里这句话说到点子上了:

“In multi-node multi-GPU deployments, these gaps appear at random positions across ranks, and each inter-rank synchronization is bottlenecked by the slowest rank.”

关键组合拳

  • 随机性:numa_balancing 是周期性异步扫描,不同 rank、不同时刻会随机中招
  • 全局同步放大伤害:LLM 推理每层 attention 之后都要 all_reduce最慢那个 rank 拖累所有人
1
2
3
所有 rank 计算 ──► 集合通信同步 ──► 下一层

最慢那个 rank 拖累所有人

如果 8 个 rank 里有 1 个刚好被 page fault 命中,那 8 个 rank 全都得等它——尾延迟灾难。对一次前向 60+ 层,只要偶尔一层有一个 rank 中招,整个 latency 就翻倍。

为啥关掉直接 +10%

  1. SGLang 已经手动 pin 好了 NUMA,调度器安排得明明白白
  2. numa_balancing 是「帮倒忙」,它假设进程没做优化,主动帮你搬,反而破坏了原有布局
  3. 代价是 hard cost:page fault、TLB shootdown、页面迁移都是不可避免的 CPU 侧开销

关掉后:内核不再主动折腾 → 页表稳定 → 没有意外 page fault → GPU kernel 之间没有随机 gap → 集合通信不再被拖尾 → 直接 +10%。

一句话总结:SGLang 已经手动做了 NUMA pin,numa_balancing 在这基础上做二次干扰,通过 page fault + 页面迁移引入随机 stall,被 GPU 集合通信同步放大成尾延迟灾难,关掉就直接 +10%。


6. Prefill 为什么也要开 MTP

默认直觉是「MTP 是加速 decode 的,跟 prefill 有啥关系」,但正是这个直觉害惨了 agentic 场景。

MTP 是什么

MTP (Multi-Token Prediction) = 多 token 预测。MiMo-V2.5 系列原生支持 3 层 MTP

1
2
3
4
主模型: 预测第 t+1 个 token
MTP Layer 1: 同时预测第 t+2 个 token
MTP Layer 2: 同时预测第 t+3 个 token
MTP Layer 3: 同时预测第 t+4 个 token

核心思想:一次 forward 出 4 个 token,然后主模型验证,接受的部分直接输出。理想情况下 decode 速度接近 4×。

关键前提:MTP 层本身有参数、有 KV cache,需要**跟着上下文一起「预热」**才能预测准。

原始问题:Prefill 不跑 MTP 的后果

默认实现:Prefill 阶段只跑主模型,跳过 MTP 层。

问题:MTP 层需要看历史 context 才能预测。如果 prefill 跳过 MTP:

1
2
3
4
5
6
7
Prefill 结束时:
主模型 KV cache: [完整 prompt 的 KV] ✅
MTP 层 KV cache: [空 / invalid] ❌

Decode 第 1 个 token:
主模型: ✅
MTP 层: context 是空的,瞎猜 → 接受率极低 ❌

MTP 需要多久「热身」:论文说是 128 个 token。因为 MTP 层需要逐个 decode 的 token 慢慢累积 KV。

这段时间 MTP 基本白搭:主模型每步生成 1 个 token,MTP 提议的 3 个 token 因为 context 不足几乎全被拒,有效加速 ≈ 1×。

为什么 agentic 场景特别惨

论文这句话点破了关键:

Since agentic scenarios involve mostly short output sequences, this limitation significantly limited MTP’s effective speedup.”

Agentic 场景的输出特征:

1
2
3
一次工具调用响应: ~20-50 tokens
一次 function call: ~30-80 tokens
一次思考步骤: ~50-150 tokens

几乎所有输出都在 128 token 之内——这正好是 MTP 需要热身的窗口。结果:

1
2
理论加速比: 3× (3层 MTP)
Agentic 场景实际加速比: ~1× (刚热身完就结束了)

等于 MTP 层白装了

解法:Prefill 阶段也跑 MTP

小米的改动很直接:prefill 时也让 MTP 层参与前向

1
2
3
4
5
6
7
Prefill 结束时:
主模型 KV cache: [完整 prompt 的 KV] ✅
MTP 层 KV cache: [完整 prompt 的 MTP KV] ✅ ← 新增

Decode 第 1 个 token:
主模型: ✅
MTP 层: 已经用整个 prompt 预热过,接受率立即很高 ✅

效果(论文数据)

  • 0-128 token 加速: 2.3× ← 从近 1× 提升到 2.3×,agentic 场景直接受益
  • 128-256 token 加速: 1.5×

工程适配

论文说:

“By introducing MTP support during prefill with dedicated adaptations and optimizations for HiCache L2/L3…”

关键点:MTP 层有自己的 KV cache,之前 HiCache 那套 offloading 系统只管主模型的 KV。要 prefill 阶段跑 MTP,就得:

  1. MTP KV 也要进 HiCache:L1 (GPU) ↔ L2 (Host) ↔ L3 (GCache)
  2. Prefix cache 得覆盖 MTP KV:前缀树节点要标 MTP 层的状态
  3. 传输和存储成本:MTP 有 3 层,KVCache 总量变多,L2/L3 得扛住

Trade-off 账

维度 不开 prefill MTP 开 prefill MTP
Prefill 计算 70 层 70 + 3 = 73 层(多 4%)
Decode 0-128 tokens 加速 ~1× 2.3×
Agentic 场景总吞吐 显著提升

核心逻辑用 prefill 的少量额外开销,换 decode 从第 1 个 token 就享受 MTP 加速。Agentic 场景短输出多,这笔账非常划算。

具体场景:Agent 调用工具,输入 5K token,输出 60 token

1
2
3
没开 prefill MTP: 100ms (prefill) + 1200ms (decode) = 1300ms
开了 prefill MTP: 104ms (prefill) + 520ms (decode) = 624ms
省 52%

一句话总结:Prefill 开 MTP 是为了给 MTP 层的 KV cache 做「上下文预热」,避免 decode 初期因为 MTP 无历史导致接受率极低。Agentic 场景输出短,正好都落在这个「未热身窗口」里,所以收益特别大——0-128 token 加速直接从 ~1× 提升到 2.3×。


总结

这篇论文的核心贡献不是某个单点优化,而是把 Hybrid SWA + MoE + 多模态这套组合架构从理论拉到生产的完整工程实践。本文聚焦在 SWA 相关 KVCache 这条主线,拆解了几个关键问题:

  1. SWA 存储:物理分离 Full/SWA pool,SWA 严格 O(W),前缀树双索引做透明分层
  2. SWA offload:Host 侧独立 SWA pool + pinned memory,搬运只按 mask 走有效位置
  3. 共享前缀命中:默认严格 O(W) 会伪命中,小米用「SWA 快照保留」策略换真实命中率
  4. numa_balancing:操作系统自动优化 vs 手动 NUMA pin 的冲突,关掉 +10%
  5. Prefill MTP:给 MTP 层 KV cache 做上下文预热,agentic 短输出场景 0-128 token 从 ~1× 提升到 2.3×

可迁移的启示

  • 架构红利不会自动兑现,工程系统必须跟着架构一起改造
  • Hybrid SWA 的价值在多用户共享场景,但需要专门的缓存策略才能兑现
  • 推理系统的 pipeline 各阶段不是孤立的,跨阶段依赖的组件需要协同优化
  • **系统级坑(numa_balancing、THP、CPU governor)**对已经手动优化过的高性能场景是纯干扰,该关就关

参考

  • 论文:MiMo Team, Xiaomi. Full-Pipeline Inference Optimization for MiMo-V2.5 Series: Pushing Hybrid SWA Efficiency to the Limit. arXiv:2607.13095, 2026.
  • 模型:Xiaomi MiMo Team. MiMo-V2.5 / MiMo-V2.5-Pro. HuggingFace, 2026.

本文基于 MiMo-V2.5 论文、个人笔记《MiMo-V2.5 SWA offload 方法》整理。

DeepSeek-V4 的架构图有一个特点:本身就是一份并行性说明书。画成并行分支的模块,基本都是在告诉系统实现者"这里有 overlap 的空间"。本文梳理 V4 中可 overlap 的算子对,对照论文承诺与 SGLang 实现现状。


一、Overlap 的三个前提

  1. 算法独立 — 两条路径无数据依赖(论文画成并行分支)
  2. 资源不冲突 — 不争用同一 SM 或同一 buffer
  3. 工程可兜住 — CUDA stream / multi-kernel launch 的硬件支持

DeepSeek 的特殊性在于:MoE + MHC(Multi-Head Compressor)双重稀疏结构,天然产生多条独立路径


二、Attention 分支内的 Overlap

2.1 wqkv_a 与 Compressor 的 Overlap

数据依赖

1
2
3
4
5
6
7
wqkv_a:  输入 = hidden (x)

输出 = q_a + kv_a + k_rope

compressor: 输入 = hidden (x) + past KV (来自 KV cache)

输出 = summary_K_4 + summary_K_128

两者都读 hidden,但 compressor 不依赖 wqkv_a 的输出,可以完全并行。

资源 pattern 互补

Kernel 性质 瓶颈资源
wqkv_a GEMM (M=4096, N=2112, K=7168) 中 GEMM MFU ~30-40%,CU 大量闲置
compressor flash_c4 / flash_c128 memory-bound HBM 带宽吃满,MFMA 单元闲置

一个 compute-bound 倾向,一个 memory-bound 倾向,硬件资源不冲突。

L2 Cache 共享

hidden 的大小:[4096, 7168] BF16 = 56 MB,接近 MI355X 的 L2 cache(32 MB)。

  • 串行wqkv_a 读完 hiddencompressor 再读——大概率已被 q_b 挤出 L2,又一次从 HBM 读
  • 并行:两个 kernel 同时从 HBM 拉 hidden,第二次的 cache line 是免费的

2.2 Indexer 的 Overlap:只依赖 q_a,不依赖 q_b

这是 V4 论文里 “MLA Co-Design with Indexer” 的核心设计。

Indexer 两条路径的依赖

1
2
3
4
5
6
7
Indexer K-side:  W_K^I · x  →  RoPE  →  summary_K

只需要 hidden (x),不需要 wqkv_a 的输出

Indexer Q-side: W_Q^I · q_lora → Hadamard + RoPE

只需要 q_lora (q_a 的输出),不需要 q_b 的输出

q_b 是 critical path 上最重的 GEMM(929 µs),indexer 不等 q_b 完成就可以启动

为什么 Q-side 用 q_lora 而不是 hidden

V4 论文的算法决策:

  • 实验结论:召回率无差异
  • 系统收益:每层省一次 [T, 7168] → [T, ...] 的大 GEMM
  • 代价:Q-side 多一个 q_lora_ready event 依赖(K-side 完全自由)

代码里精确实现了这个依赖:

1
2
3
4
# K-side:不依赖 q_lora,可以立刻启动
stream_indexer.wait_stream(current_stream)
# Q-side:只等 q_lora_ready,不等 q_b
# (在 indexer 内部用 q_lora_ready 精确控制)

2.3 三者在时间轴上的 Overlap

1
2
3
4
5
6
7
8
9
10
11
12
13
时间轴(prefill 4096 tokens,典型层):

0 µs 77 µs 83 µs 1012 µs
│ wqkv_a │ q_a │ q_b (929 µs) │
│─────────┴─────┴───────────────────────────────┤
│ │ │
│ └──► indexer K-side + Q-side │ ← 不等 q_b
│ (q_lora_ready 触发)
│ │
│ └──► compressor (c4 + c128) │ ← 完全独立
│ (只等 x)

└──► kv_write (等 wqkv_a 输出切片) │

端到端时间 = max(q_b, indexer, compressor, kv_write) = ~1012 µs(由 q_b 主导)

串行时 = wqkv_a + q_a + q_b + indexer + compressor + kv_write = ~2290 µs


三、MoE 分支内的 Overlap

3.1 Shared Expert 与 Routed Expert

两条 expert 路径完全独立,SGLang 用 alt_stream 实现,是已实现的最好案例。

3.2 MoE Wave 模型:计算与通信 Overlap

Wave 分 chunk 乒乓:dispatch → compute → combine 流水,掩盖全量通信延迟。V4 论文 Fig.5 的时序图明确画出了这一点。

3.3 Combine 通信与 Shared Expert 尾部计算

Combine(NVLink all-to-all)不需要 shared expert 的完整输出,可以和 shared expert 的 down_proj 最后一部分 overlap。SGLang 目前串行等待,未利用。


四、为什么 Multi-Stream Overlap 能赚到时间

CUDA stream 是 GPU 上的 FIFO 工作队列;stream 本身只给调度器自由度,性能要靠"互补的资源占用"赚出来。

机制 1:单 Kernel 资源利用率低(最主要)

q_b 把计算压满但 HBM 闲;swa_scatter 把 HBM 压满但计算闲。不同 stream 让它们同时跑,各用各的硬件资源。

机制 2:小 Kernel Grid 填不满 SM

MI355X 有 256 个 CU,但 trace 里很多 kernel 的 grid 很小(rocprim cumsum 只有 1 个 CU,占用率 0.4%)。多 stream 让调度器把多个小 kernel 同时塞进 CU。

机制 3:隐藏 CPU Launch Overhead

每次 hipLaunchKernel 在 host 端要 ~3-5 µs。Trace 里 ~30000 个 kernel × 4 µs ≈ 120 ms 纯 launch 开销。多 stream 让 host 预先 enqueue,GPU 不饥饿。

机制 4:同源输入的 L2 Cache 共享

wqkv_acompressor 都读 hidden,并行时第二次读直接 cache hit,省 ~20% 输入带宽。

反向判断:什么情况下多 Stream 不赚

场景 多 stream 是否有用 原因
两个高 MFU GEMM(都跑 70%+) SM 都被占满,只是 timesharing
两个 memory-bound op 读不同数据 HBM 总带宽是上限
两个 op 有数据依赖(A → B) 必须串行
一个 compute-bound + 一个 memory-bound 经典 case
一个大 GEMM + 一群 µs 级小 kernel ✅✅ 赚得最多

五、工程实现:SGLANG_OPT_USE_MULTI_STREAM_OVERLAP

这个关键优化靠一个环境变量控制,不在 --help 里:

1
2
3
4
5
# 开启
SGLANG_OPT_USE_MULTI_STREAM_OVERLAP=1 python -m sglang.launch_server ...

# 验证
python -c "import os; print(os.environ.get('SGLANG_OPT_USE_MULTI_STREAM_OVERLAP'))"

读取位置在 sglang/srt/layers/deepseek_v4.py_forward_prepare_multi_stream 方法入口,通过 os.environ.get() 判断,默认关闭。

实测效果

在 decode 阶段(batch size 中等,61 层 DeepSeek-V4),开启后整体 forward 有 ~3.5% 的端到端提升

状态 forward latency (decode) 相对提升
关闭(串行) 基准
开启(multi-stream overlap) -3.5%

提升主要来自:

  • q_b GEMM(929 µs)与 indexer + compressor 并行:decode 时 q_b 仍是瓶颈,但 indexer K-side 和 compressor 可以完全隐藏在其执行期间
  • 小 kernel(swa_scatter、cumsum 等)与主体 GEMM 并行:这些 µs 级 kernel 在串行时被 q_b 的 launch gap 放大,overlap 后基本被吸收

注意:3.5% 是 decode 阶段的收益。prefill 阶段因为 q_b 的 GEMM 更大(M=4096),overlap 的相对收益会被稀释,但绝对时间节省更显著(每层 ~1.28 ms,61 层 ~78 ms)。

没开时,整个 _forward_prepare_multi_stream 退化成串行,论文里"sparse attention 模块在 prefill 阶段近乎 free"的 claim 直接失效。


六、对照论文图:承诺 vs 现状

论文图中的结构 论文章节 承诺的并行性 SGLang 状态
Indexer ∥ wqkv_a §3.2.2 双分支并行 ✅ 算法并行,⚠️ 需 SGLANG_OPT_USE_MULTI_STREAM_OVERLAP=1 开启
Compressor c4 ∥ c128 §3.2.1 双尺度评分并行 ❌ 串行
Shared ∥ Routed Expert §3.3 双 expert 路径并行 ✅ 已实现
Wave dispatch/compute/combine §3.3.3 乒乓 overlap ✅ 已实现
KV store ∥ next layer indexer §3.2 层间 pipeline ❌ 未做

七、总结

DeepSeek-V4 论文在算法层面为 overlap 留了很大空间,尤其是 Fig.3(MHC 结构)和 Fig.5(Wave 时序图)。当前 SGLang 实现只吃到了 Shared Expert / Wave 的部分,仍有明显优化空间。

实测数据验证了 overlap 的价值:

  • decode 阶段:开启 SGLANG_OPT_USE_MULTI_STREAM_OVERLAP,整体 forward 提升 ~3.5%
  • prefill 阶段:每层节省 ~1.28 ms,61 层共 ~78 ms,sparse attention 模块接近"free"的目标

更一般的启示:算法论文画依赖图时多想一步系统实现,系统实现时多对照论文的并行性承诺。


参考:DeepSeek-V4 论文(arXiv:2606.02405)§3.2 Attention Mechanism, §3.3 Expert Mixture, §3.3.3 Dynamic Expert Routing & Wave Scheduling

背景

训练超长序列 LLM 时,单卡显存放不下完整的 KV cache,需要对序列维度做并行(Sequence Parallelism)。目前主流有两种方案:

  • Ulysses(DeepSpeed-Ulysses):两次 All-to-All,按 head 切分
  • Ring Attention:环形传递 KV blocks,分块计算

两者目标相同,但设计假设完全不同。本文从原理、通信量、代码实现到架构选择动机,做完整对比。


一、核心思想对比

Ulysses Attention

1
2
3
Step1: [N/P, h, d] ──All-to-All──→ [N, h/P, d]
Step2: 本地做标准 Attention(每张卡拿到全序列、部分 head)
Step3: 输出 ──All-to-All──→ 回到 [N/P, h, d]

关键:两次 All-to-All 把"序列切"转成"head 切",每张卡对全序列做部分 head 的 attention。

Ring Attention

1
2
3
4
Round 0: attn(Q_i, K_i, V_i)           ← 本地 attention
Round 1: 收到 K_{i-1}, V_{i-1} → attn(Q_i, K_{i-1}, V_{i-1})
...
Round P-1: 收到所有 KV → online softmax 合并结果

关键:KV 沿环形传递,每张卡只存自己的 KV chunk,计算和通信完全重叠。


二、通信量数学推导

Ulysses:All-to-All 的 (P1)/P2(P-1)/P^2

Ulysses 输入是 [N/P, h, d](序列已被切,head 完整),输出是 [N, h/P, d](序列完整,head 被切)。

每 GPU 发送 P-1 个 chunk,每个 chunk 大小:

NP×hP×d=NhdP2\frac{N}{P} \times \frac{h}{P} \times d = \frac{N \cdot h \cdot d}{P^2}

总发送量:

send per GPU=(P1)×NhdP2=NhdP1P2\text{send per GPU} = (P-1) \times \frac{N \cdot h \cdot d}{P^2} = N \cdot h \cdot d \cdot \frac{P-1}{P^2}

两个 1/P 的来源:

  • 第一个 1/P:序列维度已被切(N/P
  • 第二个 1/P:head 维度再切一次(h/P

Ring:P2P 的 (P1)/P(P-1)/P

Ring 只传 KV(Q 不动),每轮传 (N/P) × d_kv

send per GPU=(P1)×NP×dkv=NdkvP1P\text{send per GPU} = (P-1) \times \frac{N}{P} \times d_{kv} = N \cdot d_{kv} \cdot \frac{P-1}{P}

只有一个 1/P(序列切分),没有 head 切分。

对比(P=8, h=64, d=128, d_kv=576)

方法 通信量/GPU 比值
Ulysses (MHA) N×64×128×7/64=N×896N × 64 × 128 × 7/64 = N × 896 1.8×
Ring (MHA KV) N×128×2×7/8=N×224N × 128×2 × 7/8 = N × 224
Ring (MLA) N×576×7/8=N×504N × 576 × 7/8 = N × 504

注意:MHA 的 KV 是 h × d × 2 = 16384/token,MLA 的 c_kv 只有 576/token,这是后续分析的关键。


三、MLA 对 Ulysses 的致命问题

MLA 的 KV cache 结构

DeepSeek-V3/V4 使用 MLA(Multi-head Latent Attention),KV cache 不是 multi-head 的:

1
2
3
c_kv [seq, 512]     ← 压缩的共享 latent
k_rope [seq, 64] ← RoPE 位置编码部分
合计:576 dim/token(vs MHA 的 16384 dim/token)

只有 1 个"头",无法按 head 切分。

DeepSpeed Ulysses 代码验证

1
2
3
4
# deepspeed/sequence/layer.py 核心逻辑
q = _SeqAllToAll.apply(group, query, scatter_idx=2, gather_idx=0)
k = _SeqAllToAll.apply(group, key, scatter_idx=2, gather_idx=0) # K 和 Q 对称处理
v = _SeqAllToAll.apply(group, value, scatter_idx=2, gather_idx=0)

Q、K、V 走完全相同的 All-to-All 路径,没有"KV 走 All-Gather"的分支。当 num_kv_heads=1 时:

1
2
1 KV head / 4 GPUs → [1, 0, 0, 0]
GPU 1-3:没有 KV head → 无法计算 ❌

Ulysses SP 上限 = num_kv_heads

模型 num_kv_heads Ulysses SP 上限 实际可用性
DiT (视觉) 32 (MHA) 32
LLaMA-3 8 (GQA) 8 ⚠️ 受限
DeepSeek-V3/V4 1 (MLA) 1 ❌ 不可用

四、MLA 场景:Ring vs Ulysses 精确对比

通信量(P=8)

Ring Attention(MLA)

  • 只传 c_kv:每轮 (N/8) × 576,共 7 轮
  • 总计:N×504N × 504 / GPU

Ulysses(MLA,Q All-to-All + KV All-Gather)

  • Q All-to-All:7×(N/8)×64×512=N×28,6727 × (N/8) × 64 × 512 = N × 28,672
  • KV All-Gather:7×(N/8)×576=N×5047 × (N/8) × 576 = N × 504
  • Output reverse:N×28,672N × 28,672
  • 总计:N×57,848N × 57,848 / GPU

Ring 比 Ulysses 省 115×。

根本原因:MLA 的非对称性——Q 巨大(32768 dim/token)、KV 极小(576 dim/token)。Ring 只移动小的 KV,Ulysses 被迫移动巨大的 Q。

缩放性对比

P MHA Ulysses MLA Ring MLA 比 MHA 省
4 N×6,144N × 6,144 N×432N × 432 14.2×
8 N×3,584N × 3,584 N×504N × 504 7.1×
16 N×1,792N × 1,792 N×540N × 540 3.3×
64 N×428N × 428 N×567N × 567 0.75×(MHA 反超)

交叉点:P=4hd/dkv57P = 4hd/d_{kv} ≈ 57,但 Ulysses SP 上限 = 64,所以实践中 MLA Ring 几乎总是更优


五、为什么 Ulysses 在视觉模型流行,文本 LLM 不用?

四个结构性原因

1. KV head 数量趋势

1
2
3
4
2020: MHA (GPT-3) num_kv_heads = 96  → Ulysses 随便用
2022: MQA (PaLM) num_kv_heads = 1 → Ulysses 废了
2023: GQA (LLaMA-2) num_kv_heads = 8 → Ulysses 受限
2024: MLA (DeepSeek-V3) num_kv_heads = 1 → Ulysses 废了

文本 LLM 全面转向 GQA/MLA 压缩 KV heads,Ulysses 的前提条件被釜底抽薪。

2. 文本 LLM 的 TP 已经做了同样的事

1
2
Tensor Parallelism: 按 head 切权重 → 本地算 attention → All-Reduce
Ulysses: 按 head 切数据 → 本地算 attention → All-to-All

本质重叠,TP 已经切了 head 之后,Ulysses 没有额外收益。

3. 推理阶段 Decode 占 80% 时间

1
2
Prefill: ~20% 时间(可以序列并行)
Decode: ~80% 时间(Q=1 token,序列并行无用)

Ulysses 对 Decode 完全无用。视觉扩散模型没有 Decode 阶段,全程受益。

4. 视频生成 token 数极高

1
Sora 级别: (64×64) × 120 frames = 491,520 tokens

必须序列并行,且 DiT 用标准 MHA(32 heads),Ulysses 完美适配。

一句话总结

Ulysses 在视觉模型流行 = MHA(head 够多)+ 无 decode + 没被 TP 覆盖 + 序列极长。文本 LLM 四条全占不到。


六、DeepSeek 的实际选择

训练:全程 Ring/CP

DeepSeek-V3 技术报告显示,长上下文扩展(32K→128K)用的是 Ring Attention(Context Parallel),不是 Ulysses:

  • MLA 只有 1 个 KV latent → Ulysses 物理上不可用
  • KV 只有 576 dim → Ring 通信量本来就小
  • Ring 的通信-计算 overlap → 长序列时通信几乎免费

推理:SGLang 的 CP 配置

1
2
3
--enable-nsa-prefill-context-parallel
--attn-cp-size 8
--nsa-prefill-cp-mode round-robin-split

选择 Ring 风格 CP 的原因:

  • MLA 的 KV 只有 1 个共享 latent,切不了 head
  • KV 极小(576 dim),Ring 通信可接受
  • round-robin 切分 tokens 均衡负载

七、LoRA 压缩能让 Ulysses 复活吗?

思路:在压缩态做通信

MLA 的 Q 路径:hidden(7168) → q_lora(1536) → expand → Q[128, 576]

能否在 q_lora 维度(1536)做 All-Gather,而非展开后的 Q(73728)?

通信量更新(LoRA 压缩版)

通信 数据 P=8 总量
q_lora All-Gather N×1536×7/8N × 1536 × 7/8 N×1,344N × 1,344
c_kv All-Gather N×576×7/8N × 576 × 7/8 N×504N × 504
o_lora Reduce-Scatter N×1024×7/8N × 1024 × 7/8 N×896N × 896
总计 N×2,744N × 2,744

vs 原始 Ulysses(展开态):N×57,848N × 57,848压缩 21×

但仍然比 Ring 大 5.4×

1
2
Ring:              N × 504
Ulysses (LoRA): N × 2,744 ← Q 和 O 的压缩态还是要传

根本原因:Ring 的 Q 和 O 根本不过网络,留在本地计算。Ulysses 无论怎么压缩,都要把 Q/O(或其压缩态)在网络上搬一次,这是架构级差距,压缩只能缩小、不能逆转。


八、全景决策树

1
2
3
4
5
6
7
8
9
10
num_kv_heads >= SP_size?

├─ YES (DiT, ViT, 视觉 MHA)
│ └─ Ulysses ✅(All-to-All,通信量 1/P²)

└─ NO (GQA=8, MLA=1, 文本 LLM)
├─ 训练:Ring/CP ✅
└─ 推理:
├─ Prefill:Ring/CP
└─ Decode:partial attn + All-Reduce

总结

Ulysses 的核心优势是 1/P² 通信缩放,但前提是 num_kv_heads ≥ SP。MLA 把 KV 压缩到 1 个 latent,直接废掉这个前提。Ring 只传 KV(576 dim),Q 留本地,反而成了 MLA 的最优搭档。

这不是巧合——MLA 和 Ring CP 是刻意协同设计,不是将就。


相关阅读:DeepSeek-V3 技术报告、DeepSpeed-Ulysses 论文(arXiv:2309.14509)、Ring Attention 论文(arXiv:2310.01889)

从残差到 mHC:一条清晰的三代演进路

如果你训练过深层 Transformer,一定对残差连接(Residual Connection)不陌生。它是 ResNet 留下的最重要的遗产之一,也是今天所有大语言模型的标配。

但残差连接有个根本问题:它只有一条信息流

DeepSeek 在 2025 年连发两篇论文,把这个问题彻底讲透了。故事的三步是:

  1. Residual(2015)x_{l+1} = x_l + F(x_l) — 一条流,稳定但表达弱
  2. HC(2024):把残差流拓宽到 n 条并行流,表达能力暴涨,但训练直接崩
  3. mHC(2025):给 HC 加上数学约束,又强又稳

这篇文章把这三代讲清楚,重点放在 mHC 上。


残差连接的瓶颈:一条流不够用

标准 Transformer 里,每一层的计算是这样的:

1
2
x = x + Attention(Norm(x))  # 残差连接
x = x + MoE(Norm(x)) # 残差连接

Layer 0 到 Layer 42 的所有信息,全都挤在同一个向量里做加法。

浅层的语法信息、中层的语义信息、深层的推理信息,互相覆盖、互相干扰。这是残差连接的天花板。

HC(Hyper-Connections) 提出了一个很自然的想法:

为什么不把 1 条流扩成 n 条并行流,让不同深度的信息走不同的"车道"?

HC 的公式长这样:

Xl+1=BlXl+ClF(AlXl)X_{l+1} = B_l X_l + C_l F(A_l X_l)

  • A_l:把 n 条流压缩成 1 条,送给 Attention/MoE
  • F:正常的层计算
  • C_l:把 F 的输出扩回 n 条流
  • B_l:n×n 矩阵,控制 n 条流之间怎么混合(残差项)

n=4 的时候,FLOPs 几乎没增加(F 只算 1 次),但信息容量翻了 4 倍。

听起来完美,对吧?


HC 的致命缺陷:训练会崩

问题出在 B_l 上。

HC 里的 B_l 是完全可学习的,没有任何约束。当模型叠到 60 层、参数量到万亿级时,这个无约束的矩阵会出问题:

多层复合后,信号被无限放大或衰减。

数学上,HC 跨层递归展开后得到:

xL=(Bi)xl+(Bj)CiTF(...)x_L = (\prod B_i) x_l + \sum (\prod B_j) C_i^T F(...)

ΠBiΠ B_i 是多层 BB 矩阵的复合。因为 BB 无约束,这个复合矩阵的谱范数可以远大于 1,也可以接近 0。

实测结果(DeepSeek 27B 实验):
HC 的复合映射增益(Amax Gain Magnitude)达到 ~3000,意味着信号在前向传播中可以被放大 3000 倍。训练到第 12k 步,loss 直接 spike,梯度范数爆炸。

HC 的作者(ByteDance Seed 团队)在 ICLR 2025 发表了这个想法,但没有解决稳定性问题。


mHC:给 HC 加上"安全阀"

DeepSeek 团队在 2025 年提出的 mHC(Manifold-Constrained Hyper-Connections),只做了一件事:

把 HC 中的残差混合矩阵 B_l 约束在双随机矩阵流形上。

什么是双随机矩阵?

一个 n×n 矩阵,满足:

  • 所有元素非负
  • 每行之和等于 1
  • 每列之和等于 1

这样的矩阵,谱范数 永远 ≤ 1。也就是说,它永远不会放大信号。

而且,双随机矩阵在乘法下是封闭的:两个双随机矩阵相乘,结果还是双随机矩阵。这意味着叠 100 层,复合映射仍然稳定。

一句话:mHC = HC + 双随机约束,恢复了残差连接的恒等映射性质,同时保留了多流的表达能力。


怎么把矩阵变成双随机?Sinkhorn-Knopp 算法

mHC 的核心算法是 Sinkhorn-Knopp 迭代,操作很简单:

1
2
3
4
5
6
7
8
9
# 输入: B_raw (4×4, 可能是负数)
M = exp(B_raw) # 先保证非负

# 交替行列归一化 20 次
for _ in range(20):
M = M / M.sum(dim=1, keepdim=True) # 行归一化
M = M / M.sum(dim=0, keepdim=True) # 列归一化

# 输出: B (4×4 双随机矩阵)

为什么第一步用 exp()?因为 B_raw 是神经网络输出的 logit,可能有负数。直接归一化会产生负权重,违反"非负"约束。exp() 把任意实数映射到正数,完美解决。

20 次迭代是实验得出的经验值,足够收敛,又不会太慢。


4 条流到底解决了什么问题?

很多人问:为什么是 4 条流?不是 2 条或 8 条?

三个核心作用:

1. 梯度高速公路

普通残差的梯度路径是 43 个乘法项的连乘,容易消失或爆炸。
4 条流提供了 4 条并行的梯度路径,一条堵了还有其他。类似 DenseNet 的思路,但高效得多(F 只算 1 次)。

2. 信息分离存储

1
2
3
4
stream_0: 可能专注浅层信息(位置、语法)
stream_1: 可能专注中层信息(语义、实体关系)
stream_2: 可能专注深层信息(推理链)
stream_3: 可能做"工作记忆"(当前层临时计算)

Layer 5 学到的特征,通过 B 矩阵的合理混合,可以在 Layer 30 仍然清晰可用。而在 1 条流里,这个特征早就被 25 次加法淹没了。

3. 动态容量分配

B 矩阵是输入依赖的(由当前层的 hidden state 动态生成):

  • 简单 token(“the”, “a”):B ≈ 单位矩阵,4 条流基本不混合,省计算
  • 复杂 token(需要深度推理):B 大幅混合,让更多层的信息参与

这是 1 条流做不到的——普通残差对所有 token 都是 x + F(x)

为什么选 n=4?
论文实测:n=4 是性价比最优的点。n=2 效果不够,n=8 的 Sinkhorn 开销(8×8 矩阵 × 20 次迭代)显著增加,但收益递减。


工程优化:为什么只多 6.7% 开销?

mHC 看起来很重:每行代码都有矩阵运算、Sinkhorn 迭代、4 倍 hidden state……

但 DeepSeek 做了三件事,把开销压到了 仅 +6.7%

1. Kernel Fusion(算子融合)

用 TileLang 把整个 mHC 计算——RMSNorm + 线性投影 + Sigmoid + Sinkhorn——融合成一个 CUDA kernel,减少内存读写。

2. Selective Recomputing(选择性重计算)

前向时只保存每块的第一个层输入,反向时重新计算 mHC 的中间激活。最优块大小由公式给出:

LrnLn+2L_r^* \approx \sqrt{\frac{nL}{n+2}}

3. Overlapping with DualPipe

把 mHC 的计算和流水线通信重叠。在 DualPipe 调度中,F_post,res kernel 放到高优先级流,避免阻塞 All-to-All 通信。


实验效果:又强又稳

DeepSeek 用 27B MoE 模型做了对比实验:

指标 Baseline HC mHC
训练稳定性 稳定 loss spike @12k步 稳定
最终 loss 下降 +0.021 vs baseline +0.027 vs baseline
复合映射增益 1.0 ~3000 ~1.6
训练开销 0% +?% +6.7%

下游任务(BBH / DROP / GSM8K 等 8 个 benchmark)上,mHC 全面超过 baseline 和 HC。

最重要的结论: mHC 让 1.6T 参数、60+ 层、MoE 模型能够稳定训练——这是普通 HC 做不到的。


一句话总结

mHC = HC + 双随机矩阵约束,是残差连接的最终进化形态,让万亿参数 MoE 能稳定训练。

如果你正在设计超大模型,mHC 是目前最值得考虑的残差连接方式。它不改 FLOPs,不改层计算,只改残差流的拓扑结构——但就是这一点,让深层训练从"可能崩"变成了"一定稳"。


参考资料:

  • Hyper-Connections, Zhu et al., ByteDance Seed, ICLR 2025, arXiv:2409.19606
  • mHC: Manifold-Constrained Hyper-Connections, Xie et al., DeepSeek, 2025, arXiv:2512.24880
  • DeepSeek-V4 Technical Report, 2026

写于 2026-06-02,基于 DeepSeek-V4-Pro 代码和 mHC 原始论文整理。

一句话概括

MegaMoE = 把 EP(Expert Parallelism)中的 All-to-All 通信和 MoE 计算融合到一个 kernel 里,让 NVLink 传输和 Tensor Core 计算时间上完全重叠。

传统 EP 里通信和计算串行,GPU 利用率只有 50-60%;MegaMoE 通过 Symmetric Memory + 细粒度 Scheduler + 单 Kernel 状态机,把利用率推到接近 100%。


传统 EP MoE vs MegaMoE

传统 EP(如 DeepEP)

1
2
3
4
时间 →

[dispatch all-to-all] → 等完 → [GEMM1] → [SwiGLU] → [GEMM2] → 等完 → [combine all-to-all]
NVLink 通信 空闲 计算 计算 计算 空闲 NVLink 通信

问题:通信和计算串行,GPU 要么在算要么在传,利用率约 50-60%。

MegaMoE

1
2
3
4
5
6
时间 →

┌───────────────────── 一个 Mega Kernel ─────────────────────┐
│ NVLink dispatch ←→ GEMM1 ←→ SwiGLU ←→ GEMM2 ←→ NVLink combine │
│ (通信和计算同时进行,流水线式重叠) │
└──────────────────────────────────────────────────────────────┘

通信隐藏在计算背后,GPU 利用率接近 100%。


传统 DeepEP 的 Overlap 能力分析

常见误解:“传统 DeepEP dispatch/combine 是无法和 MoE 计算 overlap 的”

更准确的说法:传统 DeepEP 可以通过 low-latency kernel + 多 stream + two-batch overlap 实现部分重叠,但需要框架手工编排,重叠率有限,对 decode 小 batch 几乎无效。

DeepEP 的两种模式

DeepEP 提供了两套 kernel,overlap 能力完全不同:

DeepEP 模式 能否 overlap 原因
Normal kernel(高吞吐) 基本不能 dispatch/combine 占满 SM 跑带宽,SM 被通信占住,GEMM 抢不到资源
Low-latency kernel(纯 RDMA) 可以 只用少量 SM 发 RDMA,大部分 SM 空着,可以留给 GEMM

你提到的"无法 overlapped"更接近 Normal kernel 的情况。

框架层怎么强行 overlap(传统方案)

SGLang / TRT-LLM 的做法是 micro-batch 切分 + 多 stream

1
2
3
4
时间线 ──────────────────────────────────►
stream A: [dispatch chunk0] [combine chunk0]
stream B: [GEMM chunk0] [dispatch chunk1] [combine chunk1]
stream A: [GEMM chunk1]
  • chunk0 在算的时候,chunk1 在通信
  • 需要 2 条 CUDA stream + event 同步
  • 代价:buffer 翻倍,GEMM tile 变小导致 Tensor Core 利用率下降

这就是 TBO (Two-Batch Overlap),SGLang 里有 --enable-two-batch-overlap 开关。

为什么传统方案"能 overlap 但不彻底"

四个硬限制:

  1. Kernel 边界 = 同步点
    每次 launch 有 5–10μs 开销,小 batch 下被 launch 开销主导

  2. SM 资源争抢
    Normal dispatch 想跑满 NVLink 就要用很多 SM,GEMM 也想要全部 SM,互相挤压

  3. 依赖链串行

    1
    dispatch(x) → GEMM1 → SwiGLU → GEMM2 → combine

    单个 token 的依赖必须串行,overlap 只能靠不同 micro-batch 之间的并行

  4. Decode 阶段 batch 小
    再切一半更糟,GEMM 效率暴跌

重叠率实测

方案 重叠率 适用场景
DeepEP normal 0%(串行) Prefill 大 batch
DeepEP low-latency + TBO 30–60% Prefill 中等 batch
MegaMoE 接近 100% 所有场景(包括 decode 小 batch)

对性能标定(profiler)的意义

如果在 perf 数据库里标注 MoE 路径,至少分三档:

  1. DeepEP normal:通信和计算串行(吞吐高,延迟差)
  2. DeepEP low-latency + TBO:部分 overlap(需要 enable_two_batch_overlap=true
  3. MegaMoE:原生 overlap,不依赖切 batch

这三档在 SLA 模型里会给出显著不同的 ITL/TTFT,混在一起标定会让 profiler 拟合出错。


核心实现机制

1. Symmetric Memory(对称内存)

传统 all-to-all:

1
GPU_0 → cudaMemcpyAsync → GPU_1(需要显式同步,退出 kernel)

Symmetric Memory:

1
2
3
所有 GPU 共享一片 symmetric buffer
GPU_0 直接写入 GPU_1 的 symmetric buffer(RDMA over NVLink)
GPU_1 看到数据就开始算,不需要全局 barrier

关键:不需要等所有 token 传完再开始算。传一批就算一批(streaming)。

实现层级

  • 硬件层:NVSwitch 全连接,任何 GPU pair 都有直达链路,memory controller 识别远端地址自动走 NVLink
  • 驱动/Runtime 层:CUDA VMM 把远端 GPU 物理内存映射到本地虚拟地址(cuMemMap / nvshmem_ptr
  • 应用层:MegaMoE 直接用普通指针读写远端 buffer,一条 PTX load/store 指令完成

注意:只有 symmetric allocation 的那块 buffer 地址一致,不是整个 GPU 显存对称。

2. 单 Kernel 状态机

传统方式多个 kernel 之间有全局同步点,无法实现 fine-grained overlap:

1
kernel_1 (dispatch) → kernel 边界(全局 barrier)→ kernel_2 (GEMM1) → ...

MegaMoE 的 scheduler/mega_moe.cuh 用一个状态机让所有 SM 在同一个 kernel 内自主领任务:

1
2
3
4
5
6
7
8
9
10
enum class BlockPhase { None, Linear1, Linear2 };

while (true) {
auto block = scheduler.get_next_block();
switch (block.phase) {
case Linear1: /* 做 W1/W3 GEMM */ break;
case Linear2: /* 做 W2 GEMM */ break;
case None: return;
}
}

没有 kernel 边界,没有全局 barrier,通信和计算可以交错到 warp 级别。

3. Wave-Based Expert Processing

256 个 expert 太多,无法同时处理(Token Pool 显存有限),分 wave 分批:

1
2
3
4
Wave 0: expert [0..31]   → 处理完 → 释放 pool
Wave 1: expert [32..63] → 处理完 → 释放 pool
...
Wave 7: expert [224..255]

每个 wave 内:

  • 所有 SM 抢 L1 block(gate + up projection)
  • SwiGLU 在 SMEM 中直接做激活
  • 所有 SM 抢 L2 block(down projection)

kNumExpertsPerWave 由 SMEM 容量、BLOCK_M、pool_capacity 共同决定,让 wave 内 expert 数刚好喂饱所有 SM 又不撑爆 Token Pool。

4. Arrival Count 轮询(细粒度 Overlap 的关键)

1
2
3
4
// 轮询等待:不是等所有 token 到齐,而是等够一个 BLOCK_M 就开始算
while (volatile_count < kNumSMs * kNumRanks) {
// spin-wait 直到其他 SM/rank 报告完成
}
  • 不需要全局 barrier(不需要所有 token 都收完才开始)
  • 收到一个 block 的 token 就立刻开始算
  • 这就是 streaming overlap 的本质

5. L1/L2 交错

1
2
Expert_0: [L1 block 0][L1 block 1][L1 block 2] → SwiGLU → [L2 block 0][L2 block 1]
Expert_1: [L1 block 0][L1 block 1] → SwiGLU → [L2 block 0]

L1 结果留在 Token Pool 中,L2 直接从 pool 读(在 L2 cache 中热),一个 wave 结束后才释放 pool 空间。

对应关系:

MegaMoE 术语 权重矩阵 操作 含义
L1 (Linear 1) W1 + W3(拼在一起) gate + up projection hidden → 2×intermediate
SwiGLU 无权重 SiLU(gate) × up 激活函数
L2 (Linear 2) W2 down projection intermediate → hidden

W1 和 W3 拼在一起做一次 GEMM,可以复用 TMA 加载的 activation tile,减少 GMEM → SMEM 搬运。


SM/Warp 级别并行

GPU 硬件层次:

1
2
3
4
GPU (整块芯片)
└── SM (Streaming Multiprocessor) × 132 (H20) / 192 (B200)
└── Warp × 最多 64 resident / SM
└── Thread × 32 / Warp(SIMT 锁步执行)

在 MegaMoE 中,一个 Mega Kernel launch 占满所有 SM,每个 SM 内部:

1
2
Warp Group A(若干 warp):负责通信(polling + NVLink store/load)
Warp Group B(若干 warp):负责计算(Tensor Core MMA)

Warp Group A 和 B 在 SM 内部并行执行——这就是"通信和计算在 warp 级别交错"的含义。


为什么必须 GB200/GB300

MegaMoE 依赖三个 GB200/GB300 独有(或大幅增强)的硬件能力:

硬件特性 GB200/GB300 H20 影响
GPU 间拓扑 72 GPU NVSwitch 全连接 8 GPU NVLink mesh H20 无法 symmetric memory 跨机
Symmetric Memory ✅ kernel 内 load/store 远端 ❌ 需要 NCCL API 无法做 in-kernel 通信
NVSwitch Reduction ✅ 硬件完成 combine ❌ SM 做 reduce combine 占 SM 资源
FP4 Tensor Core ✅ 原生指令(SM100) ❌ 软件模拟(SM90) 计算慢 → 通信占比低 → overlap 收益有限
NVLink 带宽/GPU 900 GB/s(18 ports) 450 GB/s(NV18 mesh) H20 通信带宽是瓶颈

为什么计算足够快才值得 overlap

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
假设: 通信时间 = T_comm, 计算时间 = T_compute

情况 A — 计算快 (FP4 Tensor Core):
T_compute = 2ms, T_comm = 5ms
不 overlap: 7ms
overlap: max(2, 5) = 5ms ← 省 28%

情况 B — 计算慢 (BF16 软件模拟):
T_compute = 20ms, T_comm = 5ms
不 overlap: 25ms
overlap: max(20, 5) = 20ms ← 只省 20%

而且 overlap 不是免费的:
- 需要 SM 专门做通信 warp(减少计算并行度)
- Symmetric Memory 的 load/store 比本地慢 10-20x
- Scheduler 状态机有开销

只有当计算非常快、通信成为瓶颈时,overlap 的净收益才大于开销。

和 DeepEP / flashinfer_mxfp4 的区别

MegaMoE vs DeepEP

DeepEP MegaMoE
通信粒度 整个 batch dispatch 完再算 tile 级别通信和计算交错
overlap 方式 不同 CUDA stream 异步(kernel 间) 同一个 kernel 内 tile 级 pipeline
权重格式 任意 Runner backend 都能接 只支持 DeepGEMM 格式(权重需预转换)
内存模型 普通 GPU buffer NVSHMEM 对称内存 (SymmBuffer)
batch 上限 无(取决于显存) 有(NUM_MAX_TOKENS_PER_RANK

MegaMoE vs flashinfer_mxfp4

flashinfer_mxfp4 MegaMoE
解决的问题 单卡 MoE 计算效率 多卡 EP 通信+计算 overlap
Fuse 范围 GEMM1 + SwiGLU + GEMM2 dispatch + GEMM1 + SwiGLU + GEMM2 + combine
通信 不涉及(单卡或假设已 gather 好) 融合 All-to-All 通信
适用场景 TP 模式(experts 复制在每卡) EP 模式(experts 分布在不同卡)
硬件要求 任何 Hopper/Blackwell 需要 NVLink + Symmetric Memory(GB200+)
状态 已发布,可用 开发中(DeepGEMM PR #304)

两者互补:flashinfer_mxfp4 管单卡 MoE kernel 效率,MegaMoE 管 EP 多卡通信+计算融合。


Python API 调用流程

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
# 1. 多进程初始化 + 分配对称内存
buffer = deep_gemm.get_symm_buffer_for_mega_moe(
group, # 进程组 (torch.distributed)
num_experts=256,
num_max_tokens_per_rank=8192,
num_topk=6,
hidden=4096,
intermediate_hidden=2048
)

# 2. 权重预处理(FP4 packing + 布局变换)
l1_transformed, l2_transformed = deep_gemm.transform_weights_for_mega_moe(
l1_weights, # W1, W3 (gate + up)
l2_weights # W2 (down)
)

# 3. 加载输入到 symmetric buffer
buffer.x[:N].copy_(x_fp8) # FP8 激活
buffer.x_sf[:N].copy_(x_sf) # 激活的 scale factor
buffer.topk_idx[:N].copy_(topk_idx) # 路由结果
buffer.topk_weights[:N].copy_(topk_weights)

# 4. 一次调用,完成 dispatch + GEMM1 + SwiGLU + GEMM2 + combine
y = torch.empty((N, 4096), dtype=torch.bfloat16, device='cuda')
deep_gemm.fp8_fp4_mega_moe(y, l1_transformed, l2_transformed, buffer)
# y 就是最终结果,全程只有这一个 kernel launch

总结:MegaMoE 的关键创新

机制 作用 依赖硬件
Symmetric Memory kernel 内直接 load/store 远端 GPU NVSwitch 全连接
Arrival Count 轮询 收到一个 block 就开始算,不等全部 低延迟 NVLink
Wave-based Scheduler SM 自主领任务,无全局 barrier SM 数量充足(Blackwell 192 SM)
L1/L2 交错 L1 结果留在 pool,L2 直接读 大 L2 cache(Blackwell 100MB+)
FP8×FP4 GEMM 计算足够快,才值得 overlap 通信 原生 FP4 Tensor Core(SM100)

一句话:MegaMoE = Symmetric Memory(消除通信 barrier)+ 细粒度 Scheduler(一个 block 到了就算)+ 单 Kernel 状态机(避免 kernel launch 开销),把 EP MoE 的 GPU 利用率从 ~60% 推到接近 100%。


局限性与适用边界

  • 权重格式锁定:只支持 DeepGEMM 的 FP4 权重布局,需要预转换
  • 激活函数限制:依赖 SwiGLU 无状态特性,L1 结果可直接给 L2 用;若 MoE 层有跨 token 依赖(如全局 norm),pipeline 会断
  • 硬件绑定:Symmetric Memory + NVSwitch 全连接 + 原生 FP4 三个条件缺一不可,H20 及更早硬件无法使用
  • 跨机不支持:Symmetric Memory 只在单机 NVSwitch 域内有效,跨机仍需回退到 DeepEP
  • 负载不均影响fetch_expert_recv_count() 的 spin-wait 在 expert 负载不均时可能导致 SM 空转

参考

背景

DeepSeek-V4 系列模型(V4 / V4-Pro)在 MoE(Mixture of Experts)层大量使用 MXFP4 量化权重,配合 FP8/BF16 激活 实现高效推理。在 SGLang / vLLM 中,flashinfer_mxfp4 是一个专为这种量化格式设计的 MoE Runner Backend。

本文将基于 H20 (SM90 Hopper) 实测数据和源码分析,完整解析 FlashInfer MXFP4 的实现原理,以及与 Marlin MoE 的架构差异。


MXFP4 量化格式

格式定义

MXFP4 是 OCP(Open Compute Project)标准的 4-bit 浮点格式:

1
2
每个权重元素: [1 sign | 2 exponent | 1 mantissa] → 4-bit 浮点数
每 32 个元素: 共享一个 8-bit 指数缩放因子 (E8M0 block scale)

与 INT4 的对比

维度 MXFP4 INT4 (Marlin)
表示方式 浮点 (E2M1) 定点整数
动态范围 大(指数缩放) 小(线性,需精确校准)
缩放粒度 每 32 元素(block scale) 每 128 元素(group quantization)
解量化 value = fp4 × 2^(scale - 127) value = int4 × scale[channel]

优势:浮点表示让 MXFP4 在训练和推理中都有更好的数值稳定性。


FlashInfer MXFP4 架构

调用链(H20 / SM90)

1
2
3
4
5
6
7
8
9
10
11
SGLang --moe-runner-backend flashinfer_mxfp4

Python: trtllm_fp4_block_scale_moe()

C++: get_cutlass_fused_moe_module() → 加载预编译 .so

CUTLASS 3.x: cute::gemm::kernel::DefaultGemmUniversal

├─ TMA (Tensor Memory Accelerator) → 异步搬权重 tile 到 SMEM
├─ WGMMA (Warpgroup MMA) → SM90 用 FP16 TC 模拟 FP4
└─ Warp Specialization → Producer warpgroup 专搬数据,Consumer warpgroup 专计算

核心特性

  1. Expert-First 遍历:按 expert 分组 token,而非按 token 遍历 expert
  2. 全 Fuse:W1 + W3 + SwiGLU + W2 四个步骤 fuse 成一个 kernel
  3. TMA + Warp Specialization:数据搬运和计算 overlap
  4. Grouped GEMM:一次 launch 处理所有 256 个 expert

完整 Pipeline(DeepSeek-V4 为例)

模型配置

1
2
3
4
5
DeepSeek-V4-Flash:
- 256 experts, topk=6
- hidden_dim=4096, moe_intermediate=2048
- 权重: MXFP4 (4-bit), 激活: BF16
- Router: Sigmoid + Group TopK (非 Softmax)

第一阶段:Routing(路由选择)

1
2
3
4
5
6
7
8
9
10
11
12
# 文件: trtllm_fused_moe_routing_deepseek.cu

# 输入: router_logits [num_tokens, 256]
scores = torch.sigmoid(router_logits + bias) # DeepSeek-V4 用 Sigmoid, 非 Softmax

# Group TopK: 256 experts → 8 组 × 32
group_scores = scores.view(num_tokens, 8, 32).max(dim=-1) # 每组最高分
top_groups = group_scores.topk(4).indices # 选 4 组 (128 experts)

# 从 128 个 expert 中选 top-6
topk_indices = ... # [num_tokens, 6]
topk_weights = ... # 归一化后的权重

为什么用 Sigmoid?

  • 独立评分(非 Softmax 的零和游戏)
  • 多 expert 可同时激活
  • 训练更稳定

第二阶段:Gather(Token 重排)

1
2
3
4
# 文件: cutlass_fused_moe_kernels.cuh → expandInputRowsKernel

# 输入: hidden_states [num_tokens, 4096] (按 token 顺序)
# 输出: expanded_input [num_tokens × 6, 4096] (按 expert 连续排列)

目的:把路由到同一个 expert 的 token 收集到一起,形成连续内存,供后续 GEMM 使用。

1
2
3
4
5
6
7
8
9
10
原始 (token 顺序):
[t0 | t1 | t2 | t3]

路由结果:
expert_3: [t0, t1, t3]
expert_7: [t0, t2]

Gather 后 (expert 顺序):
[e3: t0, t1, t3 | e7: t0, t2 | ...]
↑ 连续 ↑ 连续

TMA 优势:连续输入让 TMA 一次加载整个 tile,达到满带宽。


第三阶段:Compute(Fused GEMM)

1
# 文件: cutlass_fused_moe_kernels.cuh → CUTLASS Grouped GEMM

这是最核心的部分,W1 + W3 + SwiGLU + W2 全部 fuse

Step 3a: GEMM1 (Gate + Up Projection)

1
2
3
4
5
6
对每个 expert_i:
M = expert_i 分到的 token 数 (例如 3)
input_i = expanded_input[offset : offset+M] # [M, 4096]

# W13 = [W1; W3] 拼接,一次 GEMM 算出
gate_up = input_i @ W13_expert_i # [M, 4096] → split → [M, 2048] + [M, 2048]

MXFP4 权重存储

1
2
3
// W1_expert_i: stored as [2048, 2048] int8
// 实际逻辑形状: [2048, 4096] FP4 (每 byte 存 2 个 FP4 元素)
// Scale: [2048, 64] float8_e8m0 (每 32 个元素一个 scale)

Step 3b: SwiGLU Activation

1
2
3
# 在 GEMM1 输出还在 SMEM/寄存器时立即执行(fuse 关键)
intermediate = SiLU(gate) * up # [M, 2048]
# DeepSeek-V4 还有 swiglu_limit=10.0 的 clamp

Step 3c: GEMM2 (Down Projection)

1
2
output_i = intermediate @ W2_expert_i  # [M, 4096]
# W2_expert_i: stored as [4096, 1024] int8 → 逻辑 [4096, 2048] FP4

CUTLASS Grouped GEMM 调度

1
2
3
4
5
6
7
8
9
10
11
12
// 一次 kernel launch 处理所有 256 个 expert
cutlass::gemm::kernel::GroupedGemmKernel<...>::run();

// 内部调度:
// SM 空闲 → 从 problem list 取下一个 expert
// expert_i 的 M=0 (没 token) → 跳过
// expert_j 的 M=12 → 分配 SM 算 GEMM

// TMA + Warp Specialization:
// Warpgroup 0 (Producer): TMA_LOAD(weight_tile, desc)
// Warpgroup 1 (Consumer): WGMMA(tile_C, tile_A, tile_B)
// Producer 搬下一个 expert 的权重时,Consumer 正在算当前 expert

Pipeline 示意(H20)

1
2
3
4
Cycle 0-100:  Producer 搬 expert_3 权重,Consumer 算 expert_1 (上一轮)
Cycle 100-200: Consumer 算 expert_3 (权重已到),Producer 搬 expert_7
Cycle 200+: Consumer 算 expert_7,Producer 搬 expert_5
→ 数据搬运和计算完全 overlap!

第四阶段:Scatter(结果写回)

1
2
3
4
# 文件: cutlass_fused_moe_kernels.cuh → finalizeMoeRoutingKernel

# 输入: expert_outputs [num_tokens × 6, 4096] (按 expert 排列)
# 输出: final_output [num_tokens, 4096] (恢复 token 原始顺序)

加权求和

1
2
3
for token_j:
output[j] = Σ (weight_k × expert_output[permute_map[j][k]])
k=1..6

示例

1
2
token_0 → expert_3 (0.6), expert_7 (0.4)
output[0] = 0.6 × out_0_e3 + 0.4 × out_0_e7

为什么叫 Scatter?

  • expert_3 的输出: [out_0_e3, out_1_e3, out_3_e3]
  • 需要写回: output[0], output[1], output[3](不连续!)
  • 这就是 scatter(分散写入),与 gather(聚集读取)相对

性能实测(H20, DeepSeek-V4-Flash, TP=4)

Benchmark 结果

Concurrency FlashInfer MXFP4 Marlin 差距
conc=1 20.2 tok/s/GPU 20.5 tok/s/GPU ≈持平
conc=2 42.1 tok/s/GPU 28.1 tok/s/GPU +50%
conc=4 32.5 tok/s/GPU 31.8 tok/s/GPU ≈持平

阶段耗时分解(估)

阶段 占比 说明
Routing ~2% 纯 element-wise,快
Gather ~3% 内存拷贝
GEMM1 (W13) ~42% 计算密集
Activation ~1% fused 在 GEMM1 后
GEMM2 (W2) ~48% 计算密集
Scatter ~4% 加权求和 + 写回

GEMM 占 90%,所以 TMA + Grouped GEMM 对总体性能影响最大。

为什么 conc=2 差距最大?

  1. Expert-First + Grouped GEMM:一次 launch 处理所有 expert,TMA 预取下一个 tile
  2. 中等 batch:每个 expert 分到 2-3 个 token → GEMM tile 够大,WGMMA 饱和
  3. Warp Specialization:Producer 搬数据,Consumer 计算,完美 overlap

Marlin 的劣势

  • Token-First:每个 token 的 topk expert 不同 → 权重反复换入换出
  • 无 TMA:用 cp.async 加载,线程要参与地址计算 → SM 利用率 < 50%
  • 无法全 fuse:中间结果 (gate, up, mid) 要写回 GMEM

为什么 conc=4 差距消失?

1
2
3
4
5
conc=4: prefill=50k tokens
→ Prefill 时间: ~1500ms (占 85%+)
→ Decode MoE 时间: ~265ms (占 15%-)

即使 MoE 快 50%,总体提升也只有: 265ms × 50% / 1765ms ≈ 7.5%

Prefill 主导后,MoE decode 的优化被稀释。


与 Marlin MoE 的架构对比

核心差异

维度 FlashInfer MXFP4 Marlin (SGLang)
量化格式 MXFP4 (浮点4位) INT4 / MXFP4
遍历顺序 Expert-First Token-First
融合程度 W1+W3+SwiGLU+W2 全 fuse W1+W3 部分 fuse,W2 单独
数据加载 TMA (硬件 DMA) cp.async / LDG (线程参与)
调度方式 Warp Specialization (Producer/Consumer) 单 warp 既搬又算
Kernel 架构 CUTLASS 3.x Grouped GEMM 手写 CUDA,Ampere 设计
TMA 利用 ✅ 异步 tile 预取 ❌ 无 TMA (用 cp.async)
多 expert 并行 Grouped GEMM 一次 launch 逐 expert 串行或小并行

Expert-First vs Token-First

Token-First (Marlin)

1
2
3
4
5
for token_0:
expert_3: W1 → W3 → SwiGLU → W2 ← 写回 GMEM
expert_7: W1 → W3 → SwiGLU → W2 ← 换权重
for token_1:
... # 重复上述,权重反复换入换出

Expert-First (FlashInfer)

1
2
3
4
5
6
7
for expert_3:
# 固定权重 W1/W3/W2,常驻 SMEM
gate_up = input_e3 @ W13_e3 # 连续 token 一起算
mid = SiLU(gate) * up # 在寄存器/SMEM 直接算
output = mid @ W2_e3 # 中间结果不落地
for expert_7:
... # 换一次权重,算所有路由到它的 token

为什么 Expert-First 能全 fuse?

  • 固定权重 → W1/W3/W2 常驻 SMEM,不换出
  • 变 batch → 同一 expert 的多个 token 连续处理
  • 中间结果 (gate, up, mid) 全在寄存器/SMEM,不写 GMEM

TMA + WGMMA 的硬件优势

TMA (Tensor Memory Accelerator)

SM90 (Hopper) 引入的硬件 DMA 引擎

  • 异步搬数据从 HBM 到 SMEM,不占用 CUDA core
  • 一条 tcgen05.1d 指令触发,后台独立运行
  • 支持多维 tensor 描述符(TMA Descriptor)

vs 传统 LDG

1
2
3
4
5
6
7
8
LDG (传统):
线程发 load → 等 HBM (~400 cycles) → 数据到 SMEM → 继续算
→ 线程 stall,SM 有空闲

TMA:
线程发 TMA_LOAD → 立即返回 → 可以去 setup 下一个 tile
→ TMA 后台搬,线程继续算上一轮结果
→ SM 利用率 80%+

WGMMA (Warpgroup MMA)

SM90 引入的 warpgroup 级矩阵乘指令

  • 一个 warpgroup = 4 个 warp = 128 线程
  • 直接读写 TMEM(Tensor Memory,SM90 新增寄存器文件)
  • 比传统 mma.sync 高一个抽象层级

SM100 (Blackwell) 的进化

1
2
SM90: TMA → SMEM → WGMMA → TMEM
SM100: TMA_GATHER4 → TMEM (跳过 SMEM) → UTCMMA (原生 FP4 TC)

Blackwell 的 TMA_GATHER4 还能一次加载 4 个不连续 token(专为 MoE 的 sparse routing 优化)。


总结

FlashInfer MXFP4 的核心优势

  1. Expert-First + 全 Fuse:中间结果不落地 GMEM,省 ~2× hidden_dim × batch 带宽
  2. TMA + Warp Specialization:搬运和计算 overlap,SM 利用率最大化
  3. Grouped GEMM:一次 launch 处理 256 experts,launch overhead 最小化
  4. Blackwell 未来兼容:cute_dsl_fused_moe_nvfp4 已支持 SM100 原生 FP4 TC

适用场景

场景 推荐后端
H20/H100, conc=2~8 ✅ FlashInfer MXFP4
A100, 或 conc=1 调试 Marlin
Blackwell (B200) FlashInfer NVFP4 (cute_dsl)
追求最大吞吐 FlashInfer MXFP4

一句话

FlashInfer MXFP4 在 Hopper 上通过 Expert-First + TMA + Warp Specialization + 全 Fuse,把 MoE 推理的瓶颈从内存带宽转移到了计算,实现了中等并发下 50% 的吞吐提升


参考资料


如果你对 MoE 推理优化感兴趣,可以看看我之前写的 Marlin MoE Kernel 深度分析,对比两种实现的差异。