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 = 1x=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 外积。

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_2k2k_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]^\topk2=[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]

  • 线性注意力S4k1=[3.0,0.6]S_4 k_1 = [3.0, 0.6],新旧 value 相互叠加,并混入 0.6 份 v2v_2 的残余;
  • delta ruleS4k1S_4 k_1S4=[2301]S_4 = \begin{bmatrix}2&3\\0&1\end{bmatrix} 的第一列 =[2.0,0.0]= [2.0, 0.0]——精确返回最新的 v3=[2,0]v_3 = [2, 0]。这一步值得展开看:拿 t=4 的 S4S_4 右乘 k1k_1k1k_1 只命中第一列(第一行的 2 来自 t=3 写入的 v3v_3,第二行的 0 恰好是 t=3 差值修正把第一列第二行清零的结果——t=2 混进来的 0.6 在 t=3 被 u3u_3 的第二分量 0.6-0.6 抵消掉了)。查询、写入用的同一个 key,写入时已保证 Sk3v3S k_3 \approx v_3,所以查询原样取回新值。

delta rule 起作用的机制:写入时先查旧值、把误差算出来再写。旧记忆里与新写入相关的部分(包括之前叠加污染的分量)被差值中的负项抵消掉,相当于先删除旧 key 上的记忆、再写入新 value,从而保证当前 key 映射到的是完整且无残余的 value。同一个 key 写两次,线性注意力读出的是两个 value 的叠加,delta rule 读出的是新的那个——这就是「记忆碰撞」(memory collision)和它的修复方式。

到这里,DeltaNet 的雏形就有了:固定大小的外积状态 + 「先删旧、再写新」的更新规则。后面 DeltaNet/GDN/KDA 做的事,都是在雏形上加约束(如 key 归一化到单位球)、加门控(衰减系数 α)和解决并行训练的问题。


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_kt=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 = -1B=2B = 2C=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 Ia>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. 出发点:固定大小的记忆

标准 softmax 注意力的 KV cache 随序列长度线性膨胀,几十万 token 的上下文可使其达到数十 GB 量级,显存与带宽同时成为瓶颈。K3 把「序列维度缩放」列为头号工程目标,93 层里 69 层换成 KDA,就是把 KV cache 换成一个固定大小的状态矩阵O(nd)O(n \cdot d)O(d2)O(d^2),上下文长度翻倍时状态大小保持不变。

但固定大小是有代价的。一个 dv×dkd_v \times d_k 的矩阵要承载任意长度的序列,记忆必须可写、可改、可遗忘:若只能写入而不能清理,序列一长信息就会相互干扰。如何设计这块「可自我管理的固定记忆」,是本文的主线。

线性注意力与 SSM 两条路线的完整推导见前置篇《线性注意力与 SSM:两条技术路线的完整推导》,本文直接使用其结论。

记号约定:两套写法互为转置

先把记号钉死。文献里状态矩阵有两种摆法,内容完全等价、互为转置,但混用会让推导看起来「中途换了个式子」:

主约定(§1.1 与 §3 起使用,与 KDA 论文一致) 转置约定(部分实现与 chunkwise 推导常用)
状态形状 StRdk×dvS_t \in \mathbb{R}^{d_k \times d_v} S^tRdv×dk\hat S_t \in \mathbb{R}^{d_v \times d_k}
读出 ot=Stqto_t = S_t^\top q_t ot=S^tqto_t = \hat S_t q_t
当前 key 的旧内容 St1ktS_{t-1}^\top k_t S^t1kt\hat S_{t-1} k_t
写入项 ktvtk_t v_t^\top(左 key 右 value) vtktv_t k_t^\top(左 value 右 key)
擦除算子 左乘:(Iβtktkt)St1(I - \beta_t k_t k_t^\top)S_{t-1} 右乘:S^t1(Iβtktkt)\hat S_{t-1}(I - \beta_t k_t k_t^\top)

两者的严格关系就是一次转置 S^t=St\hat S_t = S_t^\top。由于擦除算子 Ht=IβtktktH_t = I - \beta_t k_t k_t^\top 是对称矩阵Ht=HtH_t^\top = H_t,因为 (ktkt)=ktkt(k_tk_t^\top)^\top = k_tk_t^\top),转置可以直接穿过它:

(HtSt1+βtktvt) ⁣=St1Ht+βtvtkt=S^t1Ht+βtvtkt\big(H_t S_{t-1} + \beta_t k_t v_t^\top\big)^{\!\top} = S_{t-1}^\top H_t^\top + \beta_t v_t k_t^\top = \hat S_{t-1} H_t + \beta_t v_t k_t^\top

所以后文若看到 kvk v^\topvkv k^\top 互换、擦除算子从左边跑到右边,那是切换了约定,不是等式变了。GDN 数值验算与 chunkwise 一节为对齐论文公式会改用转置约定,届时会再次点明。

各模型按出场顺序(统一写成主约定):

  • 线性注意力:固定状态 St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top,只支持累加写入(同一个 key 写两次,读出的是两个 value 的叠加,即记忆碰撞);
  • Mamba-2:引入标量衰减门 St=αtSt1+ktvtS_t = \alpha_t S_{t-1} + k_t v_t^\top,具备遗忘能力,但衰减是全局统一的;
  • DeltaNet:改为差值写入 St=(Iβtktkt)St1+βtktvtS_t = (I - \beta_tk_tk_t^\top)S_{t-1} + \beta_t k_t v_t^\top,支持定点改写;
  • GDN:衰减门与差值写入结合 St=αt(Iβtktkt)St1+βtktvtS_t = \alpha_t(I - \beta_tk_tk_t^\top)S_{t-1} + \beta_t k_t v_t^\top
  • KDA:将标量门拆为逐通道门 Diag(αt)\operatorname{Diag}(\bm{\alpha}_t),各维度独立决定衰减速度,并加数值下界,用于 K3。

1. 演进主线:从 delta rule 到 GDN

前置篇给出两个结论:线性注意力提供了固定状态,但只能累加写入、存在记忆碰撞;SSM 推进到 Mamba-2 提供了标量衰减门 αt\alpha_t,但衰减是全局统一的。本节讨论这两者如何组合成 GDN(Gated DeltaNet)——先看写入规则怎么从加法变成差值(§1.1),再把遗忘门拼进来(§1.2),然后手算验证(§1.3)并给出可并行的 chunkwise 形式(§1.4)。KDA 对 GDN 的改动从 §3 开始。

1.1 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+ktutS_t = S_{t-1} + k_t u_t^\top

符号 形状 含义
StS_t [dk,dv][d_k, d_v] 状态矩阵(主约定)
ktk_t [dk][d_k] 当前 key,经 L2 归一化后 kt2=1|k_t|_2 = 1
voldv_{\mathrm{old}} [dv][d_v] 当前 key 从状态里读出的已有内容
βt\beta_t 标量 (0,1)\in (0,1) 内容替换强度
utu_t [dv][d_v] 实际写入的差值

严格改写:差值写入 = 先删后写。上式右侧看不出「擦除」在哪里,把 utu_t 代入展开即可,每一步只用外积的结合律 kt(ktSt1)=(ktkt)St1k_t(k_t^\top S_{t-1}) = (k_tk_t^\top)S_{t-1}

St=St1+ktut定义=St1+kt[βt(vtvold)]代入 ut=St1+βtktvtβtktvold转置展开=St1+βtktvtβtkt(St1kt) ⁣代入 vold=St1kt=St1+βtktvtβtktktSt1(St1kt)=ktSt1=(Iβtktkt)St1删除项:沿 kt 方向擦除+βtktvt新值项提取公因式 St1\begin{aligned} S_t &= S_{t-1} + k_t u_t^\top && \text{定义} \\[2pt] &= S_{t-1} + k_t\big[\beta_t(v_t - v_{\mathrm{old}})\big]^\top && \text{代入 } u_t \\[2pt] &= S_{t-1} + \beta_t k_t v_t^\top - \beta_t k_t v_{\mathrm{old}}^\top && \text{转置展开} \\[2pt] &= S_{t-1} + \beta_t k_t v_t^\top - \beta_t k_t \big(S_{t-1}^\top k_t\big)^{\!\top} && \text{代入 } v_{\mathrm{old}} = S_{t-1}^\top k_t \\[2pt] &= S_{t-1} + \beta_t k_t v_t^\top - \beta_t k_t k_t^\top S_{t-1} && (S_{t-1}^\top k_t)^\top = k_t^\top S_{t-1} \\[2pt] &= \underbrace{\big(I - \beta_t k_t k_t^\top\big) S_{t-1}}_{\text{删除项:沿 } k_t \text{ 方向擦除}} + \underbrace{\beta_t k_t v_t^\top}_{\text{新值项}} && \text{提取公因式 } S_{t-1} \end{aligned}

关键的一步是第五行:voldv_{\mathrm{old}} 自己就是由 St1S_{t-1} 算出来的,所以「减去旧值」这个动作必然能写成一个作用在 St1S_{t-1} 上的线性算子,而不是一个额外的加项。提取公因式之后,βtktkt-\beta_t k_tk_t^\top 就并入了单位阵,变成擦除算子。预测误差 vtvoldv_t - v_{\mathrm{old}} 就是 delta,Delta Rule 由此得名:每次写入的从来不是新值本身,而是新值与旧值之差。

两端均为 dk×dvd_k \times d_v,维度自洽;若换成转置约定,同一个式子写作 S^t=S^t1(Iβtktkt)+βtvtkt\hat S_t = \hat S_{t-1}(I - \beta_tk_tk_t^\top) + \beta_t v_tk_t^\top

写入差值这一形式并非经验设计,它等价于对回归损失做一步梯度下降。上面是先有「写差值」这个想法、再发现它等于先删后写;反方向走一遍会看到,差值根本不是选的,而是推出来的。

第一步:把记忆当成一个在线回归问题。状态 SS 要承担的职责是一张查询表:拿钥匙 ktk_t 来,应当取出 vtv_t,即希望 SktvtS^\top k_t \approx v_t。把这个愿望写成当前样本上的平方损失:

L(S)=12Sktvt2\mathcal{L}(S) = \tfrac{1}{2}\big\|S^\top k_t - v_t\big\|^2

SktS^\top k_t 是「用现在的记忆查 ktk_t 能取出的值」,vtv_t 是「应该取到的值」,两者之差就是残差;系数 12\tfrac12 纯为求导后消掉 2。

第二步:求梯度,得到残差与钥匙的外积。先用最熟的一维情形建立直觉:12(wxy)2\tfrac12(wx-y)^2ww 的导数是 (wxy)x(wx-y)\,x——误差乘输入。矩阵版一模一样,只是乘法变成外积。记残差 rt=SktvtRdvr_t = S^\top k_t - v_t \in \mathbb{R}^{d_v},取扰动 Δ\Delta 算一阶项:

L(S+Δ)L(S)=rt, Δkt+O(Δ2)=ktrt, Δ+O(Δ2)\mathcal{L}(S+\Delta) - \mathcal{L}(S) = \big\langle r_t,\ \Delta^\top k_t \big\rangle + O(\|\Delta\|^2) = \big\langle k_t r_t^\top,\ \Delta \big\rangle + O(\|\Delta\|^2)

Δ\Delta 配对的那个矩阵就是梯度:

SL=kt(Sktvt) ⁣=ktrt  Rdk×dv\nabla_S \mathcal{L} = k_t\,\big(S^\top k_t - v_t\big)^{\!\top} = k_t r_t^\top \ \in\ \mathbb{R}^{d_k \times d_v}

形状与 SS 一致,可直接用于更新。注意差值 SktvtS^\top k_t - v_t 是自己冒出来的——平方损失的梯度天然就长成「残差 ⊗ 钥匙」的样子。这就是「为何恰好是差值」的答案:没人规定写入要用差值,是平方损失的梯度只能是差值。

第三步:以 βt\beta_t 为步长走一步 SGD,展开重新归项:

St=St1βtSLS=St1一步梯度下降=St1βtkt(St1ktvt) ⁣代入梯度=St1βtktktSt1+βtktvt展开=(Iβtktkt)St1+βtktvt提公因式\begin{aligned} S_t &= S_{t-1} - \beta_t \nabla_S\mathcal{L}\big|_{S = S_{t-1}} && \text{一步梯度下降} \\[2pt] &= S_{t-1} - \beta_t k_t\big(S_{t-1}^\top k_t - v_t\big)^{\!\top} && \text{代入梯度} \\[2pt] &= S_{t-1} - \beta_t k_t k_t^\top S_{t-1} + \beta_t k_t v_t^\top && \text{展开} \\[2pt] &= \big(I - \beta_t k_t k_t^\top\big)S_{t-1} + \beta_t k_t v_t^\top && \text{提公因式} \end{aligned}

结果与前面从「写差值」出发得到的式子逐字相同。两条路径交汇于同一个等式,于是 βt\beta_t 的角色也明确了:它就是学习率。同样,「先删后写」不是设计直觉,而是展开式里必然的两项:βtktkt-\beta_t k_tk_t^\top 并入单位阵成为擦除,+βtktvt+\beta_t k_tv_t^\top 就是写入。

验证:写完立即读一次。取 kt2=1\|k_t\|_2 = 1(L2Norm 之后),用同一个 ktk_t 查新状态:

Stkt=[(Iβtktkt)St1+βtktvt] ⁣kt=(1βt)保留旧值St1kt+βt接纳新值vtS_t^\top k_t = \big[(I-\beta_tk_tk_t^\top)S_{t-1} + \beta_tk_tv_t^\top\big]^{\!\top}k_t = \underbrace{(1-\beta_t)}_{\text{保留旧值}}\,S_{t-1}^\top k_t + \underbrace{\beta_t}_{\text{接纳新值}} v_t

读出结果是旧值与新值的凸组合βt\beta_t 就是插值系数:βt=1\beta_t = 1Stkt=vtS_t^\top k_t = v_t,该样本的损失一步降到 0,即完全覆写;βt=0.5\beta_t = 0.5 时新旧各一半;βt=0\beta_t = 0 时不学。这也解释了 L2Norm 为何是前提:只有 kt=1\|k_t\| = 1 时上式才是干净的插值,否则系数会变成 1βtkt21 - \beta_t\|k_t\|^2,可能跌出 [0,1][0,1]

一句话总结这条链:记忆 = 在线回归 → 平方损失 → 梯度 = 残差 ⊗ 钥匙 → 一步 SGD = 先删后写。它有严格出处:Widrow-Hoff 1960 年的 delta rule / LMS 规则(名字里的 delta 正是指误差 δ\delta)在联想记忆矩阵上的应用。后面 GDN 与 KDA 做的事,就是在这个损失里再加一项正则,把「别忘了旧账」量化进目标(见 §1.2 的在线学习统一视角)。

DeltaNet 只改写入规则,两个前置条件一律保留:特征映射 ϕ\phi(实现里常取 ϕ=L2Norm\phi = \mathrm{L2Norm},即下文的单位球约束)、外积状态、逐步递推形式完全没动,变的只有一处:+kv+k(vSk)+\,k v^\top \longrightarrow +\,k(v - S^\top k)^\top

直观类比:state 是白板,key 是指针。纯加性写入相当于在白板上不断叠加便签,写满后内容互相遮盖;DeltaNet 先擦除指针 ktk_t 指向的区域,再写入新内容。βt=0\beta_t = 0:完全不写入,状态不动;βt=1\beta_t = 1:完全替换,指针指向的内容被整体替换。这种「先擦后写」的机制使状态能够覆盖错误记忆,这是 DeltaNet 与 GDN 这条路线的核心改进。

关键整理:写成转移矩阵形式。定义 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)。

HtH_t 到底是什么:秩 1 与特征值

Ht=IβtktktH_t = I - \beta_t k_t k_t^\top 是全文出现频率最高的矩阵,值得把它彻底拆开。它由两块拼成:单位阵,减去一个秩 1 矩阵

先看「秩」。矩阵的秩 = 它的列里真正不同的方向有几个,形式定义是线性无关列的最大个数,直觉是这个矩阵作为变换、输出能铺满几维空间。单位阵 IRd×dI \in \mathbb{R}^{d\times d}dd 列是 dd 根坐标轴,谁也不沾谁,秩 =d= d(满秩);而 kkk k^\top 的所有列都挤在同一条线上,秩 =1= 1

为什么 kkkk^\top 的列都是 kk 的倍数kkkk^\top(i,j)(i,j) 元素是 kikjk_ik_j,于是它的第 jj 列是

(k1kjk2kjkdkj)=kj(k1k2kd)=kjk\begin{pmatrix} k_1 k_j \\ k_2 k_j \\ \vdots \\ k_d k_j \end{pmatrix} = k_j \cdot \begin{pmatrix} k_1 \\ k_2 \\ \vdots \\ k_d \end{pmatrix} = k_j \cdot k

jj 列 = 标量 kjk_j 乘同一个向量 kk,换 jj 只换倍率、方向永远是 kk。取 k=(1,2)k = (1,2)^\top 验算:

kk=(12)(12)=(1224)k k^\top = \begin{pmatrix}1\\2\end{pmatrix}\begin{pmatrix}1&2\end{pmatrix} = \begin{pmatrix}1&2\\2&4\end{pmatrix}

第二列 (2,4)(2,4) 恰是第一列 (1,2)(1,2) 的 2 倍——形式上是 2×22\times2 矩阵,实际只携带一个方向的信息,故行列式为 0、不可逆。反过来也成立:任何秩 1 矩阵都能写成某个外积 uvuv^\top,所以「秩 1 矩阵」与「外积」基本同义。

作为变换,它把一切压到 kk 那条线上。作用在任意向量 xx 上,用结合律:

(kk)x=k(kx)=(kx)k(k k^\top)\, x = k\,(k^\top x) = (k^\top x)\cdot k

kxk^\top x 是标量,即 xxkk 方向的投影长度;乘回 kk 得到一个沿 kk 的向量。不管输入什么,输出永远落在 kk 张成的一维直线上——这就是「只张出一维」的几何含义,整个 dd 维空间被拍扁成一条线。

再看「特征值」。若存在非零向量 xx 使 Mx=λxMx = \lambda x,即 xx 经变换后方向不变(或恰好反向)、只被缩放 λ\lambda 倍,则 xx特征向量(eigenvector)、λ\lambda特征值(eigenvalue)。多数向量过一个矩阵会又转又缩,特征向量是躺在矩阵「主轴」上的特例,变换对它们只是纯缩放;特征值就是各主轴上的缩放倍率。

k=(1,2)k=(1,2)^\top 的例子算:沿 kk 方向,(kk)k=k(kk)=k2k=5k(kk^\top)k = k(k^\top k) = \|k\|^2 k = 5k,故 λ1=k2=5\lambda_1 = \|k\|^2 = 5;与 kk 垂直的 x=(2,1)x = (2,-1)^\top 满足 kx=0k^\top x = 0,于是 (kk)x=k0=0=0x(kk^\top)x = k\cdot 0 = 0 = 0\cdot x,故 λ2=0\lambda_2 = 02×22\times2 恰好两个特征值 5 与 0,正对应「秩 1 把正交方向压没、只在 kk 方向放大 k2\|k\|^2 倍」。

放回 HtH_t。结构一目了然:II 让所有方向原样保留,减去 βtktkt\beta_t k_tk_t^\top 只在 ktk_t 这一个方向上动刀,其余 d1d-1 个方向碰都不碰。特征值分两种:

Htkt=ktβtkt(ktkt)=(1βtkt2)kt,Htx=x(xkt)H_t k_t = k_t - \beta_t k_t(k_t^\top k_t) = \big(1 - \beta_t\|k_t\|^2\big)k_t, \qquad H_t x = x \quad (\forall\, x \perp k_t)

配合 L2Norm(kt=1\|k_t\| = 1)就是:沿 ktk_t 缩放 1βt1-\beta_t,正交补方向特征值为 1,即特征值落在 [1βt, 1](0,1][1-\beta_t,\ 1] \subset (0,1]——不放大、不翻转、不发散,这就是数值稳定性的特征值表述。反例也很直白:若不做归一化,取 k=(1,2)k = (1,2)^\topβ=0.6\beta = 0.6,则 1βk2=13=21-\beta\|k\|^2 = 1-3 = -2,特征值变号且模长大于 1,反复作用必然发散。

HtH_t 还有两个后文要用的性质:对称Ht=HtH_t^\top = H_t,这是前面两套约定能靠转置互换的原因);以及这类「II 减秩 1」结构与**豪斯霍尔德变换(Householder transformation)**同型,QR 分解用它做反射消元。区别在于 Householder 取 β=2/k2\beta = 2/\|k\|^2k=1\|k\|=1 时是精确反射(保长);DeltaNet 的 βt(0,1)\beta_t \in (0,1) 是「部分反射」,软化成可学习的写入强度。

秩 1 也解释了为何删除项开销低:擦除单一方向不需要满秩运算,一次外积即可。更进一步,SHtS+βtktvtS \leftarrow H_tS + \beta_t k_tv_t^\top 每步只给状态加一个秩 1 矩阵(秩 1 更新),chunk 内 CC 步就是 CC 个秩 1 更新的连乘叠加——而 WY / UT 变换正是数值线性代数里专门处理「一串秩 1 更新如何打包」的经典工具,后文 chunkwise 一节的源头就在这里。

为什么特征值这个概念到处出现。因为它回答了迭代系统最关心的问题:一个变换反复作用很多次之后会怎样MM 作用 nn 次,在特征向量方向上就是 λn\lambda^nλ>1|\lambda| > 1 的方向爆炸,λ<1|\lambda| < 1 的方向衰减消失,λ=1\lambda = 1 的方向保持不变。本文这条线上的约束几乎都在围着它转:

出现位置 特征值/缩放倍率的约束 目的
SSM/Mamba 的离散化 Aˉ\bar A 特征值(或对角衰减因子)模长 1\le 1 状态不发散
GDN / KDA 的 α(0,1)\alpha \in (0,1) 直接强制衰减算子逐方向缩放 <1< 1 可控遗忘
Ht=IβtktktH_t = I-\beta_tk_tk_t^\top 特征值 [1βt,1]\in [1-\beta_t, 1] 擦除是收缩的
1/Γ1/\Gamma 溢出(K3 加下界的原因) Γ=α\Gamma = \prod\alpha 连乘趋 0,倒数爆炸 与「$

一句话:特征向量是矩阵的「自然方向」,特征值是每个自然方向上的缩放倍率;看懂这两个数,就看懂了矩阵反复作用后的长期行为。

把递归展开,消掉时间依赖。记 Bt=βtktvtB_t = \beta_t k_t v_t^\top,递推 St=HtSt1+BtS_t = H_t S_{t-1} + B_t 逐层代入(主约定下 HH 在左侧,越晚的时间步越靠外):

S1=H1S0+B1S_1 = H_1 S_0 + B_1

S2=H2H1S0+H2B1+B2S_2 = H_2H_1 S_0 + H_2B_1 + B_2

S3=H3H2H1S0+H3H2B1+H3B2+B3S_3 = H_3H_2H_1 S_0 + H_3H_2B_1 + H_3B_2 + B_3

S4=H4H3H2H1S0+H4H3H2B1+H4H3B2+H4B3+B4S_4 = H_4H_3H_2H_1 S_0 + H_4H_3H_2B_1 + H_4H_3B_2 + H_4B_3 + B_4

规律如下:

St=(i=t1Hi)S0+it(j=ti+1Hj)BiS_t = \Big(\prod_{i=t}^{1} H_i\Big) S_0 + \sum_{i \le t} \Big(\prod_{j=t}^{i+1} H_j\Big) B_i

其中 i=t1Hi=HtHt1H1\prod_{i=t}^{1}H_i = H_tH_{t-1}\cdots H_1 表示按时间倒序左乘。展开后递归被彻底消掉了:每个 StS_t 都是初始状态、各步写入项与转移矩阵连乘的线性组合,只剩矩阵乘法和求和。而矩阵乘法满足结合律,「从左往右扫」只是众多括号化方案之一——换个括号方式(比如二叉树式两两合并),HH 的连乘与写入项的累积可以在 O(logn)O(\log n) 深度内并行完成。这就是 parallel scan / associative scan(并行扫描/结合扫描) 类方法的核心思想,也是 chunkwise 并行化和 SSM 训练并行(如 Mamba 的 selective scan)共同的理论根基。

到这里两条路线可以拼在一起了: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 那种全通道的指数衰减。

1.2 GDN:把遗忘门与定点改写拼在一起

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

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

序列长度超过状态容量时,记忆碰撞必然发生;GDN 以板擦与铅笔的组合来管理这块固定大小的白板。主约定下写作

St=αt(Iβtktkt)St1+βtktvtS_t = \alpha_t\big(I - \beta_tk_tk_t^\top\big)S_{t-1} + \beta_t k_t v_t^\top

从这里开始切换到转置约定S^=S\hat S = S^\top,读出 ot=S^tqto_t = \hat S_tq_t),以便与 GDN 论文及 chunkwise 推导逐项对齐;§1.2 余下部分与 §1.3、§1.4 的数值验算和 chunkwise 公式全部使用它。同一个式子转置后是

S^t=S^t1(αt(Iβtktkt))+βtvtkt\hat S_t = \hat S_{t-1}\big(\alpha_t(I - \beta_tk_tk_t^\top)\big) + \beta_t v_t k_t^\top

为减少符号负担,转置约定内部仍把状态记作 StS_t(即下文 StS_tS^t\hat S_t,形状 dv×dkd_v \times d_k,读出为 StqtS_tq_t、旧读出为 St1ktS_{t-1}k_t)。

其中 αt(0,1)\alpha_t \in (0,1) 是数据相关的标量门(Mamba-2 的参数化:α=exp(Softplus(Linear(xt)))\alpha = \exp(-\mathrm{Softplus}(\mathrm{Linear}(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 反射,沿 kk 方向压缩状态;标量 α\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 连乘完全同源。

统一视角。至此四个模型均已出现,它们本质上是同一个在线优化问题的闭式解,差别只在目标函数:

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

先说清这个优化问题本身。把每个时间步看成一次在线学习:已有状态 St1S_{t-1},新到一对样本 (kt,vt)(k_t, v_t),需要解出新状态 StS_t。所有四个模型的目标函数都是下面这个形式,自变量是 StS_tSt1S_{t-1}ktk_tvtv_t 均为已知量:

L(St)=StAtF2正则项2Stkt, ut拟合项\mathcal{L}(S_t) = \underbrace{\|S_t - A_t\|_F^2}_{\text{正则项}} - \underbrace{2\langle S_tk_t,\ u_t\rangle}_{\text{拟合项}}

两项各自度量的是:

  • 正则项 StAtF2\|S_t - A_t\|_F^2:新状态与锚点 AtA_t 的 Frobenius 距离,即所有矩阵元素差的平方和。它惩罚状态的改动量,对应记忆保留。锚点取 St1S_{t-1} 表示要求尽量不动,取 αtSt1\alpha_tS_{t-1} 表示允许先按 αt\alpha_t 收缩再比较,即容忍遗忘;
  • 拟合项 2Stkt, ut-2\langle S_tk_t,\ u_t\rangle:用 ktk_t 检索新状态得到 StktS_tk_t,再与写入目标 utu_t 做内积。前面的负号使内积越大、损失越小,即要求检索结果朝 utu_t 的方向对齐,对应关联学习。

该目标对 StS_t 是二次的,StL=2(StAt)2utkt=0\nabla_{S_t}\mathcal{L} = 2(S_t - A_t) - 2u_tk_t^\top = 0,因此闭式解统一为

St=At+utktS_t = A_t + u_tk_t^\top

于是四个模型的差别可以完全归结为两个量的选择:锚点 AtA_t 决定怎么遗忘,写入目标 utu_t 决定怎么写入

模型 锚点 AtA_t 写入目标 utu_t 在线学习目标 状态更新的闭式解
Linear Attn St1S_{t-1} vtv_t 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 αtSt1\alpha_tS_{t-1} vtv_t 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 St1S_{t-1} βt(vtSt1kt)\beta_t(v_t - S_{t-1}k_t) 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 αtSt1\alpha_tS_{t-1} βt(vtαtSt1kt)\beta_t(v_t - \alpha_tS_{t-1}k_t) 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 St1Diag(αt)S_{t-1}\mathrm{Diag}(\bm\alpha_t) 逐通道收缩 βt(vtSt1Diag(αt)kt)\beta_t(v_t - S_{t-1}\mathrm{Diag}(\bm\alpha_t)k_t) GDN 的逐通道化 St=St1Diag(αt)(Iβtktkt)+βtvtktS_t = S_{t-1}\mathrm{Diag}(\bm\alpha_t)(I-\beta_tk_tk_t^\top) + \beta_tv_tk_t^\top

AtA_tutu_t 代入 St=At+utktS_t = A_t + u_tk_t^\top 即可得到最后一列。以 DeltaNet 为例:ut=βt(vtSt1kt)u_t = \beta_t(v_t - S_{t-1}k_t),代入后 St=St1+βt(vtSt1kt)ktS_t = S_{t-1} + \beta_t(v_t - S_{t-1}k_t)k_t^\top,展开即 St1(Iβtktkt)+βtvtktS_{t-1}(I - \beta_tk_tk_t^\top) + \beta_tv_tk_t^\top,与前面用梯度下降推出的结果一致。

这张表读法很简单:损失里只有两样东西——vtv_t 是 ground truth(这个 key 本来该存的值),StktS_tk_t 是当前记忆对它的预测(这个 key 实际读出来的值),拟合项衡量两者是否一致;正则项则约束新状态别离锚点太远,即该保留多少旧记忆。四个模型的差别只在于:拟合项拿什么当 target(整份 vtv_t,还是残差 vtAtktv_t - A_tk_t),正则项拿什么当锚点(固定的 St1S_{t-1},还是可收缩的 αtSt1\alpha_tS_{t-1}。GDN 在两处都取强化版本,KDA 再把锚点的标量收缩换成逐通道的 Diag(αt)\mathrm{Diag}(\bm\alpha_t)

说到底,这就是把「kvk \to v 这条记忆是否还对得上」写成了一个损失函数:对不上就修(拟合项),但别为了修这一条把整张表推翻(正则项)。

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 分析其继承与改动。

1.3 GDN 数值验算:手算一遍

设定 dk=dv=2d_k = d_v = 2C=3C = 3 个 token,S0=0S_0 = 0。数据刻意构造为**k1=k2=e1k_1 = k_2 = e_1,以制造 key 碰撞**;query 取 Q=KQ = K,即每一步都用当前 token 自己的 key 去查(q1=q2=e1q_1 = q_2 = e_1q3=e2q_3 = e_2),这样读出结果能直接对照「这个地址此刻存的是什么」:

Q=K=[101001], V=[102031], α=[0.8,0.5,0.9], β=[1,1,0.6]Q = 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]

本节沿用 §1.2 声明的转置约定SRdv×dkS \in \mathbb{R}^{d_v \times d_k},读出 ot=Stqto_t = S_tq_t,旧读出为 St1ktS_{t-1}k_t

顺序递归,作为对照基准。累积衰减 γj=ijαi=[0.8,0.4,0.36]\gamma_j = \prod_{i\le j}\alpha_i = [0.8, 0.4, 0.36]

t=1k=e1k=e_1v=[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=S1e1=[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 = S_1e_1 = [1,0]

t=2k=e1k=e_1v=[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}

此时读出 o2=S2q2=S2e1=[2,0]o_2 = S_2q_2 = S_2e_1 = [2,0],即 v2v_2v1v_1 已被覆写(对照线性注意力:纯加性写入会得到 [3,0] 的叠加结果)。

t=3k=e2k=e_2v=[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=S3e2=[1.8,0.6]S_3 = \begin{bmatrix}1.8&1.8\\0&0.6\end{bmatrix}, \qquad o_3 = S_3q_3 = S_3e_2 = [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)与下界衰减都是围绕这一约束做文章。

1.4 GDN 的 Chunkwise 并行形式

推理用递归(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 连乘压缩成两个小矩阵

先说清为什么非做不可。PrP_rCCdk×dkd_k\times d_k 矩阵逐个相乘,O(Cdk3)O(C\,d_k^3)严格顺序——比原递推还贵,并行化直接失败。输出侧同理,每个 oco_c 都要「前 cc 个擦除矩阵的乘积」。

关键观察:这种乘积永远不膨胀。每个因子都是「单位阵减秩 1」,这类矩阵连乘的结果仍是「单位阵减一个低秩矩阵」,秩不超过因子个数:

Pr:=i=1r(Iβikiki)=IirwikiP_r := \prod_{i=1}^{r}(I - \beta_i k_i k_i^\top) = I - \sum_{i\le r} w_i k_i^\top

证明(对 rr 归纳)r=0r=0P0=IP_0 = I,空和成立。设 Pr1=Ii<rwikiP_{r-1} = I - \sum_{i<r} w_i k_i^\top 已成立,右乘第 rr 个因子:

Pr=Pr1(Iβrkrkr)=(Ii<rwiki)βr(Ii<rwiki)krkr=Ii<rwikiβrkrkr+βri<rwi(kikr)标量kr\begin{aligned} P_r &= P_{r-1}\big(I - \beta_r k_r k_r^\top\big) \\[2pt] &= \Big(I - \sum_{i<r} w_i k_i^\top\Big) - \beta_r\Big(I - \sum_{i<r} w_i k_i^\top\Big)k_r k_r^\top \\[2pt] &= I - \sum_{i<r} w_i k_i^\top - \beta_r k_r k_r^\top + \beta_r \sum_{i<r} w_i \underbrace{(k_i^\top k_r)}_{\text{标量}} k_r^\top \end{aligned}

第三、四项都以 krk_r^\top 结尾,合并同类项:

Pr=Ii<rwikiβr(kri<rwi(kikr))记作 wrkr=IirwikiP_r = I - \sum_{i<r} w_i k_i^\top - \underbrace{\beta_r\Big(k_r - \sum_{i<r} w_i (k_i^\top k_r)\Big)}_{\text{记作 } w_r} k_r^\top = I - \sum_{i\le r} w_i k_i^\top

归纳完成。注意 wrw_r 不是定义出来的技巧,而是「乘积保持低秩」这个要求逼出来的——要让结果保持 IiwikiI - \sum_i w_ik_i^\top 的形式,括号里那一坨只能是 wrw_r

wr=βr(kri<rwi(kikr))w_r = \beta_r\Big(k_r - \sum_{i<r}w_i\,(k_i^\top k_r)\Big)

value 侧同理。把第 ii 次写入 βiviki\beta_iv_ik_i^\top 穿过它之后所有的擦除矩阵与衰减,追踪一遍即得

u~r=βr(vri<ru~iγrγi(kikr))\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修正后的擦除向量:第 rr 步本想擦除 krk_r 方向,但若 krk_r 与之前的 kik_i 有重叠(kikr0k_i^\top k_r \ne 0),连乘展开时前面的擦除项已经顺带擦过这部分,wrw_r 把已擦的量减掉以避免重复擦除,减法权重恰是重叠度 kikrk_i^\top k_ru~r\tilde u_r 是同一修正的 value 版,唯一差别是那个衰减比 γr/γi\gamma_r/\gamma_i:第 ii 次写入到第 rr 步时已多衰减 s=i+1rαs=γr/γi\prod_{s=i+1}^{r}\alpha_s = \gamma_r/\gamma_i 倍,故其干扰要按此比例打折——越早的写入衰减越多、干扰越小

为什么 γ\gamma 只出现在 u~\tilde u(即 gating 并入的位置):GDN 的衰减是标量,标量与一切矩阵可交换,擦除连乘里的衰减可整体提到外面变成总因子 γr\gamma_r,所以 ww 的递归里看不见 γ\gamma;但 value 侧每次写入的「存活时长」不同,这个相对衰减无法外提,只能以比值留在递归里。对照 KDA:衰减变成向量后与擦除不可交换γ\gamma 再也提不出去,只能渗进内积本身(见 §4.1 的 MciM_{ci})。

第三步:UT 变换——递归变成一次下三角方程求解

两条递归里 wrw_r 只依赖 wi<rw_{i<r},是严格下三角依赖,因此可以整体写成矩阵方程。把 wrw_r 的递归移项:

wr+i<rβr(kikr)wi=βrkrw_r + \sum_{i<r} \beta_r(k_i^\top k_r)\,w_i = \beta_r k_r

WRC×dkW \in \mathbb{R}^{C\times d_k} 的第 rr 行为 wrw_r^\top,并定义严格下三角矩阵

Lri={βr(kikr),i<r0,irL=strictLower(diag(β)KK)L_{ri} = \begin{cases}\beta_r\,(k_i^\top k_r), & i < r\\ 0, & i \ge r\end{cases} \qquad\text{即}\qquad L = \mathrm{strictLower}\big(\mathrm{diag}(\beta)\,KK^\top\big)

CC 个方程可一次写成 (I+L)W=diag(β)K(I + L)\,W = \mathrm{diag}(\beta)\,K,于是

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

求这个逆很便宜,原因是 LL 幂零。严格下三角矩阵满足 LC=0L^C = 0,所以 Neumann 级数有限项精确截断

(I+L)1=IL+L2+(1)C1LC1(I+L)^{-1} = I - L + L^2 - \cdots + (-1)^{C-1}L^{C-1}

既不需要迭代、也不存在收敛性问题,一次前代法(forward substitution)即可,代价 O(C2)O(C^2)——相对于省下的 O(Cdk3)O(C\,d_k^3) 连乘完全可以忽略。这就是 UT 变换(I+L)(I+L)单位下三角(Unit Triangular,对角为 1,因为 wrw_r 完整依赖自己),求解它即「UT」名字的来源。U~\tilde U 侧只需把内积换成带衰减比的版本:

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

所以 UT 不是额外发明的东西,它就是这两条递归的矩阵形态

乘法链 → 加法链:这才是并行的真正来源。整件事的本质是把擦除矩阵的连乘 r(Iβrkrkr)\prod_r(I - \beta_rk_rk_r^\top) 换成了求和 IrwrkrI - \sum_r w_rk_r^\top,代价是求和项不再是原始的 krk_r 而是带修正的 wrw_r,修正系数由那个小三角求解预先算清。准确地说:把大的顺序依赖CCdk×dkd_k\times d_k 矩阵依次相乘)换成了小的顺序依赖C×CC\times C 三角求解)加一堆可任意并行的加法

求和为什么就是胜利:加法可交换、可结合,因而可任意分组——树形归约、分块、稠密 matmul,GPU 的全部并行性都在奖励「求和结构」;而连乘与递推必须一步一步来。

这个手法在本文这条技术路线上出现了至少四次,难度递增但模式相同:

场景 恒等式 把什么变成了求和
log 空间衰减 sαs=exp(sgs)\prod_s \alpha_s = \exp(\sum_s g_s) 累积衰减 → cumsum(最字面的一个)
SSM 的卷积形式 xt=sAˉtsBˉusx_t = \sum_s \bar A^{t-s}\bar B u_s 顺序递推 → 对历史的加权和,可用卷积/FFT
Mamba-2 的 SSD 递推 \equiv 半可分矩阵,块内 (CB)L(CB^\top)\odot L 选择性递推 → 下三角掩码 × 求和
DeltaNet/GDN/KDA 的 WY-UT (Iβkk)=Iwiki\prod(I-\beta kk^\top) = I - \sum w_ik_i^\top 秩 1 连乘 → 外积求和(本节)

一个反向的注脚:状态本身 S=iu~ikiS = \sum_i \tilde u_ik_i^\top 也是求和——推理时每步加一个秩 1,训练时 WY 把顺序过程也变成求和。累积是语义,求和是算法,这条路线的美学是自洽的。

辨析:因果掩码 ≠ 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^\topTT 三角 「WY」「UT」两个名字的学术出处(Schreiber-Van Loan 1989);DeltaNet 把同样的打包思想借到 delta rule
因果卷积的逆(去卷积) 因果线性系统 = 下三角 Toeplitz,其逆也是下三角 Toeplitz 信号处理经典:「用三角逆矩阵解顺序依赖」
Mamba-2/SSD 的半可分矩阵 块内注意力 (CB)L(CB^\top)\odot LLL 下三角衰减 同一枚硬币另一面:顺序 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 变换,分别是这个模式在「写入打包」与「依赖解耦」上的具体化。

Householder 与 WY 的出处(1958 / 1989)

WY 与 UT 都是数值线性代数的经典老物件,被 DeltaNet 系列「考古」出来复用。既然本文反复用到,把家谱交代清楚。

Householder 变换(1958)就是关于一个超平面的镜像反射。给定单位法向量 uu,反射矩阵为 H=I2uuH = I - 2uu^\top。把任意 xx 拆成沿 uu 的分量与平行镜面的分量,反射即把法向分量翻号:

Hx=x2u(ux)=x2(ux)uHx = x - 2u(u^\top x) = x - 2(u^\top x)\,u

三个性质直接从这个结构来:对称H=HH^\top = H)、正交HH=IH^\top H = I,反射保长)、对合H2=IH^2 = I,照两次镜子回到原样)。其特征值是一个 1-1(法向翻转)与 d1d-1+1+1(镜面内不动)。

它与 delta rule 的擦除矩阵是同族对象。对比 I2uuI - 2uu^\topIβtktktI - \beta_tk_tk_t^\top:取 βtkt2=2\beta_t\|k_t\|^2 = 2 时两者完全相同。结合 §1.1 算过的特征值 1βk21 - \beta\|k\|^2

β\beta(取 k=1|k|=1 沿 kk 的特征值 性质
β0\beta \to 0 1\to 1 几乎不动,不擦除
β=1\beta = 1 00 该方向完全清零(投影)
β=2\beta = 2 1-1 正交反射,即 Householder

所以 delta rule 的擦除矩阵可以理解为没照到底的半面镜子β(0,1)\beta \in (0,1) 只做收缩而非翻转,代价是不再正交,好处是「擦除强度」成了可学习的连续量。Householder 的主战场是 QR 分解:对第 1 列选一面镜子把对角线以下全照成 0,再对第 2 列选一面……mm 面镜子依次照完得到上三角 RR,镜子之积即正交阵 QQ

WY 表示(1989,Schreiber & Van Loan)解决的是「一串镜子怎么存」。QR 做完后 Q=H1H2HmQ = H_1H_2\cdots H_mmmn×nn\times n 反射之积,每次用它都重新连乘既贵又顺序。他们的观察是

Q=H1H2Hm=I+WY,W,YRn×mQ = H_1H_2\cdots H_m = I + WY^\top, \qquad W, Y \in \mathbb{R}^{n\times m}

mm 个大方阵之积压缩成两个瘦矩阵mnm \ll n),用的时候两次矩阵乘即可:Qx=x+W(Yx)Qx = x + W(Y^\top x)。推导与上面 wrw_r 的归纳一模一样:每个因子是「II 加秩 1」,连乘时秩只累加不膨胀,归纳地合并出第 jj 个修正向量。进一步可写成 Q=IYTYQ = I - YTY^\top,其中 TTm×mm\times m 上三角矩阵,由一个小递归算出——这就是 UT 变换名字的出处。这套东西在 LAPACK 的 QR 例程(xGEQRF / xORMQR)底层已经跑了几十年。

把家谱摆出来:

数值线性代数(1958 / 1989) DeltaNet / GDN / KDA(2020s)
基本砖块 I2uuI - 2uu^\top(反射) IβkkI - \beta kk^\top(擦除)
要打包的对象 mm 面镜子之积 QQ chunk 内 CC 次擦除之积 PCP_C
紧凑形式 I+WYI + WY^\top(或 IYTYI - YTY^\top IiwikiI - \sum_i w_ik_i^\top
修正向量递归 逐个镜子扣重叠 wr=βr(kri<rwikikr)w_r = \beta_r(k_r - \sum_{i<r}w_i\,k_i^\top k_r)
小三角因子 TT(上三角) (I+L)1(I+L)^{-1}(单位下三角)
目的 QQ 的应用变成 BLAS-3 顺序写入变成稠密 matmul

一句话总括:Householder 变换是一类「II 减秩 1」的反射矩阵,WY 表示是把这类矩阵的连乘压缩成两个瘦矩阵的经典技巧;DeltaNet 的作者发现 delta rule 的擦除矩阵恰是同一族对象,于是把这个三十多年前的打包技巧搬了过来,让 chunk 内的顺序写入得以并行——GDN 与 KDA 一路继承,本文的 UT 变换就是 compact-WY 里那个三角因子的计算。

回到公式。上面 TgatedT_{\text{gated}} 里的衰减感知掩码即 Γij=γi/γj (i>j)\Gamma_{ij} = \gamma_i/\gamma_j\ (i>j)WWU~\tilde U 只差内积处的一个 γ\gamma 比。下三角求解(C×CC\times C,实践中 C=64C = 64)用前代法完成。

第四步: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 1KK=[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 4S0=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 的一般情形同样对拍通过)。

2. GDN 在网络里长什么样

§1 讲的都是状态怎么更新,属于 token mixer 内部的数学。但 q,k,v,α,βq, k, v, \alpha, \beta 这些量本身从哪来、算完的 o~\tilde o 又怎么变成层输出,还没交代。本节补上这一层: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,分别提供局部上下文与归一化后的稳定擦除;α\alphaβ\beta 是标量,只需线性投影,不经过卷积;输出门沿用 Mamba 的 SiLU 门设计,K3 将其升级为满秩。

2.1 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 0t<jt < j 时零填充,看不到未来 token

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

  1. 时序维度滑动,只混合最近 W 个历史 token,开销远小于全局注意力;
  2. depthwisewkw_kxtkx_{t-k} 逐通道相乘,通道间不混——局部上下文的注入不破坏各通道独立的衰减语义(与 Diag(αt)\mathrm{Diag}(\alpha_t) 逐通道门配套);
  3. 因果tj0t-j \ge 0 保证只看历史,t<jt<j 时零填充,LLM 自回归必备。

其作用是补足线性 RNN 缺失的局部建模能力:递归状态只携带压缩后的全局历史,最近若干 token 的精细局部模式(词内字符、短程搭配)由这层卷积负责。

2.2 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}}

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

此处使用它的原因有两点:其一,平滑且非单调(负区间存在小幅下凹,xx \to -\infty 时输出趋于 0,而非 ReLU 的硬截断),梯度处处非零,深层网络训练更稳定;其二,门控形式与整个 block 的「信息通过与抑制」语义一致,Q/K/V 进入递归状态前先经过一次自门控,相当于对局部卷积混合后的特征做一次软筛选,随后由 L2Norm 归一化到单位范数。K3 将输出侧的这道门升级为满秩输入相关门,其来源即此处的 SiLU。

2.3 L2Norm

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 不使用?回到 §1.1 的结论:归一化后 kt2=1\|k_t\|_2 = 1,于是 delta rule 的擦除矩阵 IβkkI - \beta k k^\top 的特征值落在 [1β, 1][1-\beta,\ 1] 区间内,写入步长因此稳定,β=1\beta = 1 时为精确的保长反射。Q 也做归一化,目的是使读出 o=Sqo = S^\top q 的尺度可控:query 与 key 的范数均为 1,内积才具备可比性,chunkwise 公式中 QKQK^\top 的 score 也被限制在 [1,1][-1, 1] 内。V 不做归一化,因为写入内容的幅值本身携带信息,且 βtktvt\beta_t k_t v_t^\top 的稳定性由 k 侧保证,与 v 的尺度无关。这与 softmax 注意力中 QK-norm 防止 logit 爆炸的动机同源,但在线性 RNN 中它还额外承担擦除操作的数值稳定性,因此归一化并非可选优化,而是 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 论文可视为这一混合范式的早期工作。

3. 从 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 这一步引入的新问题的修复。

3.1 SSM 路线的结论

SSM 从控制论的状态方程出发,经 ZOH 离散化、S4、Mamba 到 Mamba-2,推导过程见前置篇第 2 节。本文需要的只是其结论:乘在旧状态上的标量衰减系数 αt\alpha_t,即线性注意力一族对「状态如何遗忘」的回答,但它是全局衰减、不区分通道。它与 DeltaNet 的定点覆写在 GDN 处结合,GDN 又被 KDA 逐通道化。整体骨架始终不变:

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

3.2 KDA 的递推公式

本节回到主约定StRdk×dv\mathbf{S}_t \in \mathbb{R}^{d_k\times d_v},读出 Stqt\mathbf{S}_t^\top\bm q_t),与 KDA 论文写法一致:

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 型,详见 §1.1);
  3. 新信息写入:叠加 βtktvt\beta_t \bm{k}_t \bm{v}_t^\top

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

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

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

标量门的表达力瓶颈。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。与 §1.2 的在线学习表同构,只需把正则项换成逐通道版。每步给定新样本 (kt,vt)(k_t, v_t),希望新状态 S 满足两点:其一,拟合新样本 SktvtSk_t \approx v_t;其二,不过度偏离「衰减后的历史记忆」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) 不可交换,顺序错了结果就错(与 §1.2「α 乘整个括号」的警告同源,但逐通道化之后错误更隐蔽)。FLA 参考实现里对应 S = S * g.exp() 之后立刻用衰减后的 S 做检索 v - k^T S

3.4 下界衰减(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 1iji \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.00670.0067),数十步后衰减量已足够小,不需要单步清零的能力。以一个衰减下界换取全链路 Tensor Core 化,是有利的取舍。

3.5 满秩门控(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 独立调制从循环状态读取的通道。

4. KDA 的 Chunkwise 并行形式

递推形式在推理阶段是优势——O(1)O(1) 状态更新;但在训练与 prefill 阶段成为瓶颈:每个 token 的状态依赖前一个 token,顺序循环使 GPU 的数千个核心无法并行工作。Chunkwise 并行化把序列切成长度 C 的块:块内矩阵运算并行,块间只传状态。先看通用的复杂度框架,再看 KDA 逐通道门带来的推导细节。

通用框架Chunk 内:token 交互用带衰减的因果注意力直接算,O(C2d)O(C^2 d)——C 是常数(64/128),对序列长度 N 线性。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 就是卷积的继任者。

4.1 逐通道衰减下的 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 变换。构造严格下三角 LLLci=β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 的因子——刻意保持,原因见下界衰减一节。

4.2 数值验算(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.3 与论文公式 (4) 的对照

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 实现了更强大的信息写入和遗忘控制。

4.4 速查表(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 ⊙ Γ_causal)(Ũ - 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 深度分析,对比两种实现的差异。