ggaaooppeenngg

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

TileLang 实战:KDA 从零到一–Kimi Delta Attention

前三篇分别实现了无遗忘的 chunked 线性注意力、标量衰减、以及带删除项的 Gated DeltaNet。本篇是系列收尾:把 GDN 的标量αt\alpha_t 换成逐通道向量atRdk\bm{a}_t \in \mathbb{R}^{d_k}

St=St1Diag(at)(Iβtktkt)+βtvtkt\mathbf{S}_t = \mathbf{S}_{t-1}\operatorname{Diag}(\bm{a}_t)\big(\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal}\big) + \beta_t\bm{v}_t\bm{k}_t^{\intercal}

改动只有一处:αt\alpha_t 变成了 Diag(at)\operatorname{Diag}(\bm{a}_t)。但这一处把前三篇积累的所有便利拆掉了–衰减不再能从矩阵里外提,Γ\GammaC×CC\times C 变成 C×C×dkC \times C \times d_k,而累积衰减积的通道离散度会直接把 fp16 打穿。

本篇聚焦实现。KDA 的数学推导(递推式、逐通道 WY 表示、UT 变换、下界衰减与满秩门控的动机)见《KDA 的来龙去脉》§3–§4,这里不重复;本文只做一件事:把那些公式落成能跑的 kernel,并量化每一处数值边界

前三篇见 Chunked 线性注意力标量衰减Gated DeltaNet


0. 符号与三处变化

沿用前三篇:βt\beta_t 写入强度、CC 块长、[t][t] 块序号、r[1,C]r \in [1,C] 块内位置、SRdk×dv\mathbf{S}^\intercal \in \mathbb{R}^{d_k \times d_v} 为 kernel 存储布局。

本篇的核心量改为累积 log 衰减(沿用《来龙去脉》的记号):

γi=s=1igsRdk,gs=logas<0 (逐分量)\bm{\gamma}_i = \sum_{s=1}^{i}\bm{g}_s \in \mathbb{R}^{d_k}, \qquad \bm{g}_s = \log\bm{a}_s < \bm{0}\ \text{(逐分量)}

注意 γi\bm{\gamma}_i向量–这是与前两篇最本质的区别。前三篇的 γr\gamma^r 是标量,Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 是一张 C×CC \times C 的表;本篇 γiγj\bm{\gamma}_i - \bm{\gamma}_jdkd_k 维向量,"衰减矩阵"概念上是 C×C×dkC \times C \times d_k不可能物化

三处结构性变化:

GDN(标量门) KDA(逐通道门)
累积积 γr\gamma^r 标量,cumsum 后 CC 个数 γrRdk\bm{\gamma}_r \in \mathbb{R}^{d_k},cumsum 后 C×dkC \times d_k
衰减掩码 ΓRC×C\Gamma \in \mathbb{R}^{C\times C},可物化 概念上 C×C×dkC\times C\times d_k必须融进 GEMM
KK 矩阵 KKΓ\mathbf{K}\mathbf{K}^\intercal \odot \Gamma,衰减可外提 Mci=(kceγcγi)kiM_{ci} = (\bm{k}_c \odot e^{\bm{\gamma}_c - \bm{\gamma}_i})^\intercal\bm{k}_i,衰减长在内部

第三行是全篇的技术核心。


1. 衰减长在内部:问题与出路

1.1 朴素做法的代价

Mci=(kceγcγi)kiM_{ci} = (\bm{k}_c \odot e^{\bm{\gamma}_c - \bm{\gamma}_i})^\intercal\bm{k}_i 里的指数依赖 (c,i,d)(c, i, d) 三个下标。直接算就是 C2/2C^2/2 次长度 dkd_k 的加权内积,每次都要现算 dkd_kexp

这不是常数因子问题–它把一次 GEMM 变成了 C2/2C^2/2 个独立的向量运算,完全用不上 Tensor Core。C=64C = 64dk=128d_k = 128 时是 2016 次加权内积、约 26 万次 exp2

1.2 出路:指数可分离

关键观察是指数可以按通道拆开

eγcγi=eγceγie^{\bm{\gamma}_c - \bm{\gamma}_i} = e^{\bm{\gamma}_c} \odot e^{-\bm{\gamma}_i}

于是:

Mci=d(kc[d]eγc[d])K~+[c,d](ki[d]eγi[d])K~[i,d]=(K~+K~)ciM_{ci} = \sum_{d}\underbrace{\big(k_c[d]\,e^{\gamma_c[d]}\big)}_{\widetilde{K}^{+}[c,d]}\underbrace{\big(k_i[d]\,e^{-\gamma_i[d]}\big)}_{\widetilde{K}^{-}[i,d]} = \big(\widetilde{\mathbf{K}}^{+}\widetilde{\mathbf{K}}^{-\intercal}\big)_{ci}

一次 GEMM 解决,前置两个 C×dkC \times d_k 的逐元素加权。实测拆分与朴素计算的差异 1.94×10161.94 \times 10^{-16},代数上完全等价。

1.3 代价:eγie^{-\bm{\gamma}_i} 是大于 1 的量

第二篇讲过一个教训:把比值 γi/γj\gamma_i/\gamma_j 因式分解成 γi(1/γj)\gamma_i \cdot (1/\gamma_j) 会物化一个指数增长的量,fp16 下 α0.8\alpha \le 0.8 即 NaN。这里是同一个陷阱的逐通道版本eγi[d]1e^{-\gamma_i[d]} \ge 1,且随 ii 与通道衰减强度指数增长。

但本篇的处境和第二篇不同:那里有替代方案(保留比值形式),这里没有。逐通道衰减无法外提,不拆就用不上 Tensor Core。所以问题从"要不要拆"变成了"怎样让拆分在数值上安全"。

实测 eγe^{-\bm{\gamma}} 的最大值(C=64C = 64dk=128d_k = 128):

门控下界 maxeγ\max e^{-\bm{\gamma}} fp16 fp32
0.999 1.041.04 OK OK
0.99 1.481.48 OK OK
0.95 7.147.14 OK OK
0.9 54.9554.95 OK OK
0.8 4.20×1034.20 \times 10^{3} OK OK
0.5 3.04×10103.04 \times 10^{10} 溢出 OK

门控下界在 0.9 以上时,eγe^{-\bm{\gamma}} 连 fp16 都装得下。 这就把 KDA 论文里"下界衰减(lower-bounded decay)"这个设计从模型层面的技巧,变成了 kernel 能否用 Tensor Core 的前提条件。


2. 下界衰减:不是精度调优,是可行性前提

2.1 无下界时 fp16 全线归零

eγe^{-\bm{\gamma}} 会溢出,另一头 eγe^{\bm{\gamma}} 会下溢。而逐通道门让后者严重得多–总有一些通道学到很小的 aa,它们的 γ\gamma 累积得最快。

实测 mineγC\min e^{\bm{\gamma}_C} 与 fp16 下归零的通道数(dk=128d_k = 128):

a\bm{a} 采样区间 CC mineγC\min e^{\bm{\gamma}_C} fp16 归零通道
[0.9, 0.999][0.9,\ 0.999] 64 1.70×1021.70\times10^{-2} 0 / 128
[0.9, 0.999][0.9,\ 0.999] 128 4.84×1044.84\times10^{-4} 0 / 128
[0.5, 0.999][0.5,\ 0.999] 64 2.20×10112.20\times10^{-11} 118 / 128
[0.5, 0.999][0.5,\ 0.999] 128 1.59×10201.59\times10^{-20} 128 / 128
[0.1, 0.999][0.1,\ 0.999] 64 7.39×10287.39\times10^{-28} 128 / 128
[0.01, 0.999][0.01,\ 0.999] 128 6.44×10656.44\times10^{-65} 128 / 128

C=128C = 128、下界 0.5 时全部 128 个通道在 fp16 下归零–状态被彻底清空,kernel 输出恒为块内项,跨块信息完全丢失。

加下界后:

下界 mineγC\min e^{\bm{\gamma}_C} fp16 归零
1.84×10661.84\times10^{-66} 128 / 128
0.5 9.07×10329.07\times10^{-32} 128 / 128
0.9 1.72×1061.72\times10^{-6} 0 / 128
0.95 1.43×1031.43\times10^{-3} 0 / 128

下界 0.9 是分界线。 这一个数字同时解决了两头:eγe^{\bm{\gamma}} 不下溢、eγe^{-\bm{\gamma}} 不溢出。

2.2 通道离散度:逐通道门特有的问题

标量门下 γC\gamma^C 是一个数;逐通道门下它是 dkd_k 个数,而它们的极差决定了同一个 chunk 内不同通道的数值尺度差异:

a\bm{a} 区间 CC γC\bm{\gamma}_C 极差(log 域) e极差e^{\text{极差}}
[0.9, 0.999][0.9,\ 0.999] 64 1.28 3.583.58
[0.9, 0.999][0.9,\ 0.999] 128 1.74 5.695.69
[0.5, 0.999][0.5,\ 0.999] 64 8.17 3.53×1033.53\times10^{3}
[0.5, 0.999][0.5,\ 0.999] 128 11.25 7.69×1047.69\times10^{4}
[0.01, 0.999][0.01,\ 0.999] 128 47.46 4.10×10204.10\times10^{20}

极差 4747 意味着同一个 fragment 里最强和最弱通道的数值相差 20 个数量级–任何浮点格式都无法同时表示。这是逐通道门独有的病:标量门只需要担心整体的尺度漂移,逐通道门还要担心通道之间的尺度撕裂。

下界 0.9 把极差压到 1.74(e1.74=5.7e^{1.74} = 5.7),完全可控。

2.3 sub-chunk:第二道保险

即使有下界,C=128C = 128maxeγ=1.79×103\max e^{-\bm{\gamma}} = 1.79\times10^3(下界 0.9)–fp16 装得下但余量不多。把块内再切成 sub-chunk,每个 sub-chunk 内部重新起算 γ\bm{\gamma}

划分 maxeγ\max e^{-\bm{\gamma}}
C=128C = 128 整块 1.79×1031.79\times10^{3}
sub-chunk =64= 64 5.34×1015.34\times10^{1}
sub-chunk =32= 32 9.029.02
sub-chunk =16= 16 3.413.41

指数跨度只取决于 sub-chunk 长度,与总块长无关。这与第二篇"累积积按 chunk 重置"是同一个道理,只是又下降了一层:chunk 重置控制 γ\bm{\gamma} 本身,sub-chunk 重置控制 eγe^{-\bm{\gamma}} 的动态范围。

代价是 sub-chunk 之间需要额外的状态传递,块内变成两层循环。C=64C = 64 + 下界 0.9 时 maxeγ=55\max e^{-\bm{\gamma}} = 55不需要 sub-chunkC=128C = 128 时建议切 32 或 64。


3. 四层参考实现

沿用前三篇框架。

3.1 参考 A:逐 token 递归

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
def ref_A(Q, K, V, g, beta):
"""g: (B,N,H,DK) 逐通道 log 门控,g < 0"""
B, N, H, DK = Q.shape
DV = V.shape[-1]
O = np.zeros((B, N, H, DV))
for b in range(B):
for h in range(H):
S = np.zeros((DV, DK)) # 论文方向
for t in range(N):
k, v, q = K[b, t, h], V[b, t, h], Q[b, t, h]
a, bt = np.exp(g[b, t, h]), beta[b, t, h]
# Diag(a) 在 S 与 Householder 之间
S = S @ np.diag(a) @ (np.eye(DK) - bt * np.outer(k, k)) \
+ bt * np.outer(v, k)
O[b, t, h] = S @ q
return O

Diag(at)\operatorname{Diag}(\bm{a}_t) 的位置很关键:它夹在 St1\mathbf{S}_{t-1} 与 Householder 之间。写成 Diag(a)S\operatorname{Diag}(\bm{a})\mathbf{S} 或挪到括号外都会算错。

3.2 参考 B:chunkwise + 逐通道 UT 变换

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
32
33
34
35
36
37
38
def ref_B(Q, K, V, g, beta, C):
B, N, H, DK = Q.shape
DV = V.shape[-1]
O = np.zeros((B, N, H, DV))
for b in range(B):
for h in range(H):
S = np.zeros((DK, DV)) # 转置布局
for c in range(N // C):
sl = slice(c * C, (c + 1) * C)
Qc, Kc, Vc = Q[b, sl, h], K[b, sl, h], V[b, sl, h]
gc, bt = g[b, sl, h], beta[b, sl, h]

gam = np.cumsum(gc, axis=0) # γ_i ∈ R^{DK}
gC = gam[-1]

# M[c,i] = (k_c ⊙ e^{γ_c-γ_i})·k_i, i < c
M = np.zeros((C, C))
for cc in range(C):
for i in range(cc):
M[cc, i] = np.dot(Kc[cc] * np.exp(gam[cc] - gam[i]), Kc[i])

L = np.tril(bt[None, :] * M, -1) # L[c,i] = β_i M[c,i]
Tm = np.linalg.inv(np.eye(C) + L)
Ah = bt[:, None] * Tm # diag(β) T

W = Ah @ (Kc * np.exp(gam)) # C × DK
U = Ah @ Vc # C × DV
Vt = U - W @ S # 伪值 Ṽ

# A^qk[c,j] = (q_c ⊙ e^{γ_c-γ_j})·k_j, j ≤ c
Aqk = np.zeros((C, C))
for cc in range(C):
for j in range(cc + 1):
Aqk[cc, j] = np.dot(Qc[cc] * np.exp(gam[cc] - gam[j]), Kc[j])

O[b, sl, h] = (Qc * np.exp(gam)) @ S + Aqk @ Vt
S = np.diag(np.exp(gC)) @ S + (Kc * np.exp(gC - gam)).T @ Vt
return O

这份参考刻意用朴素双循环M\mathbf{M}Aqk\mathbf{A}^{qk}–慢,但与公式逐字对应,作为基准可信。§3.3 的 kernel 镜像才用 GEMM 化写法。

Lci=βiMciL_{ci} = \beta_i M_{ci}A^=diag(β)T\hat{\mathbf{A}} = \operatorname{diag}(\beta)\mathbf{T} 这一对下标必须配套。实测另一种等价写法是 Lci=βcMciL_{ci} = \beta_c M_{ci}A^=Tdiag(β)\hat{\mathbf{A}} = \mathbf{T}\operatorname{diag}(\beta),两者都对,混用则错。

3.3 参考 C:kernel 控制流镜像(GEMM 化)

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
32
33
34
35
36
37
38
39
def ref_C(Q, K, V, g, beta, C, block_DV):
B, N, H, DK = Q.shape
DV = V.shape[-1]
O = np.zeros((B, N, H, DV))
for bv in range(DV // block_DV): # ← grid.x
dv = slice(bv * block_DV, (bv + 1) * block_DV)
for bbh in range(B * H): # ← grid.y
b, h = bbh // H, bbh % H
S = np.zeros((DK, block_DV))
for c in range(N // C): # ← T.Pipelined
sl = slice(c * C, (c + 1) * C)
Qc, Kc = Q[b, sl, h], K[b, sl, h]
Vc = V[b, sl, h, dv]
gc, bt = g[b, sl, h], beta[b, sl, h]

gam = np.cumsum(gc, axis=0)
gC = gam[-1]
Ep, Em = np.exp(gam), np.exp(-gam) # e^{γ}, e^{-γ}

# ── 指数可分离:两次 GEMM 代替 C²/2 次加权内积 ──
Kp, Km = Kc * Ep, Kc * Em
M = np.tril(Kp @ Km.T, -1) # 严格下三角
Aqk = np.tril((Qc * Ep) @ Km.T) # 含对角

L = np.tril(bt[None, :] * M, -1)
# 前向替换解 T @ [Kp | Vc],再逐行乘 β(因为 Â = diag(β)T)
RHS = np.concatenate([Kc * Ep, Vc], axis=1)
X = np.zeros_like(RHS)
for r in range(C):
X[r] = RHS[r]
for j in range(r):
X[r] -= L[r, j] * X[j]
X = bt[:, None] * X # ← β 在解完之后乘
W, U = X[:, :DK], X[:, DK:]

Vt = U - W @ S
O[b, sl, h, dv] = (Qc * Ep) @ S + Aqk @ Vt
S = Ep[-1][:, None] * S + (Kc * np.exp(gC - gam)).T @ Vt
return O

三处与参考 B 不同的实现选择:

  1. M\mathbf{M}Aqk\mathbf{A}^{qk} 走 GEMMK~+K~\widetilde{\mathbf{K}}^{+}\widetilde{\mathbf{K}}^{-\intercal}Aqk\mathbf{A}^{qk} 含对角(jcj \le c),M\mathbf{M} 不含(i<ci < c)。
  2. W\mathbf{W}U\mathbf{U} 拼成一次前向替换:两者共用同一个 (I+L)(\mathbf{I}+\mathbf{L}),拼接后只解一遍,省一半串行开销。
  3. β\beta 必须在解完三角系统之后乘A^=diag(β)T\hat{\mathbf{A}} = \operatorname{diag}(\beta)\mathbf{T} 展开是「先 T\mathbf{T} 作用、再逐行乘 β\beta」;若把 β\beta 提前乘进右端项,算的就是 Tdiag(β)\mathbf{T}\operatorname{diag}(\beta)–那是另一种配对(需搭配 Lci=βcMciL_{ci} = \beta_c M_{ci}),混用则错。这个错误在 β\beta 全部相等时完全看不出来,我第一次写就踩了:β\beta 随机时相对 L2 达 1.3×1011.3\times10^{-1}
  4. 状态更新用逐行乘代替 Diag\operatorname{Diag}Ep[-1][:, None] * S 就是 Diag(eγC)S\operatorname{Diag}(e^{\bm{\gamma}_C})\mathbf{S},不物化对角矩阵。

3.4 一致性验证

B=2,H=2,N=12,dk=dv=4,C=4B=2, H=2, N=12, d_k=d_v=4, C=4aU(0.90,0.999)dk\bm{a} \sim \mathcal{U}(0.90, 0.999)^{d_k}βU(0.1,0.9)\beta \sim \mathcal{U}(0.1, 0.9)q,k\bm{q},\bm{k} 已 L2 归一化,fp64:

比较 max abs 误差 相对 L2
B chunkwise vs A 逐 token 3.89×10163.89 \times 10^{-16} 2.46×10162.46 \times 10^{-16}
C kernel 镜像(GEMM 化)vs A 4.44×10164.44 \times 10^{-16} 2.57×10162.57 \times 10^{-16}
指数分离 GEMM vs 朴素加权内积 1.11×10161.11 \times 10^{-16}
L=βiML = \beta_i M + diag(β)T\operatorname{diag}(\beta)\mathbf{T} 1.11×10161.11 \times 10^{-16}
L=βcML = \beta_c M + Tdiag(β)\mathbf{T}\operatorname{diag}(\beta) 1.11×10161.11 \times 10^{-16}

三个退化检验(本篇比 GDN 多一个):

退化 应回到 检验的部分 实测
glogα1\bm{g} \equiv \log\alpha \cdot \bm{1}(所有通道同值) GDN 逐通道的指数分离 4.44×10164.44\times10^{-16}
β0\beta \to 0 逐通道纯衰减 整个 UT 与三角系统 1.21×10271.21\times10^{-27}
g0\bm{g} \to \bm{0}β\beta 保留 纯 DeltaNet 所有衰减权重 4.44×10164.44\times10^{-16}

三个 blockDV\text{block}_{DV}(1 / 2 / 4)的相对 L2 分别为 2.482.48 / 2.572.57 / 2.57×10162.57 \times 10^{-16}–切 DV 零依赖,与第一篇的结论一致。

第一个是本篇独有且最重要的:把逐通道门退化成标量门,必须精确回到第三篇的结果。这一步能抓住"指数分离时把通道维和位置维搞混"这类错误–而那类错误在通道值本来就相同时会隐身。


4. TileLang kernel

4.0 寄存器账

dk=dv=128d_k = d_v = 128C=64C = 64blockDV=32\text{block}_{DV} = 32、128 线程:

fragment GDN KDA 说明
S_f [dk,bDV][d_k,\text{bDV}] 32 32 不变
acc_o [C,bDV][C,\text{bDV}] 16 16 不变
Gam [C,C][C,C] 32 0 逐通道无法物化,取消
A [C,C][C,C] 32 32 Aqk\mathbf{A}^{qk}
Amat [C,C][C,C] 32 32 三角系统系数 L\mathbf{L}
Ep/Em [C,dk][C,d_k] 2×64 = 128 新增:e±γe^{\pm\bm{\gamma}}
W [C,dk][C,d_k] 64 新增
Delta/Vt [C,bDV][C,\text{bDV}] 16 16
合计 ≈ 160 ≈ 320 上限 255

320 超限了。 逐通道门带来的 [C,dk][C, d_k] 量级 fragment(e±γe^{\pm\bm{\gamma}}W\mathbf{W})比 C2C^2 表更吃寄存器–C×dk=64×128C \times d_k = 64\times128C2=642C^2 = 64^2 的两倍。

三条出路:

  1. eγe^{-\bm{\gamma}} 不常驻:它只在构造 M\mathbf{M}Aqk\mathbf{A}^{qk} 时用,用完即弃,可以放 shared memory。省 64 reg。
  2. W\mathbf{W} 走 shared:它是 UWS\mathbf{U} - \mathbf{W}\mathbf{S} 的中间量,不参与后续 GEMM 的累加器。省 64 reg。
  3. CC 降到 32:所有 CC 相关量减半。

组合 1+2 后约 192 reg,可行。下面的 kernel 采用这个方案。

4.1 声明与累积衰减

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
32
33
34
35
36
37
38
39
40
41
42
@tilelang.jit(out_idx=[5])
def kda_chunk(B, H, S, DK, DV, blk=64, block_DV=32,
num_stages=2, threads=128,
dtype=T.float16, accum_dtype=T.float32):
C = blk
NS = T.ceildiv(S, C)

@T.prim_func
def main(
Q: T.Tensor([B, S, H, DK], dtype),
K: T.Tensor([B, S, H, DK], dtype),
V: T.Tensor([B, S, H, DV], dtype),
G: T.Tensor([B, S, H, DK], accum_dtype), # 逐通道 log 门控,< 0
Beta: T.Tensor([B, S, H], accum_dtype),
O: T.Tensor([B, S, H, DV], dtype),
):
with T.Kernel(T.ceildiv(DV, block_DV), B * H, threads=threads) as (bv, bbh):
bb, bh = bbh // H, bbh % H
dv0 = bv * block_DV

Q_s = T.alloc_shared([C, DK], dtype)
K_s = T.alloc_shared([C, DK], dtype)
V_s = T.alloc_shared([C, block_DV], dtype)
S_s = T.alloc_shared([DK, block_DV], dtype)
O_s = T.alloc_shared([C, block_DV], dtype)
# 逐通道特有:e^{±γ} 与 W 放 shared,避免寄存器超限
Ep_s = T.alloc_shared([C, DK], accum_dtype) # e^{γ_r}
Em_s = T.alloc_shared([C, DK], accum_dtype) # e^{-γ_r}
Kp_s = T.alloc_shared([C, DK], dtype) # k ⊙ e^{γ}
Km_s = T.alloc_shared([C, DK], dtype) # k ⊙ e^{-γ}
W_s = T.alloc_shared([C, DK], accum_dtype)
Vt_s = T.alloc_shared([C, block_DV], dtype)

S_f = T.alloc_fragment([DK, block_DV], accum_dtype) # loop-carried
acc_o = T.alloc_fragment([C, block_DV], accum_dtype)
Aqk = T.alloc_fragment([C, C], accum_dtype)
Aq_c = T.alloc_fragment([C, C], dtype)
Lmat = T.alloc_fragment([C, C], accum_dtype) # 三角系数
Vt = T.alloc_fragment([C, block_DV], accum_dtype)
RHS_v = T.alloc_fragment([C, block_DV], accum_dtype)
gam = T.alloc_fragment([C, DK], accum_dtype) # 累积 log
be_f = T.alloc_fragment([C], accum_dtype)

累积衰减是逐通道的前缀和dkd_k 个通道各自独立 cumsum,通道之间完全并行:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
T.clear(S_f)

for i_s in T.Pipelined(NS, num_stages=num_stages):
s0 = i_s * C
T.copy(Q[bb, s0:s0+C, bh, :], Q_s)
T.copy(K[bb, s0:s0+C, bh, :], K_s)
T.copy(V[bb, s0:s0+C, bh, dv0:dv0+block_DV], V_s)

# ── 逐通道累积 log 衰减:通道维并行,位置维串行 ──
for d in T.Parallel(DK):
gam[0, d] = G[bb, s0, bh, d]
for r in T.serial(1, C):
for d in T.Parallel(DK):
gam[r, d] = gam[r - 1, d] + G[bb, s0 + r, bh, d]

for r, d in T.Parallel(C, DK):
Ep_s[r, d] = T.exp2(gam[r, d] * 1.4426950408889634) # e^{γ}
Em_s[r, d] = T.exp2(-gam[r, d] * 1.4426950408889634) # e^{-γ}
Kp_s[r, d] = T.Cast(dtype, K_s[r, d] * Ep_s[r, d])
Km_s[r, d] = T.Cast(dtype, K_s[r, d] * Em_s[r, d])
for r in T.Parallel(C):
be_f[r] = Beta[bb, s0 + r, bh]

对比 GDN:那里前缀和是 CC 步串行、每步 1 个数;这里是 CC 步串行、每步 dkd_k 个通道并行。串行长度不变,但并行度从 1 涨到 128–逐通道门在这一处反而更适合 GPU。

1.4426950408889634log2e\log_2 e,把 exe^x 转成硬件 exp2(与第三篇一致)。

4.2 指数分离:两次 GEMM 取代双重循环

1
2
3
4
5
6
7
8
9
10
11
12
# ── M = strictLower(Kp @ Km^T):衰减 KK 矩阵 ──
T.gemm(Kp_s, Km_s, Lmat, transpose_B=True, clear_accum=True)
for i, j in T.Parallel(C, C):
Lmat[i, j] = T.if_then_else(
j < i, be_f[j] * Lmat[i, j], 0.0) # L[c,i] = β_i M[c,i]

# ── A^qk = tril(Qp @ Km^T),含对角 ──
for r, d in T.Parallel(C, DK):
Q_s[r, d] = T.Cast(dtype, Q_s[r, d] * Ep_s[r, d]) # q ⊙ e^{γ}
T.gemm(Q_s, Km_s, Aqk, transpose_B=True, clear_accum=True)
for i, j in T.Parallel(C, C):
Aqk[i, j] = T.if_then_else(j <= i, Aqk[i, j], 0.0)

这是全篇最关键的两行 GEMM。 朴素写法需要 C2/2C^2/2 次长度 dkd_k 的加权内积(C=64,dk=128C=64, d_k=128 时约 26 万次 exp);分离后 e±γe^{\pm\bm{\gamma}} 各算一次(C×dk=8192C \times d_k = 8192exp2),剩下交给 Tensor Core。

注意两个掩码的边界不同:Lmatj<ij < i(严格下三角,对角会重复计入 βr\beta_r),Aqkjij \le i(含对角,自己看自己零衰减)。

4.3 前向替换:W 与 U 合并求解

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
# ── 右端项:[Kp | V],注意 β 不在这里乘 ──
for r, d in T.Parallel(C, DK):
W_s[r, d] = K_s[r, d] * Ep_s[r, d]
for r, d in T.Parallel(C, block_DV):
RHS_v[r, d] = V_s[r, d]

# ── 前向替换,一次解出 T@[Kp | V] ──
for r in T.serial(C):
for j in T.serial(r):
for d in T.Parallel(DK):
W_s[r, d] -= Lmat[r, j] * W_s[j, d]
for d in T.Parallel(block_DV):
RHS_v[r, d] -= Lmat[r, j] * RHS_v[j, d]

# ── 解完之后才逐行乘 β(Â = diag(β)T)──
for r, d in T.Parallel(C, DK):
W_s[r, d] *= be_f[r]
for r, d in T.Parallel(C, block_DV):
RHS_v[r, d] *= be_f[r]

# ── 伪值 Ṽ = U - W S ──
T.copy(S_f, S_s)
T.gemm(W_s, S_s, Vt, clear_accum=True) # W S
for r, d in T.Parallel(C, block_DV):
Vt[r, d] = RHS_v[r, d] - Vt[r, d]
T.copy(Vt, Vt_s)

W\mathbf{W}U\mathbf{U} 共用同一个 (I+L)(\mathbf{I}+\mathbf{L})拼在一次前向替换里只付一遍串行代价。这一步的并行度是 dk+blockDV=160d_k + \text{block}_{DV} = 160,比 GDN 的 32 好得多–逐通道门在这里又一次因为多了通道维而受益。

4.4 输出与状态更新

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
# ② 块间 + ③ 块内
T.gemm(Q_s, S_s, acc_o, clear_accum=True) # (Q ⊙ e^γ) S
T.copy(Aqk, Aq_c)
T.gemm(Aq_c, Vt_s, acc_o) # A^qk Ṽ

T.copy(acc_o, O_s)
T.copy(O_s, O[bb, s0:s0+C, bh, dv0:dv0+block_DV])

# ① 状态更新:Diag(e^{γ_C}) S + (K ⊙ e^{γ_C-γ_r})^T Ṽ
for d, j in T.Parallel(DK, block_DV):
S_f[d, j] *= Ep_s[C - 1, d] # 逐通道!不是标量
for r, d in T.Parallel(C, DK):
Kp_s[r, d] = T.Cast(dtype,
K_s[r, d] * Ep_s[C - 1, d] * Em_s[r, d]) # e^{γ_C-γ_r}
T.gemm(Kp_s, Vt_s, S_f, transpose_A=True)

最后一段体现了逐通道门与前三篇最直观的差别:

  • GDN:S_f[i,j] *= alpha_pow_C,一个标量乘整个状态
  • KDA:S_f[d,j] *= Ep_s[C-1,d]每个 dkd_k 行乘各自的衰减

状态的不同行按不同速率淡出–这就是"逐通道"在 kernel 层面的全部含义。

e^{\bm{\gamma}_C - \bm{\gamma}_r}Ep_s[C-1,d] * Em_s[r,d] 算,即 eγCeγre^{\gamma_C}\cdot e^{-\gamma_r}。这里复用了已有的两张表,不必重算 exp2;但它也是 eγe^{-\bm{\gamma}} 参与的第三处,进一步说明为什么下界不可省。

4.5 七处次序约束

比 GDN 多两处:

  1. 状态更新排在输出写回之后(右移语义,四篇一致)。
  2. S_f *= Ep_s[C-1,:]T.gemm 之前
  3. Q_s 被原地乘 eγe^{\bm{\gamma}} 后不可再用于块内项–但本篇的块内项恰好也用 qeγ\bm{q}\odot e^{\bm{\gamma}}Aqk\mathbf{A}^{qk} 的定义里就带),所以不需要重载 Q。这是与 GDN 的一处反差,容易照抄出错。
  4. Lmat 只取严格下三角j<ij < i),Aqk 含对角(jij \le i)。
  5. K_s 必须保持原始值Kp\mathbf{Kp}Km\mathbf{Km}、右端项、状态更新四处都从 K_s 派生,任何一处原地修改都会污染后续。本篇的做法是始终写入独立的 Kp_s/Km_sK_s 只读–比 GDN 的"污染两次再重载"更清晰。
  6. β\beta 在前向替换之后乘,不能提前混进右端项A^=diag(β)T\hat{\mathbf{A}} = \operatorname{diag}(\beta)\mathbf{T}Tdiag(β)\mathbf{T}\operatorname{diag}(\beta) 是两种不同配对,各自要搭配 Lci=βiMciL_{ci} = \beta_i M_{ci}βcMci\beta_c M_{ci}β\beta 全相等时这个错误完全隐身。
  7. T.clear(S_f) 在循环外,clear_accum=True 在循环内

第 3、5、6 条都是"照抄上一篇会错"的地方,其中第 6 条我实际踩过。


5. 四篇对照

线性注意力 标量衰减 GDN KDA
衰减 α\alpha 常数 αt\alpha_t 标量 atRdk\bm{a}_t \in \mathbb{R}^{d_k}
删除 Iβkk\mathbf{I}-\beta\bm{k}\bm{k}^\intercal 同 GDN
衰减掩码 M\mathbf{M}(0/1) Γ\GammaC2C^2 表) Γ\GammaC2C^2 表) 无表,融进 GEMM
累积积 编译期常量 CC 步串行,宽度 1 CC 步串行,宽度 dkd_k
前向替换并行度 bDV=32\text{bDV} = 32 dk+bDV=160d_k + \text{bDV} = 160
主要 fragment C2C^2 2C22C^2 3C23C^2 2C2+3Cdk2C^2 + 3Cd_k
可用 CC 128 128 64 64(需 shared 卸载)
数值命门 1/γ1/\gamma 溢出 三角系统条件数 e±γe^{\pm\bm{\gamma}} 双向 + 通道离散
关键前提 累积积按 chunk 重置 k\bm{k} L2 归一化 门控下界 ≥ 0.9

四篇的箭头量语义完全一致q\overleftarrow{\bm{q}}eγre^{\bm{\gamma}_r}k\overrightarrow{\bm{k}}eγCγre^{\bm{\gamma}_C-\bm{\gamma}_r}S\overrightarrow{\mathbf{S}}eγCe^{\bm{\gamma}_C}。从第一篇的"无衰减"到本篇的"逐通道",衰减插入的三个位置从未改变,改变的只是每个位置乘的是标量还是向量。


6. 总结

  1. 逐通道门只改一个符号,却拆掉了前三篇所有便利αtDiag(at)\alpha_t \to \operatorname{Diag}(\bm{a}_t) 让累积积从标量变向量,Γ\GammaC×CC\times C 表变成概念上的 C×C×dkC\times C\times d_k不可能物化,必须融进 GEMM。
  2. 衰减长在 KK 矩阵内部,这是 KDA chunkwise 的核心难点Mci=(kceγcγi)kiM_{ci} = (\bm{k}_c \odot e^{\bm{\gamma}_c-\bm{\gamma}_i})^\intercal\bm{k}_i 的指数依赖三个下标,无法像标量门那样外提成逐元素乘。
  3. 出路是指数可分离eγcγi=eγceγie^{\bm{\gamma}_c-\bm{\gamma}_i} = e^{\bm{\gamma}_c}\odot e^{-\bm{\gamma}_i},于是 M=K~+K~\mathbf{M} = \widetilde{\mathbf{K}}^{+}\widetilde{\mathbf{K}}^{-\intercal} 一次 GEMM 解决(实测等价,误差 1.94×10161.94\times10^{-16})。朴素写法要 C2/2C^2/2 次加权内积、约 26 万次 exp;分离后只需 Cdk=8192C d_k = 8192exp2 加两次 GEMM。
  4. 代价是必须物化 eγe^{-\bm{\gamma}}–第二篇批判过的 1/γ1/\gamma 陷阱的逐通道版本。但这次没有替代方案:不拆就用不上 Tensor Core。问题从"要不要拆"变成"如何让拆分数值安全"。
  5. 门控下界是 kernel 可行性的前提,不是精度调优。实测 C=128C=128、无下界时 fp16 下 eγCe^{\bm{\gamma}_C} 128/128 通道全部归零,状态被彻底清空。下界 0.9 时归零通道数为 0,同时 maxeγ=55\max e^{-\bm{\gamma}} = 55 不溢出–一个数字同时管住了下溢与溢出两头
  6. 通道离散度是逐通道门独有的病a[0.01,0.999]\bm{a}\in[0.01,0.999]C=128C=128γC\bm{\gamma}_C 的通道极差达 47(log 域),即同一 fragment 内最强与最弱通道相差 20 个数量级,任何浮点格式都无法同时表示。下界 0.9 把极差压到 1.74。
  7. sub-chunk 是第二道保险eγe^{-\bm{\gamma}} 的动态范围只取决于 sub-chunk 长度:C=128C=128 整块 1.79×1031.79\times10^3,切 32 后降到 9.02。这与第二篇"累积积按 chunk 重置"同理,只是又降一层。C=64C=64 + 下界 0.9 时不需要。
  8. 逐通道门在两处反而更适合 GPU。累积前缀和:GDN 是 CC 步串行 × 宽度 1,KDA 是 CC 步串行 × 宽度 dkd_k,串行长度不变而并行度从 1 涨到 128。前向替换:W\mathbf{W}U\mathbf{U} 共用 (I+L)(\mathbf{I}+\mathbf{L}) 拼成一次求解,并行度 dk+bDV=160d_k + \text{bDV} = 160,远高于 GDN 的 32。
  9. 寄存器压力换了主角。GDN 的瓶颈是三张 C2C^2 表;KDA 是三张 C×dkC \times d_ke±γe^{\pm\bm{\gamma}}W\mathbf{W}),64×12864\times12864264^2 的两倍。朴素分配约 320 reg 超限,把 eγe^{-\bm{\gamma}}W\mathbf{W} 卸载到 shared memory 后降到约 192。
  10. 三处"照抄上一篇会错":块内项的 Aqk\mathbf{A}^{qk} 定义里本就带 qeγ\bm{q}\odot e^{\bm{\gamma}}不需要像 GDN 那样重载 QK_s 应始终只读、派生量写独立 buffer,而非 GDN 的"污染再重载";β\beta 必须在解完三角系统之后乘diag(β)T\operatorname{diag}(\beta)\mathbf{T}Tdiag(β)\mathbf{T}\operatorname{diag}(\beta) 需搭配不同的 LL 下标,写反时 β\beta 全相等则隐身、β\beta 随机则相对 L2 达 1.3×1011.3\times10^{-1}(我第一次写就踩了这个)。
  11. 三个退化检验,比 GDN 多一个g\bm{g} 各通道同值 → 回到 GDN(检验指数分离是否搞混通道维与位置维)、β0\beta\to0 → 逐通道纯衰减、g0\bm{g}\to\bm{0} → 纯 DeltaNet。第一个最重要且本篇独有。
  12. 衰减的三处落点四篇未变q\overleftarrow{\bm{q}}k\overrightarrow{\bm{k}}S\overrightarrow{\mathbf{S}} 从第一篇到第四篇完全一致,变的只是每处乘标量还是乘向量。这是 GDN 那套箭头记号真正的价值–它把"衰减"隔离成了一个独立于门控形态的层面。

可迁移的启示:把一个标量参数升级成向量,代价从来不在参数量。它会连带改变哪些量能被预计算、哪些表能被物化、哪些运算能进 Tensor Core。KDA 的 αtDiag(at)\alpha_t \to \operatorname{Diag}(\bm{a}_t) 只多了 dkd_k 个数,却让衰减掩码从"一张可复用的表"变成"必须融进 GEMM 的隐式结构",并且逼出了一个模型层面的约束(门控下界)作为 kernel 可行性的前提。当一个数值技巧成为架构设计的必要条件时,它就不再是实现细节。

参考

TileLang 实战:KDA 从零到一–Gated DeltaNet

前两篇分别实现了无遗忘的 chunked 线性注意力,以及带标量衰减的版本。两者的状态更新都只做加法:新的键值对累加进状态,旧信息靠衰减系数被动遗忘。本篇补上最后一块–删除

递推式变成:

St=St1(αt(Iβtktkt))+βtvtkt\mathbf{S}_t = \mathbf{S}_{t-1}\Big(\alpha_t\big(\mathbf{I} - \beta_t \bm{k}_t\bm{k}_t^{\intercal}\big)\Big) + \beta_t\bm{v}_t\bm{k}_t^{\intercal}

比上一篇多出的是 Iβtktkt\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal} 这个广义 Householder 变换。它带来的不是又一个逐元素权重,而是一个矩阵乘在状态右侧–这一改把整个 chunkwise 算法的结构改掉了:块内不再能靠一张下三角权重表解决,需要解一个 C×CC \times C 的三角系统。

上两篇见《Chunked 线性注意力》与《标量衰减》,本文沿用其符号约定与四层参考验证框架。


0. 符号约定

沿用 GDN(arXiv:2412.06464v3)§2.2、§3.1、§3.3 与附录 A 的记号,本篇新增两个:

符号 含义 说明
βt(0,1)\beta_t \in (0,1) 写入强度(writing strength) 也是 delta rule 视角下的学习率
T[t]RC×C\mathbf{T}_{[t]} \in \mathbb{R}^{C \times C} UT 变换矩阵 下三角系统的逆,本篇的核心开销
U~[t]RC×dv\widetilde{\mathbf{U}}_{[t]} \in \mathbb{R}^{C \times d_v} 修正后的 value T\mathbf{T} 作用在 diag(β)V\operatorname{diag}(\beta)\mathbf{V}
W[t]RC×dk\mathbf{W}_{[t]} \in \mathbb{R}^{C \times d_k} 修正后的 key T\mathbf{T} 作用在 diag(β)K\operatorname{diag}(\beta)\mathbf{K}

沿用的:αt\alpha_t 单步衰减、γ[t]r=i=1rαi\gamma^r_{[t]} = \prod_{i=1}^{r}\alpha_i 累积衰减积(按 chunk 重置)、Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 衰减感知因果掩码、CC 块长、[t][t] 块序号、r[1,C]r \in [1,C] 块内位置。

一个前提:论文对 q,k\bm{q}, \bm{k} 做 L2 归一化,即 kt=1\|\bm{k}_t\| = 1。这不只是训练稳定性的考虑–下面 §1.2 会看到它直接决定了 Householder 变换会不会把状态越推越大。


1. 从加法到删除:delta rule 在做什么

1.1 三种更新方式的对照

把三篇的递推式并排放,差异一目了然:

形态 递推式 状态如何变化
线性注意力 St=St1+vtkt\mathbf{S}_t = \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} 只增不减
标量衰减 St=αtSt1+vtkt\mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} 整体等比遗忘
gated delta rule St=St1(αt(Iβtktkt))+βtvtkt\mathbf{S}_t = \mathbf{S}_{t-1}\big(\alpha_t(\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal})\big) + \beta_t\bm{v}_t\bm{k}_t^{\intercal} 定向替换 + 整体遗忘

标量衰减的问题在于它不分对象αt\alpha_t 一乘,所有历史信息按同一比例衰减。要腾出空间写入新内容,只能把无关的旧信息一起冲淡。GDN 论文的说法是,门控擅长「快速擦除」,delta rule 擅长「定向修改」,两者互补。

delta rule 的定向性来自哪里?把它拆开看:

St=St1(St1kt)vtoldkt+(βtvt+(1βt)St1kt)vtnewkt\mathbf{S}_t = \mathbf{S}_{t-1} - \underbrace{(\mathbf{S}_{t-1}\bm{k}_t)}_{\bm{v}_t^{\text{old}}}\bm{k}_t^{\intercal} + \underbrace{\big(\beta_t\bm{v}_t + (1-\beta_t)\mathbf{S}_{t-1}\bm{k}_t\big)}_{\bm{v}_t^{\text{new}}}\bm{k}_t^{\intercal}

读法是:先把 kt\bm{k}_t 这个键上原有的值 vtold=St1kt\bm{v}^{\text{old}}_t = \mathbf{S}_{t-1}\bm{k}_t 减掉,再写入新值 vtnew\bm{v}^{\text{new}}_t,而新值是旧值与目标值的凸组合,βt\beta_t 控制替换的彻底程度。βt1\beta_t \to 1 是完全覆盖,βt0\beta_t \to 0 是不动。

关键在于这个减法只作用在 kt\bm{k}_t 方向上,与 kt\bm{k}_t 正交的记忆分毫不动。这就是「定向」–相比标量衰减的一刀切,delta rule 只擦掉要覆盖的那一条。

1.2 为什么是 Householder,以及 L2 归一化的作用

Iβtktkt\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal} 是广义 Householder 变换。kt=1\|\bm{k}_t\| = 1 时,它的特征值只有两种取值,结构一目了然:

  • 沿 kt\bm{k}_t 方向:特征值 1βt1 - \beta_t
  • kt\bm{k}_t 正交的 dk1d_k - 1 个方向:特征值 11

于是 βt(0,1)\beta_t \in (0,1) 时全部特征值的绝对值都不超过 1,状态每步只会被压缩、不会被放大,递推因此稳定。这正是论文对 k\bm{k} 做 L2 归一化的深层原因–若 kt1\|\bm{k}_t\| \ne 1,沿 kt\bm{k}_t 的特征值变成 1βtkt21 - \beta_t\|\bm{k}_t\|^2βtkt2>2\beta_t\|\bm{k}_t\|^2 > 2 时就会翻到 1-1 以下,递推放大。

顺带一提,论文脚注提到可以放开到 βt(0,2)\beta_t \in (0,2) 以允许负特征值,那是为了解锁状态跟踪能力(state tracking)。本文按 (0,1)(0,1) 处理。

1.3 test-time SGD 视角

论文给了一个很有启发的解释:把状态 S\mathbf{S} 看成一个快速权重矩阵,delta rule 就是在做在线回归的一步梯度下降。目标是 L(St)=12Stktvt2\mathcal{L}(\mathbf{S}_t) = \frac{1}{2}\|\mathbf{S}_t\bm{k}_t - \bm{v}_t\|^2,那么:

StβtL(St)=Stβt(Stktvt)kt=St(Iβtktkt)+βtvtkt\mathbf{S}_t - \beta_t\nabla\mathcal{L}(\mathbf{S}_t) = \mathbf{S}_t - \beta_t(\mathbf{S}_t\bm{k}_t - \bm{v}_t)\bm{k}_t^{\intercal} = \mathbf{S}_t(\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal}) + \beta_t\bm{v}_t\bm{k}_t^{\intercal}

βt\beta_t 就是学习率,αt\alpha_t 就是 weight decay。 这个视角下 gated delta rule 没有任何神秘之处–它是带权重衰减的 test-time SGD。


2. Chunkwise 形式:为什么需要解三角系统

2.1 展开递推:转移矩阵不再是标量

按块展开 rr 步(论文式 10):

S[t]r=S[t]i=1rα[t]i(Iβ[t]ik[t]ik[t]i)F[t]r+i=1rβ[t]iv[t]ik[t]ij=i+1rα[t]j(Iβ[t]jk[t]jk[t]j)G[t]r\mathbf{S}_{[t]}^{r} = \mathbf{S}_{[t]}\underbrace{\prod_{i=1}^{r}\alpha^i_{[t]}\big(\mathbf{I} - \beta^i_{[t]}\bm{k}^i_{[t]}\bm{k}^{i\intercal}_{[t]}\big)}_{\mathbf{F}^r_{[t]}} + \underbrace{\sum_{i=1}^{r}\beta^i_{[t]}\bm{v}^i_{[t]}\bm{k}^{i\intercal}_{[t]}\prod_{j=i+1}^{r}\alpha^j_{[t]}\big(\mathbf{I} - \beta^j_{[t]}\bm{k}^j_{[t]}\bm{k}^{j\intercal}_{[t]}\big)}_{\mathbf{G}^r_{[t]}}

对比上一篇:那里的转移量是标量 αri\alpha^{r-i},可以直接查表。这里是矩阵连乘F\mathbf{F}G\mathbf{G} 都不能靠逐元素权重表达。

衰减部分可以先提出来–αi=γr\prod\alpha_i = \gamma^r 是标量,与 Householder 部分可交换:

F[t]r=γ[t]rP[t]r,P[t]r=i=1r(Iβikiki)\mathbf{F}^r_{[t]} = \gamma^r_{[t]}\,\mathbf{P}^r_{[t]}, \qquad \mathbf{P}^r_{[t]} = \prod_{i=1}^{r}\big(\mathbf{I} - \beta^i\bm{k}^i\bm{k}^{i\intercal}\big)

剩下的 P\mathbf{P} 是纯 DeltaNet 的部分,靠 WY 表示处理–这正是论文 §2.2 的内容,下一节完整展开。

2.2 论文 §2.2:无门控 DeltaNet 的 WY 表示

在处理带门控的版本之前,先把论文 §2.2 那套无门控 DeltaNet 的 chunkwise 推导完整走一遍。GDN 的做法本质上是在这套框架上打补丁,先看清基线,后面的改动才有参照。

本节 αt1\alpha_t \equiv 1(无遗忘门),递推退化为 St=St1(Iβtktkt)+βtvtkt\mathbf{S}_t = \mathbf{S}_{t-1}(\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal}) + \beta_t\bm{v}_t\bm{k}_t^{\intercal}

2.2.1 部分展开:两个连乘(论文式 3)

按块部分展开递推:

S[t]r=S[t](i=1r(Iβ[t]ik[t]ik[t]i)):=P[t]r+i=1rβ[t]iv[t]ik[t]ij=i+1r(Iβ[t]jk[t]jk[t]j):=H[t]r\mathbf{S}^r_{[t]} = \mathbf{S}_{[t]}\underbrace{\Big(\prod_{i=1}^{r}\big(\mathbf{I} - \beta^i_{[t]}\bm{k}^i_{[t]}\bm{k}^{i\intercal}_{[t]}\big)\Big)}_{\textstyle :=\mathbf{P}^r_{[t]}} + \underbrace{\sum_{i=1}^{r}\beta^i_{[t]}\bm{v}^i_{[t]}\bm{k}^{i\intercal}_{[t]}\prod_{j=i+1}^{r}\big(\mathbf{I} - \beta^j_{[t]}\bm{k}^j_{[t]}\bm{k}^{j\intercal}_{[t]}\big)}_{\textstyle :=\mathbf{H}^r_{[t]}}

两个部分的角色不同,值得分清:

形状 含义 结构
P[t]r\mathbf{P}^r_{[t]} dk×dkd_k \times d_k 历史状态的遗忘算子:入口状态经过 rr 步 Householder 后剩下什么 Householder 的纯连乘
H[t]r\mathbf{H}^r_{[t]} dv×dkd_v \times d_k 块内新写入的累积:前 rr 个 token 写进来的内容(互相已扣除重叠) 连乘的加权和

P\mathbf{P} 是纯连乘,H\mathbf{H} 的每一项后面还挂着一截连乘尾巴–两者都不能直接算,C=64C = 64P\mathbf{P} 要 64 个 dk×dkd_k\times d_k 矩阵相乘。WY 表示的作用就是把这两个连乘各自压成一次求和。

2.2.2 经典 WY:把 Householder 连乘压成秩-CC 更新(论文式 4)

这是 Bischof & Van Loan (1985) 的经典结果。核心事实:CC 个 Householder 矩阵的乘积可以写成单位矩阵减去一个秩至多 CC 的修正

P[t]r=Ii=1rw[t]ik[t]iRdk×dk,w[t]r=β[t]r(k[t]ri=1r1w[t]i(k[t]ik[t]r))Rdk\mathbf{P}^r_{[t]} = \mathbf{I} - \sum_{i=1}^{r}\bm{w}^i_{[t]}\bm{k}^{i\intercal}_{[t]} \in \mathbb{R}^{d_k\times d_k}, \qquad \bm{w}^r_{[t]} = \beta^r_{[t]}\Big(\bm{k}^r_{[t]} - \sum_{i=1}^{r-1}\bm{w}^i_{[t]}\big(\bm{k}^{i\intercal}_{[t]}\bm{k}^r_{[t]}\big)\Big) \in \mathbb{R}^{d_k}

为什么成立,看一步归纳就够:

Pr=Pr1(Iβrkrkr)=Pr1βrPr1krwrkr\mathbf{P}^{r} = \mathbf{P}^{r-1}\big(\mathbf{I} - \beta_r\bm{k}_r\bm{k}_r^{\intercal}\big) = \mathbf{P}^{r-1} - \underbrace{\beta_r\mathbf{P}^{r-1}\bm{k}_r}_{\textstyle \bm{w}_r}\bm{k}_r^{\intercal}

于是 wr=βrPr1kr\bm{w}_r = \beta_r\mathbf{P}^{r-1}\bm{k}_r,把 Pr1=Ii<rwiki\mathbf{P}^{r-1} = \mathbf{I} - \sum_{i<r}\bm{w}_i\bm{k}_i^{\intercal} 代入就得到上面那个递推。每多一个 Householder,秩只增加 1–这就是"连乘变求和"的全部内容。

2.2.3 同一个模具:H\mathbf{H} 的 WY 表示(论文式 5)

H\mathbf{H} 的推导结构与 P\mathbf{P} 完全平行

H[t]r=i=1ru[t]ik[t]iRdv×dk,u[t]r=β[t]r(v[t]ri=1r1u[t]i(k[t]ik[t]r))Rdv\mathbf{H}^r_{[t]} = \sum_{i=1}^{r}\bm{u}^i_{[t]}\bm{k}^{i\intercal}_{[t]} \in \mathbb{R}^{d_v\times d_k}, \qquad \bm{u}^r_{[t]} = \beta^r_{[t]}\Big(\bm{v}^r_{[t]} - \sum_{i=1}^{r-1}\bm{u}^i_{[t]}\big(\bm{k}^{i\intercal}_{[t]}\bm{k}^r_{[t]}\big)\Big) \in \mathbb{R}^{d_v}

w\bm{w} 的递推和 u\bm{u} 的递推并排看,会发现它们是同一个式子

wr=βr(kri<rwi(kikr)),ur=βr(vri<rui(kikr))\bm{w}_r = \beta_r\Big(\underline{\bm{k}_r} - \sum_{i<r}\bm{w}_i(\bm{k}_i^{\intercal}\bm{k}_r)\Big), \qquad \bm{u}_r = \beta_r\Big(\underline{\bm{v}_r} - \sum_{i<r}\bm{u}_i(\bm{k}_i^{\intercal}\bm{k}_r)\Big)

只有下划线处不同:一个是 kr\bm{k}_r、一个是 vr\bm{v}_r系数完全一样–这个观察是后面式 (6)(7) 能共用一个 T\mathbf{T} 的全部原因,也是 kernel 里 W\mathbf{W}U\mathbf{U} 能拼进一次前向替换的依据(第四篇 KDA 用的就是这个技巧)。

写成矩阵形式,两个连乘都消失了:

P[t]=IW[t]K[t]Rdk×dk,H[t]=U[t]K[t]Rdv×dk\mathbf{P}_{[t]} = \mathbf{I} - \mathbf{W}_{[t]}^{\intercal}\mathbf{K}_{[t]} \in \mathbb{R}^{d_k\times d_k}, \qquad \mathbf{H}_{[t]} = \mathbf{U}_{[t]}^{\intercal}\mathbf{K}_{[t]} \in \mathbb{R}^{d_v\times d_k}

2.2.4 UT 变换:递推也不必串行(论文式 6、7)

式 (4)(5) 虽然把连乘变成了求和,但 wr\bm{w}_r / ur\bm{u}_r 自身还是串行递推。Joffrain et al. (2006) 的 UT 变换把它变成一次矩阵求逆:

T[t]=[I+strictLower(diag(β[t])K[t]K[t])]1diag(β[t])RC×C\mathbf{T}_{[t]} = \Big[\mathbf{I} + \operatorname{strictLower}\big(\operatorname{diag}(\beta_{[t]})\mathbf{K}_{[t]}\mathbf{K}_{[t]}^{\intercal}\big)\Big]^{-1}\operatorname{diag}(\beta_{[t]}) \in \mathbb{R}^{C\times C}

W[t]=T[t]K[t]RC×dk,U[t]=T[t]V[t]RC×dv\mathbf{W}_{[t]} = \mathbf{T}_{[t]}\mathbf{K}_{[t]} \in \mathbb{R}^{C\times d_k}, \qquad \mathbf{U}_{[t]} = \mathbf{T}_{[t]}\mathbf{V}_{[t]} \in \mathbb{R}^{C\times d_v}

W\mathbf{W}U\mathbf{U} 共用同一个 T\mathbf{T},正是因为 §2.2.3 那两个递推的系数相同。这一步的意义在于:T\mathbf{T} 只跟 K\mathbf{K}β\beta 有关,与 V\mathbf{V} 无关。

2.2.5 代回式 3:可用 Tensor Core 的算法(论文式 8、9)

S[t+1]=S[t]P[t]+H[t]=S[t]+(U[t]W[t]S[t])K[t]Rdv×dkO[t]=Q[t]S[t]+(Q[t]K[t]M)(U[t]W[t]S[t])RC×dv\begin{aligned} \mathbf{S}_{[t+1]} &= \mathbf{S}_{[t]}\mathbf{P}_{[t]} + \mathbf{H}_{[t]} = \mathbf{S}_{[t]} + \big(\mathbf{U}_{[t]} - \mathbf{W}_{[t]}\mathbf{S}_{[t]}^{\intercal}\big)^{\intercal}\mathbf{K}_{[t]} \in \mathbb{R}^{d_v\times d_k} \\[4pt] \mathbf{O}_{[t]} &= \mathbf{Q}_{[t]}\mathbf{S}_{[t]}^{\intercal} + \big(\mathbf{Q}_{[t]}\mathbf{K}_{[t]}^{\intercal}\odot\mathbf{M}\big)\big(\mathbf{U}_{[t]} - \mathbf{W}_{[t]}\mathbf{S}_{[t]}^{\intercal}\big) \in \mathbb{R}^{C\times d_v} \end{aligned}

式 (8) 的化简值得看一眼–SP=S(IWK)=S(WS)K\mathbf{S}\mathbf{P} = \mathbf{S}(\mathbf{I}-\mathbf{W}^\intercal\mathbf{K}) = \mathbf{S} - (\mathbf{W}\mathbf{S}^\intercal)^\intercal\mathbf{K},与 H=UK\mathbf{H} = \mathbf{U}^\intercal\mathbf{K} 合并后 K\mathbf{K} 被提到右侧,括号里剩下 UWS\mathbf{U} - \mathbf{W}\mathbf{S}^\intercal这个量在式 (8) 和式 (9) 里是同一个,算一次用两回。

M\mathbf{M} 是下三角全 1 的因果掩码。注意此处没有任何衰减–这正是 GDN 要改的地方。

2.2.6 数值验证:九个等式逐一核对

dk=6,dv=5,C=7d_k=6, d_v=5, C=7βU(0.1,0.9)\beta \sim \mathcal{U}(0.1,0.9)k\bm{k} 已 L2 归一化,入口状态 S00\mathbf{S}_0 \ne \mathbf{0}(比论文附录 A 的首块假设更严格),fp64:

论文式 内容 max abs 误差
(3) Sr=S0Pr+Hr\mathbf{S}^r = \mathbf{S}_0\mathbf{P}^r + \mathbf{H}^r 4.44×10164.44\times10^{-16}
(4) Pr=Iiwiki\mathbf{P}^r = \mathbf{I} - \sum_i\bm{w}_i\bm{k}_i^\intercal 4.44×10164.44\times10^{-16}
(5) Hr=iuiki\mathbf{H}^r = \sum_i\bm{u}_i\bm{k}_i^\intercal 3.33×10163.33\times10^{-16}
矩阵形式 P=IWK\mathbf{P} = \mathbf{I} - \mathbf{W}^\intercal\mathbf{K} 2.22×10162.22\times10^{-16}
矩阵形式 H=UK\mathbf{H} = \mathbf{U}^\intercal\mathbf{K} 2.78×10162.78\times10^{-16}
(6)(7) W=TK\mathbf{W} = \mathbf{T}\mathbf{K}(UT 变换 vs 式 4 递推) 5.55×10175.55\times10^{-17}
(6)(7) U=TV\mathbf{U} = \mathbf{T}\mathbf{V} 2.22×10162.22\times10^{-16}
(8) S[t+1]\mathbf{S}_{[t+1]} 完整式 4.44×10164.44\times10^{-16}
(9) O[t]\mathbf{O}_{[t]} 完整式 8.88×10168.88\times10^{-16}

另外验证了 P=IWK\mathbf{P} = \mathbf{I}-\mathbf{W}^\intercal\mathbf{K} 的特征值绝对值分别为 0.9918,0.9082,0.5805,0.3131,0.1951,0.06760.9918, 0.9082, 0.5805, 0.3131, 0.1951, 0.0676全部 1\le 1,与 §1.2 说的"每步只压缩不放大"一致。

2.2.7 GDN 改了哪两处

把门控加回来(αt1\alpha_t \ne 1),论文 §3.3 的做法只动了式 (6)(7) 两个地方

无门控(式 6、7) GDN(门控版)
T\mathbf{T} 里的 KK 矩阵 KK\mathbf{K}\mathbf{K}^{\intercal} ΓKK\Gamma \odot \mathbf{K}\mathbf{K}^{\intercal}
W\mathbf{W} 的输入 K\mathbf{K} diag(γ)K=K\operatorname{diag}(\gamma)\mathbf{K} = \overleftarrow{\mathbf{K}}
U\mathbf{U} 的输入 V\mathbf{V} V\mathbf{V}(不变)
式 (8) 的旧状态项 S[t]\mathbf{S}_{[t]} γCS[t]=S\gamma^C\mathbf{S}_{[t]} = \overrightarrow{\mathbf{S}}
式 (8) 的 K\mathbf{K} K\mathbf{K} γCγrK=K\frac{\gamma^C}{\gamma^r}\mathbf{K} = \overrightarrow{\mathbf{K}}
式 (9) 的 Q\mathbf{Q} 与掩码 Q\mathbf{Q}M\mathbf{M} γrQ=Q\gamma^r\mathbf{Q} = \overleftarrow{\mathbf{Q}}Γ\Gamma

实测这个替换的正确性:门控版 S[t+1]\mathbf{S}_{[t+1]} 与逐 token 递归差 3.33×10163.33\times10^{-16};令 α1\alpha\equiv1 时,Γ\Gamma 与下三角全 1 掩码 M\mathbf{M} 完全相等(差 0.0)T\mathbf{T} 与无门控版完全相等(差 0.0)W\mathbf{W} 也完全相等。三个 0.0 说明门控版是无门控版的严格推广,没有引入任何额外近似。

换个角度看Γ\Gamma 就是把式 (9) 那个 0/1 因果掩码 M\mathbf{M} 升级成了"带衰减的因果掩码"。Mij{0,1}\mathbf{M}_{ij} \in \{0,1\} 只回答"jj 能否影响 ii",Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 还回答"影响衰减了多少"。前两篇反复出现的那个 Γ\Gamma,在这个视角下就是 M\mathbf{M} 的自然推广。


2.3 手推一遍:Householder 连乘怎么变成三角系统

直接推 C=4C = 4 的情形最清楚。定义每步的修正量 dr\bm{d}_r,使得递推写成加法形式:

Sr=αrSr1+drkr,dr=βrvrαrβrSr1kr\mathbf{S}^r = \alpha_r\mathbf{S}^{r-1} + \bm{d}_r\bm{k}_r^{\intercal}, \qquad \bm{d}_r = \beta_r\bm{v}_r - \alpha_r\beta_r\mathbf{S}^{r-1}\bm{k}_r

这一步只是把 Sr1(αr(Iβrkk))+βrvk\mathbf{S}^{r-1}(\alpha_r(\mathbf{I}-\beta_r\bm{k}\bm{k}^\intercal)) + \beta_r\bm{v}\bm{k}^\intercal 重新分组,恒等变形。注意 dr\bm{d}_r 里带 αr\alpha_r–删除项作用在已经衰减过的状态上,这个细节写错不会报错、只会算错。

有了加法形式,就能像上一篇那样倒代换(每项系数是「距块尾的步数」):

S[t+1]=γCS[t]+r=1CγCγrdrkr\mathbf{S}_{[t+1]} = \gamma^C\mathbf{S}_{[t]} + \sum_{r=1}^{C}\frac{\gamma^C}{\gamma^r}\,\bm{d}_r\bm{k}_r^{\intercal}

问题在于 dr\bm{d}_r 依赖 Sr1\mathbf{S}^{r-1},而 Sr1\mathbf{S}^{r-1} 又依赖 d1,,dr1\bm{d}_1, \ldots, \bm{d}_{r-1}–串行依赖,无法并行。把 Sr1\mathbf{S}^{r-1} 也倒代换开(利用 αrγr1=γr\alpha_r\gamma^{r-1} = \gamma^r):

αrSr1=γrS[t]+i<rγrγidiki\alpha_r\mathbf{S}^{r-1} = \gamma^r\mathbf{S}_{[t]} + \sum_{i<r}\frac{\gamma^r}{\gamma^i}\bm{d}_i\bm{k}_i^{\intercal}

代回 dr\bm{d}_r 的定义:

dr=βr(vrγrS[t]kri<rγrγi(kikr)di)\bm{d}_r = \beta_r\Big(\bm{v}_r - \gamma^r\mathbf{S}_{[t]}\bm{k}_r - \sum_{i<r}\frac{\gamma^r}{\gamma^i}\big(\bm{k}_i^{\intercal}\bm{k}_r\big)\bm{d}_i\Big)

这是一个下三角线性系统dr\bm{d}_r 只依赖 di (i<r)\bm{d}_i\ (i < r),系数是 βrγrγi(kikr)\beta_r\frac{\gamma^r}{\gamma^i}(\bm{k}_i^\intercal\bm{k}_r)。写成矩阵形式,令 Ari=βrγrγi(kikr)\mathbf{A}_{ri} = \beta_r\frac{\gamma^r}{\gamma^i}(\bm{k}_i^\intercal\bm{k}_r) 的严格下三角部分:

(I+strictLower(A))Δ=diag(β)(Vdiag(γ)KS[t])(\mathbf{I} + \operatorname{strictLower}(\mathbf{A}))\,\Delta = \operatorname{diag}(\beta)\big(\mathbf{V} - \operatorname{diag}(\gamma)\mathbf{K}\mathbf{S}_{[t]}^{\intercal}\big)

其中 ΔRC×dv\Delta \in \mathbb{R}^{C \times d_v} 的第 rr 行是 dr\bm{d}_r。注意 A=diag(β)(ΓKK)\mathbf{A} = \operatorname{diag}(\beta)\big(\Gamma \odot \mathbf{K}\mathbf{K}^{\intercal}\big)衰减感知掩码 Γ\Gamma 在这里第二次出现,这次是嵌在三角系统的系数矩阵里。

2.4 论文附录 A:扩展 WY 表示的归纳证明

上面 §2.3 是"把递推硬拆开"的推法。论文附录 A 给了一个更漂亮的等价路线:先猜出闭式,再用数学归纳法证明。这一节按论文原文复述(GDN 论文附录 A,为减少符号负担,论文同样只考虑首块,即 S0=0\mathbf{S}_0 = \mathbf{0})。

命题(扩展 WY 表示).St\mathbf{S}_t

St=i=1tγtγiuiki,ut=βt(vti=1t1γtγiuikikt)\mathbf{S}_t = \sum_{i=1}^{t}\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}, \qquad \bm{u}_t = \beta_t\Big(\bm{v}_t - \sum_{i=1}^{t-1}\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}\bm{k}_t\Big)

证明.tt 作归纳。

St+1=St(αt+1(Iβt+1kt+1kt+1))+βt+1vt+1kt+1=αt+1(i=1tγtγiuiki)αt+1βt+1(i=1tγtγiuikikt+1kt+1)+βt+1vt+1kt+1=i=1tγt+1γiuiki+βt+1(vt+1i=1tγt+1γiuikikt+1)ut+1kt+1=i=1tγt+1γiuiki+γt+1γt+11ut+1kt+1  =  i=1t+1γt+1γiuiki\begin{aligned} \mathbf{S}_{t+1} &= \mathbf{S}_t\Big(\alpha_{t+1}\big(\mathbf{I} - \beta_{t+1}\bm{k}_{t+1}\bm{k}_{t+1}^{\intercal}\big)\Big) + \beta_{t+1}\bm{v}_{t+1}\bm{k}_{t+1}^{\intercal} \\[4pt] &= \alpha_{t+1}\Big(\sum_{i=1}^{t}\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}\Big) - \alpha_{t+1}\beta_{t+1}\Big(\sum_{i=1}^{t}\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}\bm{k}_{t+1}\bm{k}_{t+1}^{\intercal}\Big) + \beta_{t+1}\bm{v}_{t+1}\bm{k}_{t+1}^{\intercal} \\[4pt] &= \sum_{i=1}^{t}\frac{\gamma_{t+1}}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal} + \underbrace{\beta_{t+1}\Big(\bm{v}_{t+1} - \sum_{i=1}^{t}\frac{\gamma_{t+1}}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}\bm{k}_{t+1}\Big)}_{\textstyle \bm{u}_{t+1}}\bm{k}_{t+1}^{\intercal} \\[4pt] &= \sum_{i=1}^{t}\frac{\gamma_{t+1}}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal} + \underbrace{\frac{\gamma_{t+1}}{\gamma_{t+1}}}_{\textstyle 1}\bm{u}_{t+1}\bm{k}_{t+1}^{\intercal} \;=\; \sum_{i=1}^{t+1}\frac{\gamma_{t+1}}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal} \end{aligned}

\square

第二个等号到第三个等号是全部关键,值得拆开看:

  • αt+1γtγi=γt+1γi\alpha_{t+1}\cdot\frac{\gamma_t}{\gamma_i} = \frac{\gamma_{t+1}}{\gamma_i}–衰减系数被吸收进累积积的比值,这是 γ\gamma 定义为连乘才有的性质,也是四篇一直在用的那条恒等式。
  • 第二项里 αt+1\alpha_{t+1} 同样被吸收进 γt+1γi\frac{\gamma_{t+1}}{\gamma_i},然后整项与第三项合并、共同提出右侧的 kt+1\bm{k}_{t+1}^{\intercal}–括号里剩下的就是 ut+1\bm{u}_{t+1} 的定义。
  • 最后一步只是注意 γt+1/γt+1=1\gamma_{t+1}/\gamma_{t+1} = 1,于是新项能并入求和,下标从 tt 推到 t+1t+1,归纳闭合。

这个证明独立印证了 §2.3 那个坑。 注意第二项的系数是 αt+1βt+1\alpha_{t+1}\beta_{t+1}α\alphaβ\beta 都在,因为删除项 Iβkk\mathbf{I}-\beta\bm{k}\bm{k}^\intercal 作用在已经乘过 αt+1\alpha_{t+1} 的状态上。我在 §2.3 用倒代换推 dr\bm{d}_r 时漏掉这个 αr\alpha_r,refB 就与逐 token 递归差了 2.66×1012.66\times10^{-1}。两条路线在同一个位置要求同一个因子,可以互为校验。

实测验证这个命题(dk=5,dv=4,C=6d_k=5, d_v=4, C=6,fp64,逐步比对 St\mathbf{S}_t):

检验 结果
WY 闭式 vs 逐 token 递推(t=1..6t=1..6 逐步) 最大 1.67×10161.67\times10^{-16}
ut\bm{u}_t 是否等于 §2.3 的修正量 dt\bm{d}_t 1.11×10161.11\times10^{-16}
dt\bm{d}_t 漏掉 αt\alpha_t 偏差 2.99×1022.99\times10^{-2}

第二行说明论文的 ut\bm{u}_t 与我 §2.3 手推的 dt\bm{d}_t 是同一个量,只是推导路径不同:论文归纳法从闭式出发验证,§2.3 从递推倒代换构造。两者都给出同一个下三角系统,下面的 UT 变换对二者通用。

关于跨块:论文只推首块(S0=0\mathbf{S}_0=\mathbf{0})。实际 kernel 里每块的入口状态非零,ut\bm{u}_t 的定义要补上 βtγtS0kt-\beta_t\gamma_t\mathbf{S}_0^{\intercal}\bm{k}_t 这一项,即 §2.3 里那个 γrS[t]kr-\gamma^r\mathbf{S}_{[t]}\bm{k}_r照抄附录 A 而漏掉这项,是我最初 refB 失配(max abs 2.66×1012.66\times10^{-1}、rel L2 6.01×1026.01\times10^{-2})的另一个来源–首块测试全对、第二块开始错,这种症状基本可以直接定位到跨块项。

2.5 UT 变换:把解系统变成矩阵乘

定义(论文 §2.2 与 §3.3):

T[t]=[I+strictLower(diag(β[t])(Γ[t]K[t]K[t]))]1diag(β[t])\mathbf{T}_{[t]} = \Big[\mathbf{I} + \operatorname{strictLower}\big(\operatorname{diag}(\beta_{[t]})(\Gamma_{[t]} \odot \mathbf{K}_{[t]}\mathbf{K}_{[t]}^{\intercal})\big)\Big]^{-1}\operatorname{diag}(\beta_{[t]})

于是修正量一次算出:

Δ=TVU~Tdiag(γ)KWS[t]\Delta = \underbrace{\mathbf{T}\mathbf{V}}_{\widetilde{\mathbf{U}}} - \underbrace{\mathbf{T}\operatorname{diag}(\gamma)\mathbf{K}}_{\overleftarrow{\mathbf{W}}}\mathbf{S}_{[t]}^{\intercal}

这正是论文式 (11)(12) 里的 (U~[t]W[t]S[t])\big(\widetilde{\mathbf{U}}_{[t]} - \overleftarrow{\mathbf{W}_{[t]}}\mathbf{S}_{[t]}^{\intercal}\big)。完整的 chunkwise 算法:

S[t+1]=S[t]+ΔK[t]O[t]=Q[t]S[t]+(Q[t]K[t]Γ[t])Δ\begin{aligned} \mathbf{S}_{[t+1]} &= \overrightarrow{\mathbf{S}_{[t]}} + \Delta^{\intercal}\overrightarrow{\mathbf{K}_{[t]}} \\ \mathbf{O}_{[t]} &= \overleftarrow{\mathbf{Q}_{[t]}}\mathbf{S}_{[t]}^{\intercal} + \big(\mathbf{Q}_{[t]}\mathbf{K}_{[t]}^{\intercal} \odot \Gamma_{[t]}\big)\Delta \end{aligned}

三个箭头量与前两篇完全一致:qr=γrqr\overleftarrow{\bm{q}^r} = \gamma^r\bm{q}^rkr=γCγrkr\overrightarrow{\bm{k}^r} = \frac{\gamma^C}{\gamma^r}\bm{k}^rS=γCS\overrightarrow{\mathbf{S}} = \gamma^C\mathbf{S}

对比上一篇,结构上只有一处变化:块内那一项的右乘对象从 V[t]\mathbf{V}_{[t]} 变成了 Δ\Delta,而 Δ\Delta 需要解一个三角系统才能得到。β0\beta \to 0T0\mathbf{T} \to \mathbf{0}Δ0\Delta \to \mathbf{0},退化为纯衰减;α1\alpha \equiv 1ΓM\Gamma \to \mathbf{M},退化为纯 DeltaNet。

2.6 数值验证

B=2,H=2,N=12,dk=dv=4,C=4B=2, H=2, N=12, d_k=d_v=4, C=4αtU(0.90,0.999)\alpha_t \sim \mathcal{U}(0.90, 0.999)βtU(0.10,0.90)\beta_t \sim \mathcal{U}(0.10, 0.90)k\bm{k} 已 L2 归一化,fp64:

比较 max abs 误差 相对 L2
chunkwise + UT vs 逐 token 递归 8.88×10168.88 \times 10^{-16} 2.49×10162.49 \times 10^{-16}
α1\alpha \equiv 1(退化为纯 DeltaNet) 8.88×10168.88 \times 10^{-16}
β0\beta \to 0(退化为纯衰减) 1.62×10271.62 \times 10^{-27}

两个退化检验都必须做,理由与上一篇同构:α1\alpha \equiv 1Γ\Gamma 退化成 0/1 掩码,Γ\Gamma 里任何指数写错都看不出来;β0\beta \to 0 时整个三角系统消失,UT 部分的错误全部隐身。


3. 三角系统的数值性质

T\mathbf{T} 要求一个 C×CC \times C 矩阵的逆,这是本篇最值得担心的地方。但实际上它比看起来温和得多。

3.1 单位下三角,条件数可控

I+strictLower(A)\mathbf{I} + \operatorname{strictLower}(\mathbf{A})单位下三角矩阵–对角线恒为 1,严格下三角才是 A\mathbf{A}。这有两个直接后果:

  1. 行列式恒为 1,永不奇异。 不存在需要 pivoting 的情形。
  2. 求逆可以用前向替换,不需要通用矩阵求逆。

实测条件数(C=64C = 64k\bm{k} L2 归一化,αtU(0.9,0.999)\alpha_t \sim \mathcal{U}(0.9, 0.999),20 组随机采样):

βt\beta_t 采样区间 cond 中位数 cond 最大
[0.05, 0.2][0.05,\ 0.2] 1.76 1.91
[0.1, 0.9][0.1,\ 0.9] 5.67 6.86
[0.5, 0.99][0.5,\ 0.99] 7.88 9.05
[0.9, 0.999][0.9,\ 0.999] 10.38 11.99
[1.0, 1.99][1.0,\ 1.99] 34.27 40.85

β(0,1)\beta \in (0,1) 时条件数不超过 12,fp16 完全够用。这与上一篇那个「因式分解会溢出」的结论形成有意思的对比:那里是代数变形引入了 1/γ1/\gamma 这种指数增长量,而这里虽然要求逆,但矩阵结构本身保证了良态。

放开到 β(0,2)\beta \in (0,2)(论文脚注提到的负特征值情形)条件数跳到 34,仍可接受,但已需要留意。

3.2 前向替换代替求逆

T\mathbf{T} 从不需要显式求逆。逐行前向替换:

T[r,:]=erj<rA[r,j]T[j,:]\mathbf{T}[r,:] = \bm{e}_r - \sum_{j<r}\mathbf{A}[r,j]\,\mathbf{T}[j,:]

实测与 np.linalg.inv 的差异在 101710^{-17} 量级(三组随机种子分别 7.81,5.55,8.33×10177.81, 5.55, 8.33 \times 10^{-17}),即机器精度。

更进一步,T\mathbf{T} 本身也不必物化。 需要的只是 Δ=TR\Delta = \mathbf{T}\mathbf{R}(其中 R=diag(β)(Vdiag(γ)KS)\mathbf{R} = \operatorname{diag}(\beta)(\mathbf{V} - \operatorname{diag}(\gamma)\mathbf{K}\mathbf{S}^\intercal)),直接对 R\mathbf{R} 做前向替换:

Δ[r,:]=R[r,:]j<rA[r,j]Δ[j,:]\Delta[r,:] = \mathbf{R}[r,:] - \sum_{j<r}\mathbf{A}[r,j]\,\Delta[j,:]

这省掉一个 C×CC \times C 的中间量。代价是引入了 CC 步串行–这是本篇 kernel 与前两篇最本质的区别,下面 §5 会看到它如何限制并行度。


4. 四层参考实现

沿用前两篇的框架:

参考 实现方式 验证目标
A 逐 token 递归 递推式定义 St=St1(αt(Iβtkk))+βtvk\mathbf{S}_t = \mathbf{S}_{t-1}(\alpha_t(\mathbf{I}-\beta_t\bm{k}\bm{k}^\intercal)) + \beta_t\bm{v}\bm{k}^\intercal
B chunkwise + UT 变换 §2.3 的三角系统推导、§2.4 的 WY 表示与式 (11)(12)
C kernel 控制流镜像(切 DV、S\mathbf{S}^\intercal 布局、前向替换) kernel 结构
D 退化检验(α1\alpha \equiv 1 / β0\beta \to 0 与前两篇的一致性

4.0 状态方向

与前两篇一致:论文的状态是 SRdv×dk\mathbf{S} \in \mathbb{R}^{d_v \times d_k},kernel 存转置 SRdk×dv\mathbf{S}^\intercal \in \mathbb{R}^{d_k \times d_v}。本篇多一层要注意的是 ΔRC×dv\Delta \in \mathbb{R}^{C \times d_v}–它的布局与 V\mathbf{V} 相同,所以 ΔK\Delta^\intercal\overrightarrow{\mathbf{K}} 在转置布局下写作 KΔ\overrightarrow{\mathbf{K}}^\intercal\Delta

4.1 参考 A:逐 token 递归

1
2
3
4
5
6
7
8
9
10
11
12
13
14
def ref_A(Q, K, V, alpha, beta):
B, N, H, DK = Q.shape
DV = V.shape[-1]
O = np.zeros((B, N, H, DV))
for b in range(B):
for h in range(H):
S = np.zeros((DV, DK)) # 论文方向
for t in range(N):
k, v, q = K[b, t, h], V[b, t, h], Q[b, t, h]
a, be = alpha[b, t, h], beta[b, t, h]
# 广义 Householder:注意 α 乘在整个括号上
S = S @ (a * (np.eye(DK) - be * np.outer(k, k))) + be * np.outer(v, k)
O[b, t, h] = S @ q
return O

4.2 参考 B:chunkwise + UT 变换

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
def ref_B(Q, K, V, alpha, beta, C):
B, N, H, DK = Q.shape
DV = V.shape[-1]
O = np.zeros((B, N, H, DV))
for b in range(B):
for h in range(H):
S = np.zeros((DV, DK))
for c in range(N // C):
sl = slice(c * C, (c + 1) * C)
Qc, Kc, Vc = Q[b, sl, h], K[b, sl, h], V[b, sl, h]
a, be = alpha[b, sl, h], beta[b, sl, h]

g = np.cumprod(a) # γ^r,chunk 内重置
gC = g[-1]
i = np.arange(C)
Gam = np.where(i[:, None] >= i[None, :],
g[i][:, None] / g[i][None, :], 0.0) # Γ

# UT 变换:T = [I + strictLower(diag(β)(Γ ⊙ K K^T))]^{-1} diag(β)
A = np.diag(be) @ (Gam * (Kc @ Kc.T))
T = np.linalg.inv(np.eye(C) + np.tril(A, -1)) @ np.diag(be)

# Δ = Ũ - \overleftarrow{W} S^T
Delta = T @ Vc - T @ (np.diag(g) @ (Kc @ S.T))

O[b, sl, h] = (Qc * g[:, None]) @ S.T + ((Qc @ Kc.T) * Gam) @ Delta
S = gC * S + Delta.T @ (Kc * (gC / g)[:, None])
return O

与上一篇参考 B 的差别只有两处:多出 T 的构造,以及块内项的右乘对象从 Vc 变成 Delta衰减权重的三处落点完全没变

4.3 参考 C:kernel 控制流镜像

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
def ref_C(Q, K, V, alpha, beta, C, block_DV):
B, N, H, DK = Q.shape
DV = V.shape[-1]
O = np.zeros((B, N, H, DV))
for bv in range(DV // block_DV): # ← grid.x
dv = slice(bv * block_DV, (bv + 1) * block_DV)
for bbh in range(B * H): # ← grid.y
b, h = bbh // H, bbh % H
S = np.zeros((DK, block_DV)) # 转置布局
for c in range(N // C): # ← T.Pipelined
sl = slice(c * C, (c + 1) * C)
Qc, Kc, Vc = Q[b, sl, h], K[b, sl, h], V[b, sl, h, dv]
a, be = alpha[b, sl, h], beta[b, sl, h]
g = np.cumprod(a); gC = g[-1]
i = np.arange(C)
Gam = np.where(i[:, None] >= i[None, :],
g[i][:, None] / g[i][None, :], 0.0)

A = (Kc * be[:, None]) @ Kc.T * Gam # diag(β)(Γ ⊙ K K^T)
Tm = np.linalg.inv(np.eye(C) + np.tril(A, -1))
R = be[:, None] * Vc - be[:, None] * ((g[:, None] * Kc) @ S)
Delta = Tm @ R # C × block_DV

O[b, sl, h, dv] = (Qc * g[:, None]) @ S + ((Qc @ Kc.T) * Gam) @ Delta
S = gC * S + (Kc * (gC / g)[:, None]).T @ Delta
return O

注意 R 的写法:diag(β) 被展开成逐行乘 be[:, None],这与 kernel 里的做法一致(不物化对角矩阵)。

4.4 一致性验证

B=2,H=2,N=12,dk=dv=4,C=4B=2, H=2, N=12, d_k=d_v=4, C=4blockDV=2\text{block}_{DV}=2,fp64:

比较 max abs 误差 相对 L2
B chunkwise + UT vs A 逐 token 8.88×10168.88 \times 10^{-16} 2.49×10162.49 \times 10^{-16}
C kernel 镜像 vs A 逐 token 8.88×10168.88 \times 10^{-16} 2.61×10162.61 \times 10^{-16}
α1\alpha \equiv 1(纯 DeltaNet)B vs A 8.88×10168.88 \times 10^{-16}
β1012\beta \to 10^{-12}(纯衰减)B vs A 1.62×10271.62 \times 10^{-27}

另验证了 T\mathbf{T} 的两种等价写法(diag(β) @ (Γ * KKᵀ)(K * β) @ Kᵀ * Γ)差异 5.55×10175.55 \times 10^{-17},后者省一次对角矩阵构造。


5. TileLang kernel

grid 划分沿用前两篇:切 (bv, bh)(bv,\ b \cdot h),序列轴走 kernel 内的 T.Pipelined。但本篇多了一个 C×CC \times CA 矩阵和 CC 步串行的前向替换,寄存器与并行度都要重新算。

5.0 寄存器账

dk=dv=128d_k = d_v = 128C=64C = 64、128 线程:

fragment 上一篇 本篇 说明
S_f [dk,bDV][d_k, \text{bDV}] 32 reg 32 reg 不变
acc_o [C,bDV][C, \text{bDV}] 16 reg 16 reg 不变
Dtri / Gam [C,C][C,C] 32 reg 32 reg 衰减掩码
A [C,C][C,C] 32 reg 32 reg QK\mathbf{Q}\mathbf{K}^\intercal 复用
Amat [C,C][C,C] 32 reg 新增:三角系统系数
Delta [C,bDV][C, \text{bDV}] 16 reg 新增:修正量
合计 ≈ 112 reg ≈ 160 reg 上限 255

C=64C = 64 时仍有余量。C=128C = 128Gam + A + Amat 三张 C2C^2 表就是 384 reg,直接超限–本篇的 CC 实际上被限制在 64。

5.1 声明与权重预计算

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
32
33
34
35
36
37
38
39
@tilelang.jit(out_idx=[5])
def gdn_chunk(B, H, S, DK, DV, blk=64, block_DV=32,
num_stages=2, threads=128,
dtype=T.float16, accum_dtype=T.float32):
C = blk
NS = T.ceildiv(S, C)

@T.prim_func
def main(
Q: T.Tensor([B, S, H, DK], dtype),
K: T.Tensor([B, S, H, DK], dtype),
V: T.Tensor([B, S, H, DV], dtype),
Alpha: T.Tensor([B, S, H], accum_dtype), # α_t,数据依赖
Beta: T.Tensor([B, S, H], accum_dtype), # β_t,数据依赖
O: T.Tensor([B, S, H, DV], dtype),
):
with T.Kernel(T.ceildiv(DV, block_DV), B * H, threads=threads) as (bv, bbh):
bb, bh = bbh // H, bbh % H
dv0 = bv * block_DV

Q_s = T.alloc_shared([C, DK], dtype)
K_s = T.alloc_shared([C, DK], dtype)
V_s = T.alloc_shared([C, block_DV], dtype)
S_s = T.alloc_shared([DK, block_DV], dtype)
O_s = T.alloc_shared([C, block_DV], dtype)
D_s = T.alloc_shared([C, block_DV], dtype) # Δ 的 f16 副本

S_f = T.alloc_fragment([DK, block_DV], accum_dtype) # loop-carried
acc_o = T.alloc_fragment([C, block_DV], accum_dtype)
A = T.alloc_fragment([C, C], accum_dtype) # Q K^T
A_cast = T.alloc_fragment([C, C], dtype)
Amat = T.alloc_fragment([C, C], accum_dtype) # 三角系统系数
Delta = T.alloc_fragment([C, block_DV], accum_dtype)
R = T.alloc_fragment([C, block_DV], accum_dtype) # 右端项
Gam = T.alloc_fragment([C, C], accum_dtype) # Γ
lg = T.alloc_fragment([C], accum_dtype) # log2 γ^r
g_r = T.alloc_fragment([C], accum_dtype) # γ^r
w_decay = T.alloc_fragment([C], accum_dtype) # γ^C/γ^r
be_f = T.alloc_fragment([C], accum_dtype) # β_r

αt\alpha_t 现在是数据依赖的,累积积必须在 kernel 内算。用 log 域 cumsum 而非 cumprod–上一篇实测 fp16 直接连乘在 C=128C = 128 时完全下溢:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
for i_s in T.Pipelined(NS, num_stages=num_stages):
s0 = i_s * C
T.copy(Q[bb, s0:s0+C, bh, :], Q_s)
T.copy(K[bb, s0:s0+C, bh, :], K_s)
T.copy(V[bb, s0:s0+C, bh, dv0:dv0+block_DV], V_s)

# ── 累积衰减积:log 域前缀和 ──
for r in T.serial(C):
lg[r] = T.log2(Alpha[bb, s0 + r, bh])
for r in T.serial(1, C): # 串行前缀和
lg[r] += lg[r - 1]
for r in T.Parallel(C):
g_r[r] = T.exp2(lg[r]) # γ^r
w_decay[r] = T.exp2(lg[C - 1] - lg[r]) # γ^C/γ^r
be_f[r] = Beta[bb, s0 + r, bh]

# ── Γ[i,j] = γ^i/γ^j (j≤i),log 域相减后取指数 ──
for i, j in T.Parallel(C, C):
Gam[i, j] = T.if_then_else(
j <= i, T.exp2(lg[i] - lg[j]), 0.0)

Γ 的构造沿用上一篇的教训:log 域相减再 exp2,不能先各自取指数再相除,否则 1/γj1/\gamma_j 会溢出。条件求值也不能省成「先全算再掩」,j>ij > i 处指数为正会先溢出成 inf。

5.2 三角系统:构造系数并前向替换

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
# ── Amat = diag(β)(Γ ⊙ K K^T) 的严格下三角 ──
T.gemm(K_s, K_s, Amat, transpose_B=True, clear_accum=True)
for i, j in T.Parallel(C, C):
Amat[i, j] = T.if_then_else(
j < i, be_f[i] * Gam[i, j] * Amat[i, j], 0.0)

# ── 右端项 R = diag(β)(V - diag(γ) K S^T) ──
T.copy(S_f, S_s)
for r, d in T.Parallel(C, DK):
K_s[r, d] *= g_r[r] # diag(γ) K,原地
T.gemm(K_s, S_s, R, clear_accum=True) # (γK) S
for r, d in T.Parallel(C, block_DV):
R[r, d] = be_f[r] * (V_s[r, d] - R[r, d])

# ── 前向替换:Δ[r] = R[r] - Σ_{j<r} Amat[r,j] Δ[j] ──
for r in T.serial(C): # C 步串行,无法避免
for d in T.Parallel(block_DV):
Delta[r, d] = R[r, d]
for j in T.serial(r):
for d in T.Parallel(block_DV):
Delta[r, d] -= Amat[r, j] * Delta[j, d]

前向替换是本篇唯一的串行段。C=64C = 64 时是 64 步,每步内部 block_DV = 32 个元素并行–并行度只有 32,远低于 128 线程。这是 gated delta rule 相比前两篇的固有代价:三角系统的依赖链无法打破。

有两个缓解方向:一是增大 block_DV(但吃寄存器),二是把 CC 切成更小的子块做分块前向替换(blocked forward substitution),用矩阵乘处理子块间的耦合。后者是 fla 库实际采用的做法,但实现复杂度显著上升。

5.3 两项输出与状态更新

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
# ② 跨块项:\overleftarrow{Q} S^T
for r, d in T.Parallel(C, DK):
Q_s[r, d] *= g_r[r] # 位置三:query 侧
T.gemm(Q_s, S_s, acc_o, clear_accum=True)

# ③ 块内项:(Q K^T ⊙ Γ) Δ —— 注意要用未加权的原始 Q、K
T.copy(Q[bb, s0:s0+C, bh, :], Q_s)
T.copy(K[bb, s0:s0+C, bh, :], K_s) # K_s 被 diag(γ) 污染过
T.gemm(Q_s, K_s, A, transpose_B=True, clear_accum=True)
for i, j in T.Parallel(C, C):
A[i, j] *= Gam[i, j]
T.copy(A, A_cast)
T.copy(Delta, D_s)
T.gemm(A_cast, D_s, acc_o)

T.copy(acc_o, O_s)
T.copy(O_s, O[bb, s0:s0+C, bh, dv0:dv0+block_DV])

# ① 状态更新:放在输出写回之后
for i, j in T.Parallel(DK, block_DV):
S_f[i, j] *= T.exp2(lg[C - 1]) # γ^C
for r, d in T.Parallel(C, DK):
K_s[r, d] *= w_decay[r] # 位置二:γ^C/γ^r
T.gemm(K_s, D_s, S_f, transpose_A=True) # \overrightarrow{K}^T Δ

5.4 五处次序约束

比上一篇多一处,全部写错都不报错、只算错:

  1. 状态更新排在输出写回之后SfS_f 在步骤②③被读时代表进块状态 S[t]\mathbf{S}_{[t]}
  2. S_f *= γ^CT.gemm 之前。先衰减旧状态再累加新贡献。
  3. K_s 被复用三次、污染两次。构造 Amat 用原始 K\mathbf{K},算 R 时被乘上 diag(γ)\operatorname{diag}(\gamma),块内项要用原始 K\mathbf{K}(须重载),状态更新时又要乘 γC/γr\gamma^C/\gamma^r这是本篇最容易错的地方–上一篇 K_s 只污染一次。
  4. Amat 必须只取严格下三角j<ij < i,不含对角)。含对角就变成了求解 (I+Afull)(\mathbf{I} + \mathbf{A}_{\text{full}}),对角上的 βrkr2=βr\beta_r\|\bm{k}_r\|^2 = \beta_r 会被重复计入。
  5. T.clear(S_f) 在流水线循环外,clear_accum=True 在循环内

5.5 与上一篇的改动汇总

位置 上一篇 本篇 新增开销
输入 Q, K, V + αt\alpha_t, βt\beta_t 两个张量 2 次 HBM 读 / chunk
累积积 编译期常量 kernel 内 log 域 cumsum CC 步串行前缀和
Γ\Gamma 编译期可算 运行时 exp2(lg[i]-lg[j]) C2C^2exp2
三角系统 Amat 构造 + 前向替换 C2C^2 reg + CC 步串行
块内右乘 V\mathbf{V} Δ\Delta 一次 C×bDVC \times \text{bDV} f16 转换
K 复用 污染 1 次 污染 2 次,重载 1 次 一次 shared 写入

6. 三篇对照

线性注意力 标量衰减 Gated DeltaNet
递推 S+vk\mathbf{S} + \bm{v}\bm{k}^\intercal αS+vk\alpha\mathbf{S} + \bm{v}\bm{k}^\intercal S(α(Iβkk))+βvk\mathbf{S}(\alpha(\mathbf{I}-\beta\bm{k}\bm{k}^\intercal)) + \beta\bm{v}\bm{k}^\intercal
块内掩码 M\mathbf{M}(0/1) Γ\Gamma(衰减感知) Γ\Gamma + 三角系统
块内右乘 V\mathbf{V} V\mathbf{V} Δ\Delta
串行段 前缀和 + 前向替换(各 CC 步)
C2C^2 fragment 1(A\mathbf{A} 2(+ Γ\Gamma 3(+ Amat
可用 CC 128 128(切 DV 后) 64
对应架构 RetNet / Lightning-Attn Gated DeltaNet / KDA

三篇的衰减权重落点完全一致q\overleftarrow{\bm{q}}k\overrightarrow{\bm{k}}S\overrightarrow{\mathbf{S}} 三处,从第一篇到第三篇没有变过。delta rule 加进来的是块内那一项的内容(VΔ\mathbf{V} \to \Delta),而不是衰减的结构。这是 GDN 论文那套箭头记号的价值:它把「衰减」这件事隔离成了一个可以独立理解的层面。


7. 总结

  1. delta rule 的本质是定向替换Iβtktkt\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^\intercal 先减掉 kt\bm{k}_t 键上的旧值 St1kt\mathbf{S}_{t-1}\bm{k}_t,再写入新值 βtvt+(1βt)St1kt\beta_t\bm{v}_t + (1-\beta_t)\mathbf{S}_{t-1}\bm{k}_t;与 kt\bm{k}_t 正交的记忆完全不受影响。这与标量衰减的一刀切互补:门控负责快速擦除,delta rule 负责精确修改。
  2. L2 归一化 k\bm{k} 不只是训练技巧kt=1\|\bm{k}_t\| = 1 时 Householder 变换的特征值是 {1βt}{1}dk1\{1-\beta_t\} \cup \{1\}^{d_k-1}βt(0,1)\beta_t \in (0,1) 时它们的绝对值都不超过 1,状态不会被越推越大;若不归一化,βtkt2>2\beta_t\|\bm{k}_t\|^2 > 2 会让特征值翻到 1-1 以下导致发散。
  3. test-time SGD 视角L=12Skv2\mathcal{L} = \frac{1}{2}\|\mathbf{S}\bm{k}-\bm{v}\|^2 的一步梯度下降就是 delta rule,βt\beta_t 是学习率、αt\alpha_t 是 weight decay。
  4. WY 表示的核心事实:CC 个 Householder 的乘积只是秩至多 CC 的修正Pr=Pr1(Iβrkrkr)=Pr1βrPr1krwrkr\mathbf{P}^r = \mathbf{P}^{r-1}(\mathbf{I}-\beta_r\bm{k}_r\bm{k}_r^\intercal) = \mathbf{P}^{r-1} - \underbrace{\beta_r\mathbf{P}^{r-1}\bm{k}_r}_{\bm{w}_r}\bm{k}_r^\intercal,每多一个 Householder 秩只增 1,于是连乘变求和:P=IWK\mathbf{P} = \mathbf{I}-\mathbf{W}^\intercal\mathbf{K}H=UK\mathbf{H} = \mathbf{U}^\intercal\mathbf{K} 同理。
  5. 矩阵值转移矩阵迫使块内计算变成解三角系统。前两篇的转移量是标量,可以查表;Householder 连乘不行。把递推重写成 Sr=αrSr1+drkr\mathbf{S}^r = \alpha_r\mathbf{S}^{r-1} + \bm{d}_r\bm{k}_r^\intercal 后倒代换,dr\bm{d}_r 的依赖关系构成单位下三角系统,系数矩阵恰是 diag(β)(ΓKK)\operatorname{diag}(\beta)(\Gamma \odot \mathbf{K}\mathbf{K}^\intercal)Γ\Gamma 在这里第二次出现
  6. 两条推导路线互为校验。论文附录 A 的归纳法从闭式 St=iγtγiuiki\mathbf{S}_t = \sum_i\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^\intercal 出发验证,§2.3 的倒代换从递推构造,两者给出同一个量(实测 ut=dt\bm{u}_t = \bm{d}_t,差 1.11×10161.11\times10^{-16})。归纳法第二项的系数 αt+1βt+1\alpha_{t+1}\beta_{t+1} 独立确认了下一条那个必须带的 α\alpha。另注意附录 A 只推首块,跨块要补 βtγtS0kt-\beta_t\gamma_t\mathbf{S}_0^\intercal\bm{k}_t——照抄会出现「首块全对、第二块起错」的症状。
  7. dr\bm{d}_r 的定义里带 αr\alpha_rdr=βrvrαrβrSr1kr\bm{d}_r = \beta_r\bm{v}_r - \alpha_r\beta_r\mathbf{S}^{r-1}\bm{k}_r。删除项作用在已衰减的状态上,漏掉 αr\alpha_r 不会报错,α1\alpha \equiv 1 时也看不出来,只在两者都非退化时表现为数值不符。
  8. 三角系统是良态的,这是好消息I+strictLower(A)\mathbf{I} + \operatorname{strictLower}(\mathbf{A}) 是单位下三角,行列式恒为 1、永不奇异、不需 pivoting。实测 C=64C = 64β(0,1)\beta \in (0,1) 时条件数不超过 12,fp16 够用。放开到 β(0,2)\beta \in (0,2) 升到 34。
  9. T\mathbf{T} 与其逆都不必物化,直接对右端项做前向替换即可,实测与 inv 差异在 101710^{-17} 量级。省一个 C×CC \times C 中间量。
  10. 代价是并行度。前向替换的 CC 步依赖链无法打破,每步只有 block_DV 个元素并行(32 < 128 线程)。加上累积积的串行前缀和,本篇有两段 CC 步串行–这是 gated delta rule 相比前两篇的固有开销。
  11. CC 被限制在 64。三张 C2C^2 的 f32 fragment(Γ\GammaQK\mathbf{Q}\mathbf{K}^\intercalAmat)在 C=128C = 128 时合计 384 reg/thread,超过 255 上限。
  12. K_s 污染两次是最容易错的地方:构造 Amat 用原始 K\mathbf{K},算右端项时乘 diag(γ)\operatorname{diag}(\gamma),块内项要重载原始 K\mathbf{K},状态更新再乘 γC/γr\gamma^C/\gamma^r。上一篇只污染一次。
  13. 两个退化检验都必须做α1\alpha \equiv 1 回到纯 DeltaNet(实测 8.88×10168.88\times10^{-16})、β0\beta \to 0 回到纯衰减(1.62×10271.62\times10^{-27})。前者让 Γ\Gamma 的指数错误隐身,后者让整个 UT 部分隐身,缺一不可。
  14. 衰减的三处落点三篇未变q=γrq\overleftarrow{\bm{q}} = \gamma^r\bm{q}k=γCγrk\overrightarrow{\bm{k}} = \frac{\gamma^C}{\gamma^r}\bm{k}S=γCS\overrightarrow{\mathbf{S}} = \gamma^C\mathbf{S} 从第一篇到第三篇完全一致,delta rule 只改变块内项右乘的内容。

可迁移的启示:引入一个"看起来只是多一项"的机制,实际代价往往不在 FLOPs 而在依赖结构。delta rule 的算术开销并不大(一个 C×CC\times C 的三角系统),但它把块内计算从"纯矩阵乘"变成了"带串行依赖的求解",并行度从 128 掉到 32。评估一个改动时,先问它引入了什么依赖,再问它加了多少乘法。

参考

  • Gated Delta Networks(GDN):Yang, Kautz & Hatamizadeh, Gated Delta Networks: Improving Mamba2 with Delta Rule, arXiv:2412.06464,ICLR 2025。§3.1 给出 gated delta rule、§3.3 给出 chunkwise 算法与 UT 变换;§2.2 给出本文 §2.2 复述的无门控 WY/UT 推导(式 3–9)、附录 A(Extended WY Representation for Gated Delta Rule)给出本文 §2.4 复述的归纳法证明,原文只考虑首块(S0=0\mathbf{S}_0=\mathbf{0}
  • DeltaNet 的硬件高效 chunkwise 算法:Yang et al., 2024b(arXiv:2406.06484)
  • WY 表示:Bischof & Van Loan, 1985;UT 变换:Joffrain et al., 2006
  • delta rule 溯源:Widrow & Hoff, 1960;用于线性 Transformer:Schlag et al., 2021a

TileLang 实战:KDA 从零到一–标量衰减

上一篇实现了不带任何遗忘机制的 chunked 线性注意力,状态单调累加。本篇引入第一个衰减因子:把递推式改为 St=αSt1+vtkt\mathbf{S}_t = \alpha \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^\intercalα\alpha 是一个标量常数。

这个衰减项来自哪里?Vanilla 线性注意力(St=St1+vtkt\mathbf{S}_t = \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^\intercal)在语言建模上远不如 Transformer,为此 Mamba2(Dao & Gu, 2024a)引入了一个数据依赖的逐步遗忘门 αt(0,1)\alpha_t \in (0, 1) 来有选择地丢弃历史信息。本文取它的常数特例 αtα\alpha_t \equiv \alpha–这一档对应的正是 RetNet 与 Lightning-Attention,详见 §1.0。

上一篇见《TileLang 实战:KDA 从零到一–Chunked 线性注意力》,本文沿用其参考实现的验证框架。


0. 符号约定:与 GDN 论文对齐

本文全程采用 GDN(arXiv:2412.06464v3)§2.1 与式 (1)(2) 的记号。这里先把几个容易混淆的符号定下来。

符号 含义 说明
αt(0,1)\alpha_t \in (0,1) 单步衰减系数 本文用常数,记作 αtα\alpha_t \equiv \alpha
γj=i=1jαi\gamma_j = \prod_{i=1}^{j}\alpha_i 累积衰减积 常数情形下 γj=αj\gamma_j = \alpha^{\,j}
Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 衰减感知因果掩码 iji \ge j 时有值,否则 0
CC chunk 长度 代码里对应 BC / blk
[t][t] chunk 序号 Q[t]\mathbf{Q}_{[t]} 即第 tt
r[1,C]r \in [1, C] chunk 位置(1-based) q[t]r:=qtC+r\bm{q}_{[t]}^r := \bm{q}_{tC+r}

两个必须说清的点:

一、γ\gamma 是累积积,不是单步衰减。 这是阅读 GDN 时最容易混淆的一处–很多二次资料把 γ\gamma 当成逐步遗忘系数,但论文里逐步系数是 αt\alpha_tγj\gamma_j 是它们的前缀积。

二、累积积按 chunk 重置,不是从序列起点算。 论文式 (1) 的脚注明确写了 γ[t]j=j=tC+1tC+jαj\gamma_{[t]}^{j} = \prod_{j=tC+1}^{tC+j}\alpha_j,并自认“略微滥用了 γ\gamma 的记号”–每个 chunk 从自己的第一个位置重新开始累乘。这不是细节:若从序列起点算,γj\gamma_j 会随 jj 单调下溢到零(§2.5 有实测),而按 chunk 重置后指数永远不超过 CC,这才是分块算法数值可控的前提。


1. 递推式与分块重写

1.0 Mamba2 的遗忘门:衰减项的来处

上一篇实现的是 vanilla 线性注意力(Katharopoulos et al., 2020),状态只增不减:

St=St1+vtktRdv×dk,ot=StqtRdv\mathbf{S}_t = \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} \in \mathbb{R}^{d_v \times d_k}, \qquad \bm{o}_t = \mathbf{S}_t\bm{q}_t \in \mathbb{R}^{d_v}

它在语言建模上明显弱于 Transformer。GDN 论文 §2.1 给出的诊断很直接:缺少遗忘历史信息的手段。以 Mamba2(Dao & Gu, 2024a)为例,补救方式是在状态上乘一个逐步衰减项(论文原话是 “up to specific parameterization”,即忽略具体参数化细节):

St=αtSt1+vtkt,ot=Stqt\mathbf{S}_t = \alpha_t \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal}, \qquad \bm{o}_t = \mathbf{S}_t\bm{q}_t

其中 αt(0,1)\alpha_t \in (0, 1)tt 变化的、数据依赖的标量衰减项(a data-dependent scalar-valued decay term that varies with tt)。这一个乘法带来两件事:

  • 有界性。无衰减时 St\|\mathbf{S}_t\|tt 无界增长(每步加一个外积);有了 αt<1\alpha_t < 1,状态范数被压在一个稳态附近。
  • 选择性αt0\alpha_t \to 0 可以快速清空状态(适合上下文切换),αt1\alpha_t \to 1 则保持记忆。Mamba 把这类机制称为 selective mechanism,本质是 gated RNN 里的遗忘门在矩阵值状态上的推广。

这个递推结构不是 Mamba2 独有的,同样出现在 Gated RFA、xLSTM、Gated RetNet 中。而当 αt\alpha_t 与数据无关、退化为常数时,该形式就是 RetNet 与 Lightning-Attention。

本文取的正是这个常数情形:

形态 衰减项 对应架构
无衰减 vanilla 线性注意力(上一篇)
常数标量 αtα\alpha_t \equiv \alpha RetNet / Lightning-Attention(本文)
数据依赖标量 αt=f(xt)\alpha_t = f(x_t) Mamba2 / GLA / Gated RetNet

本文取常数是为了简化:α\alpha 作为编译期常量,权重表可以预计算,kernel 改动最小。

1.1 本文的递推式与显式求和

本文在上一篇的基础上给状态加入一个遗忘系数 α(0,1]\alpha \in (0, 1]

St=αSt1+vtktRdv×dk,ot=StqtRdv\mathbf{S}_t = \alpha \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} \in \mathbb{R}^{d_v \times d_k}, \qquad \bm{o}_t = \mathbf{S}_t\bm{q}_t \in \mathbb{R}^{d_v}

展开成显式求和,注意每个 viki\bm{v}_i\bm{k}_i^{\intercal} 被后续每一步各乘一次 α\alpha,从 iitt 共乘 tit - i 次:

St=itαtiviki,ot=itαtivi(kiqt)\mathbf{S}_t = \sum_{i \le t} \alpha^{\,t-i}\, \bm{v}_i\bm{k}_i^{\intercal}, \qquad \bm{o}_t = \sum_{i \le t} \alpha^{\,t-i}\, \bm{v}_i\,(\bm{k}_i^{\intercal}\bm{q}_t)

用累积积写就是论文 §2.1 的形式(γj=αj\gamma_j = \alpha^j,所以 γt/γi=αti\gamma_t/\gamma_i = \alpha^{t-i}):

ot=itγtγivi(kiqt)\bm{o}_t = \sum_{i \le t} \frac{\gamma_t}{\gamma_i}\, \bm{v}_i\,(\bm{k}_i^{\intercal}\bm{q}_t)

对比无衰减时的 ot=itvi(kiqt)\bm{o}_t = \sum_{i \le t} \bm{v}_i(\bm{k}_i^{\intercal}\bm{q}_t),唯一变化是每一项多了权重 αti\alpha^{t-i}–距离越远权重越小,这就是「衰减」的含义。α=1\alpha = 1 时权重恒为 1,退化为不遗忘。

1.2 三处需要插入衰减权重的位置

沿用上一篇的分块框架,序列按 CC 切分为 NCNC 个 chunk。把求和拆成跨块与块内两部分,衰减权重会分别落到三个位置。沿用论文的 chunk 内局部下标 r[1,C]r \in [1, C]

位置一–块内下三角(论文的 Γ[t]\Gamma_{[t]}。同块内 jij \le i,相对距离就是局部下标之差:

(Γ[t])ij=γ[t]iγ[t]j={αijji0j>i(\Gamma_{[t]})_{ij} = \frac{\gamma_{[t]}^{\,i}}{\gamma_{[t]}^{\,j}} = \begin{cases} \alpha^{\,i-j} & j \le i \\ 0 & j > i \end{cases}

上一篇这里是 0/1 因果掩码 M\mathbf{M},本文变成衰减感知掩码 Γ\Gamma注意权重全部落在 (0,1](0, 1] 区间–对角线是 α0=1\alpha^0 = 1,左下角最小值是 αC1\alpha^{C-1}

位置二–每块写入状态时的块尾对齐(论文的 k\overrightarrow{\bm{k}}。跨块状态需要统一的时间基准,取所属 chunk 的末尾。第 rr 个 token 的贡献衰减到块末要乘 γ[t]C/γ[t]r=αCr\gamma_{[t]}^{C}/\gamma_{[t]}^{r} = \alpha^{\,C-r}

k[t]r=γ[t]Cγ[t]rk[t]r,Δ[t]=V[t]K[t]Rdv×dk\overrightarrow{\bm{k}_{[t]}^{r}} = \frac{\gamma_{[t]}^{C}}{\gamma_{[t]}^{r}}\,\bm{k}_{[t]}^{r}, \qquad \Delta_{[t]} = \mathbf{V}_{[t]}^{\intercal}\, \overrightarrow{\mathbf{K}_{[t]}} \in \mathbb{R}^{d_v \times d_k}

这个权重从哪来?用 C=4C = 4 手推一遍最清楚。块内每走一格执行一次递推,S[t]\mathbf{S}_{[t]} 为进块状态、S[t+1]\mathbf{S}_{[t+1]} 为出块状态:

Sr=αrSr1+vrkr,S0=S[t],SC=S[t+1]\mathbf{S}_r = \alpha_r \mathbf{S}_{r-1} + \bm{v}_r\bm{k}_r^{\intercal}, \qquad \mathbf{S}_0 = \mathbf{S}_{[t]},\quad \mathbf{S}_C = \mathbf{S}_{[t+1]}

正向走一遍没什么可看的。有意思的是站在 S4\mathbf{S}_4 这端逐层倒代换,把中间状态一个个拆掉:

S4=α4S3+v4k4=α4α3S2+α4v3k3+v4k4=α4α3α2S1+α4α3v2k2+α4v3k3+v4k4=α1α2α3α4S[t]+α2α3α4v1k1+α3α4v2k2+α4v3k3+v4k4\begin{aligned} \mathbf{S}_4 &= \alpha_4\mathbf{S}_3 + \bm{v}_4\bm{k}_4^{\intercal} \\ &= \alpha_4\alpha_3\mathbf{S}_2 + \alpha_4\bm{v}_3\bm{k}_3^{\intercal} + \bm{v}_4\bm{k}_4^{\intercal} \\ &= \alpha_4\alpha_3\alpha_2\mathbf{S}_1 + \alpha_4\alpha_3\bm{v}_2\bm{k}_2^{\intercal} + \alpha_4\bm{v}_3\bm{k}_3^{\intercal} + \bm{v}_4\bm{k}_4^{\intercal} \\ &= \alpha_1\alpha_2\alpha_3\alpha_4\,\mathbf{S}_{[t]} + \alpha_2\alpha_3\alpha_4\,\bm{v}_1\bm{k}_1^{\intercal} + \alpha_3\alpha_4\,\bm{v}_2\bm{k}_2^{\intercal} + \alpha_4\,\bm{v}_3\bm{k}_3^{\intercal} + \bm{v}_4\bm{k}_4^{\intercal} \end{aligned}

盯着最后一行的系数:v1\bm{v}_1 带的是 α2α3α4\alpha_2\alpha_3\alpha_4注意没有 α1\alpha_1,因为它自己就是第 1 格写入的,写完后一路经受的是"身后"三道门;历史状态 S[t]\mathbf{S}_{[t]} 则带满 α1α2α3α4\alpha_1\alpha_2\alpha_3\alpha_4没有任何一项的因子个数取决于它的绝对位置,全部取决于「离第 CC 格还差几步」。

用累积积改写,公共前缀立刻现形:α2α3α4=γ4/γ1\alpha_2\alpha_3\alpha_4 = \gamma_4/\gamma_1α3α4=γ4/γ2\alpha_3\alpha_4 = \gamma_4/\gamma_2α4=γ4/γ3\alpha_4 = \gamma_4/\gamma_3,而 v4\bm{v}_4 那项是 γ4/γ4=1\gamma_4/\gamma_4 = 1。于是打包收口:

S[t+1]=γ[t]CS[t]+r=1Cγ[t]Cγ[t]rvrkr\mathbf{S}_{[t+1]} = \gamma_{[t]}^{C}\,\mathbf{S}_{[t]} + \sum_{r=1}^{C} \frac{\gamma_{[t]}^{C}}{\gamma_{[t]}^{r}}\, \bm{v}_r\bm{k}_r^{\intercal}

这就是 k[t]r\overrightarrow{\bm{k}_{[t]}^{r}} 的全部内容,也就是论文说的「decaying each vector to the last position」。

位置三–跨块状态递推与 query 侧乘累积衰减积(论文的 S\overrightarrow{\mathbf{S}}q\overleftarrow{\bm{q}}。相邻块之间隔了整块 CC 步,因此状态递推是:

S[t]=γ[t]CS[t]=αCS[t],S[t+1]=S[t]+V[t]K[t]\overrightarrow{\mathbf{S}_{[t]}} = \gamma_{[t]}^{C}\,\mathbf{S}_{[t]} = \alpha^{C}\mathbf{S}_{[t]}, \qquad \mathbf{S}_{[t+1]} = \overrightarrow{\mathbf{S}_{[t]}} + \mathbf{V}_{[t]}^{\intercal}\,\overrightarrow{\mathbf{K}_{[t]}}

而第 rr 个 token 读取这个状态时,要衰减到本块首位置的基准上,乘 γ[t]r=αr\gamma_{[t]}^{r} = \alpha^{r}

q[t]r=γ[t]rq[t]r,O[t]=Q[t]S[t]+(Q[t]K[t]Γ[t])V[t]RC×dv\overleftarrow{\bm{q}_{[t]}^{r}} = \gamma_{[t]}^{r}\,\bm{q}_{[t]}^{r}, \qquad \mathbf{O}_{[t]} = \overleftarrow{\mathbf{Q}_{[t]}}\,\mathbf{S}_{[t]}^{\intercal} + \big(\mathbf{Q}_{[t]}\mathbf{K}_{[t]}^{\intercal} \odot \Gamma_{[t]}\big)\mathbf{V}_{[t]} \in \mathbb{R}^{C \times d_v}

这正是论文式 (1)。三个位置的权重汇总–注意每个权重本质上都是门控的累积连乘,常数情形下才收成 α\alpha 的幂:

位置 论文记号 一般形式(累积积) 常数特例 取值范围 作用对象
块内下三角 (Γ[t])ij(\Gamma_{[t]})_{ij} γ[t]iγ[t]j=r=j+1iαr\dfrac{\gamma_{[t]}^{\,i}}{\gamma_{[t]}^{\,j}} = \prod_{r=j+1}^{i}\alpha_r αij\alpha^{\,i-j} [αC1, 1][\alpha^{C-1},\ 1] C×CC \times C 分数矩阵,逐元素
写入状态 k[t]r\overrightarrow{\bm{k}_{[t]}^{r}} γ[t]Cγ[t]r=u=r+1Cαu\dfrac{\gamma_{[t]}^{C}}{\gamma_{[t]}^{r}} = \prod_{u=r+1}^{C}\alpha_u αCr\alpha^{\,C-r} [αC1, 1][\alpha^{C-1},\ 1] K[t]\mathbf{K}_{[t]} 的行,逐行加权
跨块递推 S[t]\overrightarrow{\mathbf{S}_{[t]}} γ[t]C=u=1Cαu\gamma_{[t]}^{C} = \prod_{u=1}^{C}\alpha_u αC\alpha^{C} 标量 整个状态矩阵
读取状态 q[t]r\overleftarrow{\bm{q}_{[t]}^{r}} γ[t]r=u=1rαu\gamma_{[t]}^{r} = \prod_{u=1}^{r}\alpha_u αr\alpha^{r} [αC, α][\alpha^{C},\ \alpha] Q[t]\mathbf{Q}_{[t]} 的行,逐行加权

四个权重全部 1\le 1,指数均为非正。因果约束 iji \ge j 使 Γij=r=j+1iαr1\Gamma_{ij} = \prod_{r=j+1}^{i}\alpha_r \le 1,四个权重同理–论文的写法全程不需要物化任何大于 1 的量,这是 §3 讨论的前提。


2. 累积衰减积:从连乘到矩阵并行形式

上一节的推导是直接展开求和得到的,但衰减机制有一个更本质的表述方式,GDN 论文用它统一了递归形式与并行形式。

2.1 累积衰减积的定义

把衰减退回 §1.0 那个一般形式–依赖数据的 αt(0,1)\alpha_t \in (0, 1),每个时刻的遗忘强度由输入决定(本文的常数 α\alphaαtα\alpha_t \equiv \alpha 的特例)。定义累积衰减积

γj=i=1jαi\gamma_j = \prod_{i=1}^{j} \alpha_i

γj\gamma_j 的含义是从起点衰减到第 jj 步的总折扣。有了它,递推式的展开可以写得非常紧凑。展开 St=αtSt1+vtkt\mathbf{S}_t = \alpha_t \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal}

St=it(r=i+1tαr)viki=itγtγiviki\mathbf{S}_t = \sum_{i \le t} \Big( \prod_{r=i+1}^{t} \alpha_r \Big) \bm{v}_i\bm{k}_i^{\intercal} = \sum_{i \le t} \frac{\gamma_t}{\gamma_i}\, \bm{v}_i\bm{k}_i^{\intercal}

中间那个连乘 r=i+1tαr\prod_{r=i+1}^{t} \alpha_r 正好是两个累积积的比值 γt/γi\gamma_t / \gamma_i–这是累积积定义的全部价值:把「从 iitt 的区间连乘」化归为「两个前缀量之比」,于是任意区间的衰减都可以由一个前缀数组 O(1)O(1) 查得,不必对每个 (i,t)(i, t) 对重新连乘。

上面写的是全序列版本。到了分块算法里,累积积要按 chunk 重置,即论文式 (1) 脚注的 γ[t]j=j=tC+1tC+jαj\gamma_{[t]}^{\,j} = \prod_{j=tC+1}^{tC+j}\alpha_j。这一步不是为了好看–全序列累乘的 γj\gamma_j 会随 jj 单调下溢到零(§2.5 有实测:N=512N = 512α[0.5,0.9]\alpha \in [0.5, 0.9] 时 fp32 已进入非正规数),而按 chunk 重置后指数永远不超过 CC。后文出现 γ\gamma 时默认指 chunk 内的版本。

2.2 两种等价形式

代入 ot=Stqt\bm{o}_t = \mathbf{S}_t\bm{q}_t,同一个结果可以写成两种形式(论文 §2.1):

向量形式(vector form)–逐时刻递归,对应推理阶段:

St=αtSt1+vtkt,ot=Stqt=itvi(γtγikiqt)\mathbf{S}_t = \alpha_t \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal}, \qquad \bm{o}_t = \mathbf{S}_t\bm{q}_t = \sum_{i \le t} \bm{v}_i\Big(\frac{\gamma_t}{\gamma_i}\,\bm{k}_i^{\intercal}\bm{q}_t\Big)

矩阵并行形式(matrix parallel form)–整块一次算出,对应训练与 prefill:

O=((QK)Γ)V,Γij={γiγjij0i<j\mathbf{O} = \big( (\mathbf{Q} \mathbf{K}^{\intercal}) \odot \Gamma \big) \mathbf{V}, \qquad \Gamma_{ij} = \begin{cases} \dfrac{\gamma_i}{\gamma_j} & i \ge j \\[2mm] 0 & i < j \end{cases}

Γ\Gamma 是一个衰减感知因果掩码(decay-aware causal mask)–把上一篇的 0/1 因果掩码 M\mathbf{M} 换成了衰减比值。验证第 (i,j)(i,j) 元素:

[(QK)Γ]ij=γiγj(kjqi)(ji)\big[(\mathbf{Q} \mathbf{K}^{\intercal}) \odot \Gamma\big]_{ij} = \frac{\gamma_i}{\gamma_j} (\bm{k}_j^{\intercal}\bm{q}_i) \quad (j \le i)

与 §2.1 展开式逐项一致。这个形式的价值是把逐 token 的递归变成一次稠密 GEMM 加一次逐元素乘,完全并行,这正是「parallel within each chunk」的含义。递归形式与并行形式的这种等价在 Mamba2 中被称为状态空间对偶性(state space duality, SSD)

Γ\Gamma 的双重身份:一个哈达玛积承载两重语义

Γ\Gamma 常被笼统理解为“一个带衰减的掩码”,但它实际上是两个正交语义的乘积,只是恰好能合并成一张表:

Γ=M因果性:能不能看D遗忘门:看得多清,Mij={1ij0i<j,Dij=γiγj\Gamma = \underbrace{\mathbf{M}}_{\text{因果性:能不能看}} \odot \underbrace{\mathbf{D}}_{\text{遗忘门:看得多清}}, \qquad \mathbf{M}_{ij} = \begin{cases}1 & i \ge j\\ 0 & i<j\end{cases}, \quad \mathbf{D}_{ij} = \frac{\gamma_i}{\gamma_j}

  • M\mathbf{M}离散的、与数据无关的结构约束–token ii 不能看到未来的 j>ij > i。这是 causal 语言模型的硬性要求,α\alpha 取什么值都不影响它。
  • D\mathbf{D}连续的、由门控决定的权重–iijj 相距越远,γi/γj\gamma_i/\gamma_j 越小。这才是遗忘门起作用的地方。

将该矩阵可视化后,结构十分清晰:下三角,且每个值只由 token 序号差 iji-j 决定

图中三点值得对着公式确认:

  1. 对角线恒为 α0=1\alpha^0 = 1–自己看自己,零衰减。本文的 tril\operatorname{tril} 含对角(§1.2 位置一的 jij \le i),与图一致。
  2. 同一行自右向左指数缩小–同行内 ii 固定,jj 越小则间隔 iji-j 越大、权重越小。所以「衰减」在矩阵上表现为沿对角线方向的等值带iji-j 相同的格子取值相同。
  3. 上三角 j>ij > i 整块置 0–那是还没发生的 token。这就是「下三角」的全部含义。

面板 ② 用 α=0.5\alpha = 0.5C=4C = 4 给了可验算的数字:γ=(0.5, 0.25, 0.125, 0.0625)\gamma = (0.5,\ 0.25,\ 0.125,\ 0.0625),格 (4,1)(4,1)γ4/γ1=0.0625/0.5=0.125=α3\gamma^4/\gamma^1 = 0.0625/0.5 = 0.125 = \alpha^3。整张表我用 fp64 复核过,与 αij\alpha^{i-j} 逐格一致、对角恒为 1。

图中「整张表没有一处幂运算,每格只是一次除法」一句还揭示了比值写法的另一重好处,同时解释了为什么论文写 γi/γj\gamma_i/\gamma_j 而不直接写 αij\alpha^{i-j}:后者只在 α\alpha 为常数时才成立,一旦门控逐 token 变化,「唯一底数」就不存在了,r=j+1iαr\prod_{r=j+1}^{i}\alpha_r 无法写成任何数的幂;而比值形式原样成立。本文因 α\alpha 为常数而两种写法皆可,论文则必须采用比值形式。

回到实现。 上面的分解还隐含一个容易忽略的陷阱:若把 M\mathbf{M}D\mathbf{D} 真的分开算再相乘,D\mathbf{D} 的上三角是大于 1 的i<ji<j 时指数为正)。取 α=0.9\alpha=0.9C=4C=4

D=[11.1111.2351.3720.911.1111.2350.810.911.1110.7290.810.91]   M  Γ=[10000.91000.810.9100.7290.810.91]\mathbf{D} = \begin{bmatrix}1&\mathbf{1.111}&\mathbf{1.235}&\mathbf{1.372}\\0.9&1&\mathbf{1.111}&\mathbf{1.235}\\0.81&0.9&1&\mathbf{1.111}\\0.729&0.81&0.9&1\end{bmatrix} \ \xrightarrow{\ \odot\ \mathbf{M}\ }\ \Gamma = \begin{bmatrix}1&0&0&0\\0.9&1&0&0\\0.81&0.9&1&0\\0.729&0.81&0.9&1\end{bmatrix}

C=4C = 4 时上三角最大才 1.372,但 C=64C = 64α=0.8\alpha = 0.8 时右上角是 0.863=1.27×1060.8^{-63} = 1.27 \times 10^{6},早已溢出 fp16。于是:

写法 中间量范围 结果
先构造完整 D\mathbf{D},再乘 M\mathbf{M} [αC1, α(C1)][\alpha^{C-1},\ \alpha^{-(C-1)}] fp16 下上三角可能已 inf,inf × 0 = NaN
只在 iji \ge j 处求值,否则直接置 0 (0, 1](0,\ 1] 安全

因此 kernel 里必须写成条件求值而非「算完再掩」(§5.1 的 T.if_then_else(j <= i, ...) 就是这个原因)–哪怕数学上 MD\mathbf{M} \odot \mathbf{D} 与「只算下三角」完全等价。M\mathbf{M} 的存在恰好保证了 Γ\Gamma 中每个有效元素的指数 ij0i-j \ge 0,这才是「Γij1\Gamma_{ij} \le 1」这个结论成立的前提。

换个角度理解:softmax 注意力里掩码是加 -\infty(因为后面要过 exp),线性注意力里掩码是乘 0(因为没有 exp,直接置零即可)–但两者都不该「先算全量再掩」,原因分别是数值溢出与计算浪费。

2.3 对照:Mamba2 官方 kernel 怎么写这个衰减

上面几节的结论不是纸上推演–TileLang 官方 examples/linear_attention/example_mamba_chunk_state.py 就是这么做的。它算的是 Mamba2 chunkwise 的状态更新那一步,对应本文 §1.2 的位置二(K\overrightarrow{\mathbf{K}},把整块贡献衰减到块尾)。

参考实现一行 einsum 就说完了:

1
2
3
decay_states = torch.exp((dA_cumsum[:, :, :, -1:] - dA_cumsum))
return torch.einsum("bclhn,bhcl,bhcl,bclhp->bchpn",
B, decay_states, dt, x)

对着本文的记号读,dA_cumsum 就是 logγ[t]r\log \gamma_{[t]}^{r}ΔA\Delta A 的累积和,AA 已经是负的),于是:

decay_states[r]=exp(logγ[t]Clogγ[t]r)=γ[t]Cγ[t]r\texttt{decay\_states}[r] = \exp\big(\log\gamma_{[t]}^{C} - \log\gamma_{[t]}^{r}\big) = \frac{\gamma_{[t]}^{C}}{\gamma_{[t]}^{r}}

这正是 §1.2 位置二的 k[t]r\overrightarrow{\bm{k}_{[t]}^{r}} 权重,一字不差。Mamba2 里 k\bm{k} 的角色由 B 承担、v\bm{v}x 承担,dt 是离散化步长(本文的常数情形没有这一项)。

kernel 里对应的三行是这样的:

1
2
3
4
5
6
7
8
9
p = 1.44269504                     # log2(e)

dA_cs_last[0] = dA_cumsum[batch_idx, bz, chunk_idx, chunk_size - 1] # log γ_C,循环外取一次
...
for i in T.Parallel(block_K):
scale[i] = T.exp2(dA_cs_last[0] * p - dA_cumsum_local[i] * p) * dt_local[i]
for i, j in T.Parallel(block_M, block_K):
xt_local[i, j] = x_local[j, i] * scale[j] # 逐行乘衰减比值,同时转置
T.gemm(xt_local, B_shared, acc_o) # 加权后才进 Tensor Core

四处细节和本文的结论逐条对上:

官方写法 为什么这么写
输入是 dA_cumsum(log 域累积和),不是 γ\gamma 本身 连乘会下溢:fp16 直接 cumprodC=128C = 128α[0.5, 0.9]\alpha \in [0.5,\ 0.9] 时归零。log 域改成加法就没这个问题
exp2(log γ_C - log γ_r)先在 log 域相减再取指数 相减的结果因 rCr \le C0\le 0exp2 输出恒 1\le 1。若先各自取指数再相除,就要物化 1/γr1/\gamma_r 这个大于 1 的量
p = 1.44269504exe^x 转成硬件 exp2 1.44269504=log2e1.44269504 = \log_2 e,于是 ex=2xlog2ee^x = 2^{x\log_2 e}exp2 有单指令实现,pow 通常展开成多条
dA_cs_lastT.Pipelined 循环外读一次 logγC\log\gamma_C 是整块共用的常量,每轮重取只是白读一次 shared memory

最值得注意的是第二行。它完全可以写成 exp2(log γ_C * p) / exp2(log γ_r * p)–数学上等价,还能把 γC\gamma_C 提到循环外省一次减法。但那就等于物化了 1/γr1/\gamma_r,正是 §3 实测会 NaN 的写法。官方选择在 log 域相减,指数结果因 rCr \le Clogγ\log\gamma 单调递减而恒 0\le 0exp2 输出恒不大于 1,不存在溢出可能。

顺带一个和本文互补的观察:官方 kernel 把 scale 直接乘到了 xt_local(即 v\bm{v} 侧)而非 B_sharedk\bm{k} 侧)。两者数学等价–衰减是标量,挂在哪一侧都行;选 x 是因为那一步本来就要做转置 xt_local[i,j] = x_local[j,i]将衰减乘法合并进转置的同一个 T.Parallel 中,省去一趟对 shared memory 的读写。这类「把逐元素操作合并进已有的数据搬运」是 tile 编程中常见的优化手法。


3. 定量分析:因式分解在 fp16 下的失效点

论文的写法(先 GEMM 再逐元素乘 Γ\Gamma)中间量恒在 (0,1](0,1],是数值安全的。但块内下三角 (Γ[t])ij=αij(\Gamma_{[t]})_{ij} = \alpha^{\,i-j} 存在一个看似有利的代数变形–指数可以拆开:

αij=αiαj\alpha^{\,i-j} = \alpha^{\,i} \cdot \alpha^{-j}

于是块内那一项可以写成:

(Q[t]K[t]Γ[t])=tril((Q[t]αi)(K[t]αj))\big(\mathbf{Q}_{[t]}\mathbf{K}_{[t]}^{\intercal} \odot \Gamma_{[t]}\big) = \operatorname{tril}\big( (\mathbf{Q}_{[t]} \odot \alpha^{\,i}) (\mathbf{K}_{[t]} \odot \alpha^{-j})^{\intercal} \big)

两种实现方式的差别:

方案 做法 块内额外开销 权重数值范围
比值形式(论文) 先 GEMM 得分数,再逐元素乘 Γ\Gamma 一次 C×CC \times C 逐元素乘 + 一张 C2C^2 权重表 (0,1](0, 1]
因式分解 先给 QQKK 的行分别乘上衰减因子,再单次 GEMM 两次 C×dkC \times d_k 逐行加权,无 C2C^2 开销 αj\alpha^{-j} 最大 α(C1)\alpha^{-(C-1)}

因式分解在 FLOPs 与寄存器占用上都更优:省掉一张 C×CC \times C 的权重表,逐元素乘的规模从 C2C^2 降到 2Cdk2Cd_kC=64C = 64dk=64d_k = 64 时前者是 4096 次乘法,后者 8192 次–乘法次数反而增加,但权重表不必常驻寄存器,这在 C=128C = 128 时是实质性的压力缓解。

问题在数值范围。αj\alpha^{-j}大于 1 的量,且随 jj 指数增长:

α\alpha α31\alpha^{-31}CC=32) α63\alpha^{-63}CC=64) α127\alpha^{-127}CC=128) fp16 安全上限 fp32 安全上限
0.99 1.371.37 1.881.88 3.583.58 C<1103C < 1103 C<8827C < 8827
0.95 4.904.90 25.325.3 675675 C<216C < 216 C<1729C < 1729
0.90 26.226.2 763763 6.47×1056.47 \times 10^{5} C<105C < 105 C<842C < 842
0.80 1.01×1031.01 \times 10^{3} 1.27×1061.27 \times 10^{6} 2.03×10122.03 \times 10^{12} C<49C < 49 C<397C < 397
0.50 2.15×1092.15 \times 10^{9} 9.22×10189.22 \times 10^{18} 1.70×10381.70 \times 10^{38} C<15C < 15 C<127C < 127

fp16 上限是 65504。α=0.8\alpha = 0.8C=64C = 64α63=1.27×106\alpha^{-63} = 1.27 \times 10^6 已经溢出;α=0.5\alpha = 0.59.22×10189.22 \times 10^{18} 连 fp32 都接近极限。

3.1 实测:单块块内计算的精度对比

单块块内计算的精度对比(C=64C = 64dk=dv=64d_k = d_v = 64,fp16 存储 + fp32 累加,20 组随机输入取中位数,参考值为 fp64 精确计算):

α\alpha 比值形式 相对 L2 因式分解 相对 L2 α(C1)\alpha^{-(C-1)}
0.99 4.39×1044.39 \times 10^{-4} 4.15×1044.15 \times 10^{-4} 1.881.88
0.95 4.57×1044.57 \times 10^{-4} 4.14×1044.14 \times 10^{-4} 25.325.3
0.90 4.47×1044.47 \times 10^{-4} 4.15×1044.15 \times 10^{-4} 763763
0.80 4.60×1044.60 \times 10^{-4} NaN 1.27×1061.27 \times 10^{6}
0.50 4.14×1044.14 \times 10^{-4} NaN 9.22×10189.22 \times 10^{18}

结论清晰:

  1. α0.9\alpha \ge 0.9 时两者精度相当,因式分解甚至略优(少一次逐元素乘引入的舍入);
  2. α0.8\alpha \le 0.8 时因式分解在 fp16 下彻底失效,产生 NaN 而非精度下降–αj\alpha^{-j} 溢出成 inf,随后 inf 乘 0 得 NaN;
  3. 比值形式的误差在全部 α\alpha 取值下稳定在 4.14.6×1044.1\text{--}4.6 \times 10^{-4},与 α\alpha 无关。

比值形式的误差之所以稳定,是因为 Γij\Gamma_{ij} 恒在 (0,1](0,1]这是一个与 α\alpha 无关的数值保证。实测各 α\alpha 下衰减矩阵的取值:

α\alpha CC 最大值 最小非零值 下溢为 0 的比例
0.99 128 1.000 2.79×1012.79 \times 10^{-1} 0.0%
0.90 128 1.000 1.55×1061.55 \times 10^{-6} 0.0%
0.50 128 1.000 5.88×10395.88 \times 10^{-39} 0.0%

即使 α=0.5\alpha = 0.5C=128C = 128,最小权重 5.88×10395.88 \times 10^{-39} 在 fp32 下仍是正规数,无下溢。下溢比溢出安全得多:权重下溢为 0 意味着「这个远距离贡献可以忽略」,语义上正确;而溢出为 inf 会污染整行输出。

3.2 比值形式自身的下溢边界

比值形式也不是完全没有约束。Γ\Gamma 本身若用 fp16 存储,αC1\alpha^{C-1} 可能低于 fp16 最小正规数 6.10×1056.10 \times 10^{-5}

α\alpha C=64C=64 fp16 状态 C=128C=128 fp16 状态
0.99 5.31×1015.31 \times 10^{-1} 正常 2.79×1012.79 \times 10^{-1} 正常
0.95 3.95×1023.95 \times 10^{-2} 正常 1.48×1031.48 \times 10^{-3} 正常
0.90 1.31×1031.31 \times 10^{-3} 正常 1.55×1061.55 \times 10^{-6} 非正规数
0.80 7.85×1077.85 \times 10^{-7} 非正规数 4.93×10134.93 \times 10^{-13} 下溢为 0
0.50 1.08×10191.08 \times 10^{-19} 下溢为 0 5.88×10395.88 \times 10^{-39} 下溢为 0

处理方式很简单:权重表用 f32 fragment 保存,只在喂入 MMA 前把乘完的分数矩阵降到 f16。分数矩阵本身量级正常,降精度无损。这也是下面 kernel 采用的做法。

一句话总结:因式分解能省一次 C2C^2 逐元素乘,但把中间量范围从 (0,1](0,1] 推到 [1,α(C1)][1, \alpha^{-(C-1)}]α0.8\alpha \le 0.8 时 fp16 直接 NaN–这个 FLOPs 优化不值得,论文的比值形式应原样保留。


4. 四层参考实现

沿用上一篇的三层框架,本文增加一层专门验证因式分解写法:

参考 实现方式 验证目标
A 逐 token 递归 递推式定义 St=αSt1+vtkt\mathbf{S}_t = \alpha\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal}
B 分块向量化,三处衰减权重 §1.2 三个权重位置的推导
C 逐块独立重算,模拟 grid kernel 控制流
D 因式分解版块内计算 §3 两种写法的代数等价性

4.0 一个必要的转置:论文的 Rdv×dk\mathbb{R}^{d_v \times d_k} vs kernel 的 Rdk×dv\mathbb{R}^{d_k \times d_v}

这里要先交代一个容易造成混乱的差异。论文(以及本文 §1–§3 的全部推导)的状态是 SRdv×dk\mathbf{S} \in \mathbb{R}^{d_v \times d_k}

St=αtSt1+vtktRdv×dk,ot=StqtRdv\mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} \in \mathbb{R}^{d_v \times d_k}, \qquad \bm{o}_t = \mathbf{S}_t\bm{q}_t \in \mathbb{R}^{d_v}

注意是 vk\bm{v}\bm{k}^{\intercal}(value 在外、key 转置在内)、读出是 Sq\mathbf{S}\bm{q}(状态左乘 query)。而下面所有代码里的 S 存的都是它的转置

S =^ SRdk×dv\texttt{S} \ \widehat{=}\ \mathbf{S}^{\intercal} \in \mathbb{R}^{d_k \times d_v}

于是两边的写法逐行对应:

论文(SRdv×dk\mathbf{S} \in \mathbb{R}^{d_v \times d_k} 代码(S Rdk×dv\in \mathbb{R}^{d_k \times d_v}
St=αSt1+vtkt\mathbf{S}_t = \alpha\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} S = g*S + outer(k, v)
ot=Stqt\bm{o}_t = \mathbf{S}_t\bm{q}_t o = q @ S
S[t+1]=S[t]+V[t]K[t]\mathbf{S}_{[t+1]} = \overrightarrow{\mathbf{S}_{[t]}} + \mathbf{V}_{[t]}^{\intercal}\overrightarrow{\mathbf{K}_{[t]}} S = g**C * S + Kd.T @ V
O[t]=Q[t]S[t]+\mathbf{O}_{[t]} = \overleftarrow{\mathbf{Q}_{[t]}}\mathbf{S}_{[t]}^{\intercal} + \dots acc = (Qb * w_query) @ S + ...

为何不直接按论文的方向存?因为 dk×dvd_k \times d_v 布局下,跨块读出就是 T.gemm(Q_s, S_s, acc_o)[C,dk]×[dk,dv][C, d_k] \times [d_k, d_v]不需任何 transpose 标志;若按论文方向存,每次读出都要 transpose_B=True。同理状态更新 VK\mathbf{V}^{\intercal}\overrightarrow{\mathbf{K}} 在转置布局下变成 KV\overrightarrow{\mathbf{K}}^{\intercal}\mathbf{V},正好是 transpose_A=True 一个标志就能表达的形式。这是纯工程选择,不影响任何数学结论–但读代码时必须心里有这个转置,否则会觉得和论文对不上。

4.1 参考 A:逐 token 递归

1
2
3
4
5
6
7
8
9
10
11
def ref_recurrent(Q, K, V, g):
"""严格照抄递推式:S = g*S + k v^T"""
B, N, H, D = Q.shape
O = torch.empty(B, N, H, D, dtype=torch.float64, device=Q.device)
for b in range(B):
for h in range(H):
S = torch.zeros(D, D, dtype=torch.float64, device=Q.device)
for t in range(N):
S = g * S + torch.outer(K[b, t, h].double(), V[b, t, h].double())
O[b, t, h] = Q[b, t, h].double() @ S
return O

与上一篇的唯一差别是 S = g * S + ... 而非 S += ...

4.2 参考 B:分块向量化

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
def ref_chunked(Q, K, V, g, BC):
"""三处衰减权重分别对应 §1.2 的位置一、二、三"""
B, N, H, D = Q.shape
NC = N // BC
Qc = Q.permute(0, 2, 1, 3).double().reshape(B, H, NC, BC, D)
Kc = K.permute(0, 2, 1, 3).double().reshape(B, H, NC, BC, D)
Vc = V.permute(0, 2, 1, 3).double().reshape(B, H, NC, BC, D)

i = torch.arange(BC, dtype=torch.float64, device=Q.device)
# 位置一:块内衰减下三角 g^(i-j),j>i 处置 0
decay_tri = torch.where(i[:, None] >= i[None, :],
g ** (i[:, None] - i[None, :]),
torch.zeros((), dtype=torch.float64, device=Q.device))
w_decay = g ** (BC - 1 - i) # 位置二:写入状态时衰减到块尾
w_query = g ** (i + 1) # 位置三:读取状态时补上距上块末尾的步数

# 每块对状态的贡献(已按块尾对齐)
contrib = torch.einsum("bhcmd,bhcmv->bhcdv", Kc * w_decay[:, None], Vc)

# 跨块前缀状态递推:S_prev[c] = g^BC * S_prev[c-1] + contrib[c-1]
states_prev = torch.zeros(B, H, NC, D, D, dtype=torch.float64, device=Q.device)
for c in range(1, NC):
states_prev[:, :, c] = g ** BC * states_prev[:, :, c - 1] + contrib[:, :, c - 1]

O_inter = torch.einsum("bhcnd,bhcdv->bhcnv", Qc * w_query[:, None], states_prev)
A = torch.einsum("bhcnd,bhcmd->bhcnm", Qc, Kc) * decay_tri # 比值形式
O_intra = torch.einsum("bhcnm,bhcmv->bhcnv", A, Vc)

return (O_inter + O_intra).reshape(B, H, N, D).permute(0, 2, 1, 3).contiguous()

跨块状态在上一篇可以用 torch.cumsum 一次算完,本文不行–cumsum 是无权重的前缀和,而本文的递推带系数 αC\alpha^{C}。这里保留显式循环;若要向量化,需改用加权前缀和的写法:把 Δ[t]\Delta_{[t]} 先除以 αtC\alpha^{tC}cumsum、最后乘回,代价是又引入 αtC\alpha^{-tC} 这个溢出源,tt 大时比 §3 的块内因式分解更危险。这是同一个取舍在跨块层面的重演–它的本质仍是 §3 那个 1/γ1/\gamma 物化问题,只是尺度从 chunk 内部换到了 chunk 之间。

4.3 参考 C:kernel 控制流镜像

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
def ref_kernel_mimic(Q, K, V, g, BC):
"""外层循环模拟 grid,内层 range(bx) 模拟 T.Pipelined(bx)"""
B, N, H, D = Q.shape
NC = N // BC
O = torch.zeros(B, N, H, D, dtype=torch.float64, device=Q.device)
i = torch.arange(BC, dtype=torch.float64, device=Q.device)
decay_tri = torch.where(i[:, None] >= i[None, :],
g ** (i[:, None] - i[None, :]),
torch.zeros((), dtype=torch.float64, device=Q.device))
w_decay = g ** (BC - 1 - i)
w_query = g ** (i + 1)

for bz in range(B):
for by in range(H):
for bx in range(NC):
sl = slice(bx * BC, (bx + 1) * BC)
# ① 流式累加:每轮先整块衰减 g^BC,再累加本块贡献
S = torch.zeros(D, D, dtype=torch.float64, device=Q.device)
for c in range(bx):
cs = slice(c * BC, (c + 1) * BC)
Kd = K[bz, cs, by, :].double() * w_decay[:, None]
S = g ** BC * S + Kd.T @ V[bz, cs, by, :].double()

Qb = Q[bz, sl, by, :].double()
Kb = K[bz, sl, by, :].double()
Vb = V[bz, sl, by, :].double()
acc = (Qb * w_query[:, None]) @ S # ② 跨块
acc += ((Qb @ Kb.T) * decay_tri) @ Vb # ③ 块内
O[bz, sl, by, :] = acc
return O

与上一篇参考 C 的差别只有两处:状态累加从 S += 变成 S = g**BC * S + ...,以及三处权重乘法。循环结构完全没变,这正是 §1 结论「分块恒等式结构不变」的代码体现。

4.4 四层参考的一致性验证

B=2,H=2,N=12,D=4,BC=4,α=0.9B=2, H=2, N=12, D=4, BC=4, \alpha=0.9,fp64(numpy 复现):

比较 max abs 误差 相对 L2
B 分块向量化 vs A 逐 token 递归 3.55×10153.55 \times 10^{-15} 1.82×10161.82 \times 10^{-16}
C kernel 结构镜像 vs A 逐 token 递归 2.67×10152.67 \times 10^{-15} 1.83×10161.83 \times 10^{-16}
D 因式分解 vs A 逐 token 递归 3.55×10153.55 \times 10^{-15} 1.98×10161.98 \times 10^{-16}
B(α=1.0\alpha = 1.0)vs 上一篇参考 A 3.55×10153.55 \times 10^{-15} 1.62×10161.62 \times 10^{-16}

另外用逐 token 变化的 αtU(0.85, 0.999)\alpha_t \sim \mathcal{U}(0.85,\ 0.999) 复验过一遍(同规模,fp64),确认 §2.2 的矩阵并行形式在门控非常数时同样成立:矩阵形式 vs 向量形式相对 L2 2.03×10162.03 \times 10^{-16}Γij\Gamma_{ij} 取值范围 [0.628, 1.000][0.628,\ 1.000] 全部 1\le 1

前三行确认四份实现数学等价,误差均在 fp64 机器精度量级。第四行是退化检验:令 α=1\alpha = 1 应当精确回到上一篇的无衰减实现,这条验证能同时捕获三处权重中任何一处的指数写错–例如把 αCr\alpha^{\,C-r} 误写成 αCr+1\alpha^{\,C-r+1}α=1\alpha = 1 时两者都是 1,退化检验通过但 α=0.9\alpha = 0.9 时参考 B 与 A 不符。两条验证必须都做。


5. TileLang kernel 的改动

沿用上一篇 §6.4.2 的 grid 划分:(bv, bh)(bv,\ b\cdot h),序列轴不进 grid、退回 kernel 内的 T.Pipelined 顺序循环,状态作为 loop-carried fragment 常驻寄存器,每 chunk 只做一次 KVK^\top V。本文在此基础上新增一张 f32 权重表与两个权重向量。

先说清为什么必须切 DV。本文比上一篇多出一个 C×CC \times CDtri,而它在 DV 方向是共享的(衰减权重只跟 chunk 内位置有关,与 value 通道无关),因此不随 DV 切分而缩小。dk=dv=128d_k = d_v = 128C=64C = 64、128 线程下点一遍账:

fragment 不切 DV blockDV=32\text{block}_{DV} = 32
S_f [128,128][128,128] → 128 reg [128,32][128,32] → 32 reg
acc_o [64,128][64,128] → 64 reg [64,32][64,32] → 16 reg
Dtri + A 2×[64,64]2 \times [64,64] → 64 reg 64 reg(不变
合计 ≈ 256 reg/thread,已超 255 上限 ≈ 112 reg/thread

不切 DV 在这个配置下直接编译不过。 代价是 Dtri 被每个 bvbv block 各算一遍–DV/blockDV=4DV/\text{block}_{DV} = 4 时同一张 4096 元素的表算了 4 遍,合计 16384 次 exp2。这是纯冗余计算,但它不占额外寄存器,且 exp2 是单指令,相比编译不过是划算的交换。

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
32
33
34
35
@tilelang.jit(out_idx=[3])
def linattn_decay_chunk(B, H, S, DK, DV, alpha, blk=64, block_DV=32,
num_stages=2, threads=128,
dtype=T.float16, accum_dtype=T.float32):
C = blk
NS = T.ceildiv(S, C)
alpha_pow_C = alpha ** C # 编译期常量

@T.prim_func
def main(
Q: T.Tensor([B, S, H, DK], dtype),
K: T.Tensor([B, S, H, DK], dtype),
V: T.Tensor([B, S, H, DV], dtype),
O: T.Tensor([B, S, H, DV], dtype),
):
# grid:DV 切块 × (batch·head);序列轴不在这里
with T.Kernel(T.ceildiv(DV, block_DV), B * H, threads=threads) as (bv, bbh):
bb, bh = bbh // H, bbh % H
dv0 = bv * block_DV

Q_s = T.alloc_shared([C, DK], dtype)
K_s = T.alloc_shared([C, DK], dtype)
V_s = T.alloc_shared([C, block_DV], dtype)
S_s = T.alloc_shared([DK, block_DV], dtype)
O_s = T.alloc_shared([C, block_DV], dtype)

S_f = T.alloc_fragment([DK, block_DV], accum_dtype) # loop-carried
acc_o = T.alloc_fragment([C, block_DV], accum_dtype)
A = T.alloc_fragment([C, C], accum_dtype)
A_cast = T.alloc_fragment([C, C], dtype)
# 新增:衰减矩阵 Γ,f32 保存以避免 §3.2 的 fp16 下溢
Dtri = T.alloc_fragment([C, C], accum_dtype)
# 新增:两个长度为 C 的权重向量(位置二、位置三)
w_decay = T.alloc_fragment([C], accum_dtype) # α^(C-1-m),写入状态用
w_query = T.alloc_fragment([C], accum_dtype) # α^(i+1),读取状态用

5.1 预计算衰减权重

三张权重表都在 T.Pipelined 之前算一次,整条序列共用:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
lg = T.log2(T.Cast(accum_dtype, alpha))       # log2(α),三处复用

# 位置一:Dtri[i,j] = α^(i-j) for j<=i, else 0
for i, j in T.Parallel(C, C):
Dtri[i, j] = T.if_then_else(
j <= i,
T.exp2(lg * T.Cast(accum_dtype, i - j)),
0.0)

# 位置二:写入状态时,第 m 行要衰减到块尾 → α^(C-1-m)
for m in T.Parallel(C):
w_decay[m] = T.exp2(lg * T.Cast(accum_dtype, C - 1 - m))

# 位置三:读取状态时,第 i 行距上块末尾 i+1 步 → α^(i+1)
for i in T.Parallel(C):
w_query[i] = T.exp2(lg * T.Cast(accum_dtype, i + 1))

三者的指数都取自 §1.2 的权重表:

变量 形状 元素 指数含义 取值范围
Dtri[i,j] C×CC \times C αij\alpha^{\,i-j}jij \le i,否则 0) 同块内两 token 的间隔 [αC1, 1][\alpha^{C-1},\ 1]
w_decay[m] CC αC1m\alpha^{\,C-1-m} mm 行到块尾的步数 [αC1, 1][\alpha^{C-1},\ 1]
w_query[i] CC αi+1\alpha^{\,i+1} ii 行到上块末尾的步数 [αC, α][\alpha^{C},\ \alpha]

注意 w_decayw_query方向相反w_decay[C-1] = α^0 = 1(块尾那一行本身就是基准,不需衰减),而 w_query[0] = α^1(块首那一行距上块末尾也有 1 步)。这两处若把指数写反或差一,α=1\alpha = 1 时完全看不出来–§4.4 的退化检验正是为此设计的。

指数用 exp2(log2(α) · k) 而非 pow(α, k)exp2log2 都有单指令硬件实现,而 pow 通常展开成多条指令。

DtriT.if_then_else 不能省成「先全算再掩」–如 §2.2 所述,j>ij > i 处的 αij\alpha^{i-j} 指数为正,C=64C = 64α=0.8\alpha = 0.8 时右上角已达 1.27×1061.27 \times 10^{6},fp16 下溢出成 inf 后再乘 0 会得到 NaN。这里的条件求值同时是正确性和数值安全两重保障。

5.2 序列内循环:三处权重的落点

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
32
33
        T.clear(S_f)                     # 循环外只清一次,整条序列共用状态

for i_s in T.Pipelined(NS, num_stages=num_stages):
s0 = i_s * C
T.copy(Q[bb, s0:s0+C, bh, :], Q_s) # 全 DK
T.copy(K[bb, s0:s0+C, bh, :], K_s) # 全 DK
T.copy(V[bb, s0:s0+C, bh, dv0:dv0+block_DV], V_s) # 仅 dv 竖条

# ② 跨块项:此刻 S_f 仍是 S_prev(更新在循环末尾)
T.copy(S_f, S_s)
for i, d in T.Parallel(C, DK):
Q_s[i, d] *= w_query[i] # 位置三:query 侧
T.gemm(Q_s, S_s, acc_o, clear_accum=True)

# ③ 块内项:要用未加权的原始 Q,故重新载入
T.copy(Q[bb, s0:s0+C, bh, :], Q_s)
T.gemm(Q_s, K_s, A, transpose_B=True, clear_accum=True)
for i, j in T.Parallel(C, C):
A[i, j] *= Dtri[i, j] # 位置一:逐元素乘 Γ
T.copy(A, A_cast)
T.gemm(A_cast, V_s, acc_o)

T.copy(acc_o, O_s)
T.copy(O_s, O[bb, s0:s0+C, bh, dv0:dv0+block_DV]) # 已是最终值

# ① 状态更新:放在输出写回之后 = 右移语义
for i, j in T.Parallel(DK, block_DV):
S_f[i, j] *= alpha_pow_C # 位置三:整块衰减 α^C
for m, d in T.Parallel(C, DK):
K_s[m, d] *= w_decay[m] # 位置二:衰减到块尾
T.gemm(K_s, V_s, S_f, transpose_A=True)

return main

四处次序约束,写错任何一处都不会报错、只会算错:

  1. 状态更新必须排在输出写回之后SfS_f 在步骤②被读取时代表 S[t]\mathbf{S}_{[t]}(进块状态),本块自己的贡献由 tril\operatorname{tril} 那一项负责。提前更新就变成了 inclusive 前缀,块内项被重复计入。
  2. S_f *= alpha_pow_C 必须在 T.gemm(K_s, V_s, S_f) 之前。递推式是 S[t+1]=αCS[t]+Δ[t]\mathbf{S}_{[t+1]} = \alpha^{C}\mathbf{S}_{[t]} + \Delta_{[t]},先衰减旧状态再累加新贡献;若顺序颠倒,本块贡献会被多衰减一次。
  3. 步骤③必须重载 Q。步骤②已把 Q_s 原地乘上 αr\alpha^{r},而块内项用的是未加权的原始 Q[t]\mathbf{Q}_{[t]}。原地加权省了一个 buffer,代价是重载一次;若寄存器有余量,另开 Q_scaled 更安全。
  4. T.clear(S_f) 在循环外,clear_accum=True 在循环内。状态要跨迭代累加,只能循环外清一次;acc_oA 每 chunk 都是全新的,进了流水线循环就不能再用 T.clear(会被排到错误阶段),必须靠 clear_accum=True 在 MMA 那一刻覆盖累加器。

K_s 在循环末尾被原地乘过 w_decay,而步骤③用的是同一个 K_s–这里没有冲突,因为③在①之前执行。但若调整次序或复用,必须重新载入:本文的 K_s 残留内容已被乘过权重,比上一篇「残留内容不确定」更危险。

5.3 与上一篇 kernel 的改动汇总

位置 上一篇 本文 新增开销
grid (NS,H,B)(NS, H, B) (DV/blockDV, BH)(DV/\text{block}_{DV},\ B \cdot H) 序列不进 grid,状态常驻寄存器
权重表 Dtri[BC, BC] f32 fragment C2C^2 个 f32 寄存器
状态累加 T.gemm 直接累加 S_f *= alpha_pow_C 再 gemm dkdvd_k d_v 次乘法 / 轮
K 写入状态 原样 逐行乘 αCr\alpha^{\,C-r} C×dkC \times d_k 次乘法 / 轮
Q 读取状态 原样 逐行乘 αr\alpha^{r} C×dkC \times d_k 次乘法
块内掩码 if_then_else 置 0 逐元素乘 Dtri C2C^2 次乘法
Q 复用 步骤②③共用 步骤③必须重载 一次 shared 写入

切 DV 后 S_f[128,128][128,128] 降到 [128,32][128,32],寄存器压力最大的反而不是 Dtridk=dv=128d_k = d_v = 128C=64C = 64blockDV=32\text{block}_{DV} = 32 时逐项数:

fragment 上一篇 本文 变化
S_f [128,128][128,128] → 128 reg [128,32][128,32] → 32 reg 下降 96 reg
acc_o [64,128][64,128] → 64 reg [64,32][64,32] → 16 reg 下降 48 reg
Dtri + A 2×[64,64]2 \times [64,64] → 64 reg 新增 64 reg
合计 约 224 reg 约 112 reg ↓ 一半

切 DV 省下的寄存器刚好覆盖 Dtri 的开销,还有余量。C=128C = 128Dtri + A 升到 256 reg,即使切 DV 到 32 也超过 255 上限–那才是因式分解唯一真正有吸引力的地方(它不需要 C2C^2 权重表),但 §3.1 的 NaN 结论表明这个吸引力不成立。


6. 衰减带来的新维度:有效记忆长度

α\alpha 引入了一个上一篇不存在的语义参数–状态的遗忘速度。定义有效记忆长度为权重衰减到 10310^{-3} 所需的 token 数,即 αL=103\alpha^{L} = 10^{-3}

L=ln103lnαL = \frac{\ln 10^{-3}}{\ln \alpha}

α\alpha α64\alpha^{64} 有效记忆长度
0.999 9.38×1019.38 \times 10^{-1} 6904 token
0.99 5.26×1015.26 \times 10^{-1} 687 token
0.95 3.75×1023.75 \times 10^{-2} 135 token
0.90 1.18×1031.18 \times 10^{-3} 66 token
0.50 5.42×10205.42 \times 10^{-20} 10 token

这张表解释了为什么实际模型里的门控值普遍接近 1:α=0.9\alpha = 0.9 的有效记忆只有 66 token,连一个 chunk(C=64C = 64)都刚刚覆盖,长程依赖完全丢失。α0.99\alpha \ge 0.99 恰好落在 §3 表格中因式分解仍然安全的区间–这解释了为什么部分实现敢做这个分解:它们隐含假设了门控接近 1。

这个假设在标量、常数门控下可以接受–反正全局就一个 α\alpha,调到 0.99 以上就行。但它是一个隐含假设,不是数值保证:一旦门控变成数据依赖的(尤其是每个通道各自学一个),训练中总会有部分值降到 0.9 以下以实现快速遗忘,「全部接近 1」就不再成立,因式分解的溢出从例外变成常态。这也是 §3 实测中 α0.8\alpha \le 0.8 即溢出的原因。结论不变:不要做那个分解,不要依赖「门控总是接近 1」这个前提。

一句话总结α\alpha 的安全区间(接近 1)与有用区间(提供实际遗忘能力)方向相反,常数门控下两者尚可兼顾,但这个兼顾靠的是假设而非保证。


7. 数值验证

验证层次沿用上一篇结构,新增两项:

  1. 四层参考互验(fp64):A/B/C/D 两两对照,误差应在 101510^{-15} 量级;
  2. 退化检验α=1\alpha = 1 时参考 B 应精确回到上一篇实现–这条能捕获三处权重的指数偏移错误;
  3. fp16 失效点复现α0.8\alpha \le 0.8 时因式分解写法应产生 NaN,确认 §3.1 结论;
  4. kernel vs 参考 C:fp16 输入 + f32 累加,阈值取相对 L2 <2×102< 2 \times 10^{-2}
  5. 延迟对比走 CUDA event 中位数。

第 1、2、3 层已在 numpy fp64 上完成(见 §3.1 与 §4.4 表格)。第 4、5 层需要 CUDA 设备,待真卡跑通后单独补充实测数据,此处不做性能推测。

本文 kernel 的预期主导误差源与上一篇相同–T.copy(S_f, S_s) 的 f32→f16 降精度。但衰减带来一个有利变化:αC\alpha^{C} 每轮把旧状态压缩,早期块的累积误差也随之衰减,因此误差不再像上一篇那样随 bxbx 单调增长,而是趋于一个稳态。α=0.9\alpha = 0.9C=64C = 64αC=1.18×103\alpha^{C} = 1.18 \times 10^{-3},约三轮之后早期误差已不可见。衰减机制顺带改善了数值稳定性,这是一个反直觉但合理的副作用。


8. 总结

  1. 衰减项来自 Mamba2 的遗忘门。Vanilla 线性注意力的状态只增不减,语言建模上明显弱于 Transformer;Mamba2 的补救是乘一个数据依赖的标量 αt(0,1)\alpha_t \in (0,1),即 St=αtSt1+vtkt\mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal},同时带来有界性(状态范数不再无界增长)与选择性αt0\alpha_t \to 0 快速清空、1\to 1 保持记忆)。本文取其常数特例 αtα\alpha_t \equiv \alpha–这正是 RetNet / Lightning-Attention 的形态。
  2. 标量衰减不改变分块恒等式的结构,只在三个位置插入指数权重(即论文式 (2) 的三个箭头量):块内下三角 (Γ[t])ij=αij(\Gamma_{[t]})_{ij} = \alpha^{\,i-j}、写入状态时的块尾对齐 k\overrightarrow{\bm{k}}αCr\alpha^{\,C-r}、跨块递推 S\overrightarrow{\mathbf{S}}αC\alpha^{C} 与 query 侧 q\overleftarrow{\bm{q}}αr\alpha^{r}。四个权重全部 1\le 1,这是本文数值安全的根本原因。
  3. 累积衰减积 γj=i=1jαi\gamma_j = \prod_{i=1}^{j} \alpha_i 是统一两种形式的代数工具。它把区间连乘 r=i+1tαr\prod_{r=i+1}^{t} \alpha_r 化归为前缀量之比 γt/γi\gamma_t / \gamma_i,于是递归的向量形式与并行的矩阵形式 O=((QK)Γ)V\mathbf{O} = ((\mathbf{Q}\mathbf{K}^{\intercal}) \odot \Gamma)\mathbf{V} 可以相互转换(即 Mamba2 所谓的状态空间对偶性,实测相对 L2 2.03×10162.03 \times 10^{-16})。需注意论文中 αt\alpha_t 为单步衰减、γj\gamma_j 为累积积,二者不可混淆;常数情形下 γj=αj\gamma_j = \alpha^{\,j},此时 Γij=αij\Gamma_{ij} = \alpha^{\,i-j}。另外累积积按 chunk 重置(论文式 (1) 脚注),不是从序列起点算。
  4. 论文的衰减矩阵 Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 本身是数值安全的–因果约束 iji \ge j 使它等于 r=j+1iαr1\prod_{r=j+1}^{i}\alpha_r \le 1,全程不需要物化大于 1 的量(实测 Γij[0.628, 1.000]\Gamma_{ij} \in [0.628,\ 1.000])。论文的箭头记号 q\overleftarrow{\bm{q}}k\overrightarrow{\bm{k}}S\overrightarrow{\mathbf{S}} 把「衰减到首 / 末位置」直接编码在箭头方向上,基准点选得当就能保证所有指数非正。
  5. 真正的陷阱是把比值因式分解成 γi(1/γj)\gamma_i \cdot (1/\gamma_j)。这个分解能把两步合成单次 GEMM 并免去 C2C^2 权重表,但要求物化 1/γj1/\gamma_j,把中间量从 (0,1](0,1] 推到 [1, 1/γC][1,\ 1/\gamma_C]。实测(C=64C=64,fp16 存储加 fp32 累加):α0.9\alpha \ge 0.9 时两者精度相当(4.14.6×1044.1\text{--}4.6 \times 10^{-4}),α0.8\alpha \le 0.8 时分解写法产生 NaNα=0.8\alpha = 0.8α63=1.27×106\alpha^{-63} = 1.27 \times 10^{6},已远超 fp16 上限 65504。data-dependent αt[0.8, 0.95]\alpha_t \in [0.8,\ 0.95]C=128C = 1281/γj1/\gamma_j 最大达 3.40×1083.40 \times 10^{8},同样溢出。结论是保留论文的比值形式,不做分解。
  6. 累积积本身也必须在 log 域计算。fp16 直接 cumprodC=128C = 128α[0.5, 0.9]\alpha \in [0.5,\ 0.9] 时已完全下溢为 0,fp32 在 C=512C = 512 时进入非正规数区间;logγj=ijlogαi\log \gamma_j = \sum_{i \le j} \log \alpha_i 是线性增长的负数,表示范围安全。这就是门控全程存 log 值、用 exp2 还原的原因。
  7. 退化检验是本文新增的关键验证手段:令 α=1\alpha = 1 应精确回到上一篇实现(实测相对 L2 1.62×10161.62 \times 10^{-16})。但它无法单独定案–三处权重中任何一处的指数偏移在 α=1\alpha = 1 时都不可见,必须同时做 α=0.9\alpha = 0.9 的四层参考互验(实测 1.822.03×10161.82\text{--}2.03 \times 10^{-16})。
  8. 四处次序约束(写错都不报错、只算错):状态更新必须排在输出写回之后(否则 inclusive 前缀导致块内项重复计入);S_f *= alpha_pow_C 必须在 T.gemm 之前(否则本块贡献被多衰减一次);步骤③必须重载 Q(步骤②已把 Q_s 原地乘了 αr\alpha^{r});T.clear(S_f) 在流水线循环外、clear_accum=True 在循环内。
  9. 必须切 DV:本文新增的 C×CC \times C 权重表在 DV 方向共享、不随切分缩小。dk=dv=128d_k = d_v = 128C=64C = 64 时不切 DV 约需 256 reg/thread,已超 255 上限直接编译不过;切到 blockDV=32\text{block}_{DV} = 32 降至约 112 reg。代价是 Dtri 被每个 bvbv block 各算一遍,属纯冗余但不占额外寄存器。
  10. 有效记忆长度暴露了一个内在矛盾α\alpha 的数值安全区间(接近 1)与实际有用区间(提供遗忘能力)方向相反。α=0.9\alpha = 0.9 的有效记忆仅 66 token,连一个 chunk 都刚覆盖;而 α0.99\alpha \ge 0.99 才落在因式分解仍然安全的区间。常数门控下两者尚可兼顾,但这是假设不是保证–部分实现敢做因式分解,正是隐含依赖了「门控总是接近 1」。

可迁移的启示:代数上等价的两种写法,数值行为可以完全不同。γi/γj\gamma_i / \gamma_jγi(1/γj)\gamma_i \cdot (1/\gamma_j) 在实数域相等,但前者(jij \le i 时)恒不大于 1、后者上界随块长指数增长。论文把衰减写成比值而非乘积并不是记号习惯,而是保证数值安全的必要安排–那套 \overleftarrow{\cdot} / \overrightarrow{\cdot} 箭头记号就是在提醒读者每个衰减都有明确的基准点(chunk 首或 chunk 末)。实现时不要对论文公式做看似等价的代数改写,先问改写后的中间量落在什么范围。


参考

  • Gated Delta Networks(GDN):arXiv:2412.06464,ICLR 2025。§2.1 以 Mamba2 为例给出带衰减的线性递推 St=αtSt1+vtkt\mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^\intercal、累积衰减积 γj=i=1jαi\gamma_j = \prod_{i=1}^{j}\alpha_i、向量 / 矩阵并行两种形式与衰减感知掩码 Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j;式 (1)(2) 给出 chunkwise 形式与 \overleftarrow{\cdot} / \overrightarrow{\cdot} 箭头记号
  • Mamba2:Transformers are SSMs(Dao & Gu, 2024),arXiv:2405.21060。衰减项 αt\alpha_t 与状态空间对偶性(SSD)的出处
  • RetNet:arXiv:2307.08621。常数(数据无关)衰减的代表,即本文实现对应的形态
  • Kimi Linear / KDA 论文:arXiv:2510.26692
  • flash-linear-attentionfla/ops/simple_gla/(标量衰减的参考实现)
  • Mamba2 / GLA:标量与门控衰减的 chunkwise 形式
  • 本站《TileLang 实战:KDA 从零到一–Chunked 线性注意力》《KDA 的来龙去脉》《TileLang 编程基本知识点

本文的四层参考互验、退化检验与 fp16 失效点在 numpy fp64/fp16 上复现;GPU 实测数据待补。

TileLang 实战:KDA 从零到一–Chunked 线性注意力

KDA(Kimi Delta Attention)的递推式一行就写完,但直接照着它写 kernel 需要同时处理四个机制:逐通道门控、delta rule 的三角求解、log 域 cumsum、跨 chunk 状态传递。四个机制耦合在一个 kernel 里,数值出错时无法判断误差来自哪一层。 本文采用递进式实现路径,每一级只引入一个新机制、每一级都可独立运行并做数值验证。本文覆盖第一级–移除全部衰减因子与删除因子,只保留 SS+KVS \leftarrow S + K^\top V

这一级的价值不在性能,而在于将 chunkwise 分解恒等式单独隔离验证。本文同时给出一个反直觉的结论:线性注意力的全部优势是把 O(N2D)O(N^2 D) 换成 O(ND2)O(N D^2),但本文采用的逐块独立重算策略会让 FLOPs 退回 O(N2)O(N^2)–与因果 FlashAttention 同量级。§6 会说明这个退化的根源在于把序列轴放进了 grid,并对照 TileLang 官方 chunk_delta_h 的做法:序列轴不进 grid,递推就退回单 block 内的顺序循环,NN 的次数才能保住

数学背景见《KDA 的来龙去脉》,TileLang 语言基础见《TileLang 编程基本知识点》,本文复用的 tile 切分与 fragment 累加模式见《TileLang 实战:FlashAttention 前向 Kernel》。


1. 递进路径:把 KDA 拆成可验证的增量

KDA 的完整递推式(Dt=diag(egt)D_t = \operatorname{diag}(e^{g_t})gtR<0dkg_t \in \mathbb{R}^{d_k}_{<0}):

St=St1Dt(Iβtktkt)+βtvtkt,ot=StqtS_t = S_{t-1} D_t (I - \beta_t k_t k_t^\top) + \beta_t v_t k_t^\top, \qquad o_t = S_t q_t

记号说明:Kimi Linear 论文(arXiv:2510.26692)原文写作 St=(Iβtktkt)Diag(αt)St1+βtktvtS_t = (I - \beta_t k_t k_t^\top)\operatorname{Diag}(\alpha_t) S_{t-1} + \beta_t k_t v_t^\topαt[0,1]dk\alpha_t \in [0,1]^{d_k}。本文取其转置形式以匹配 kernel 中 SS(dv,dk)(d_v, d_k) 内存布局,并沿用 GDN / Mamba 的 log 域记号 αt=egt\alpha_t = e^{g_t},因为实现里门控本来就存在 log 域。

按机制拆解,每一级只放开一个自由度:

级别 递推式 新增机制 新出现的实现结构
第一级(本文) SS+KVS \leftarrow S + K^\top V 分块恒等式本身、块内因果掩码
第二级 SγS+KVS \leftarrow \gamma S + K^\top V 标量衰减 块内权重从 0/1 变 γij\gamma^{i-j} 指数下三角
第三级 Sdiag(egt)S+S \leftarrow \operatorname{diag}(e^{g_t}) \cdot S + \cdots 逐 token 门控 log 域 cumsum、exp2 硬件指令
第四级 SS(Iβkk)+βvkS \leftarrow S(I - \beta k k^\top) + \beta v k^\top delta rule UT 变换、三角求解(wy_fast 雏形)
第五级 门控 + delta rule 二者耦合 五阶段流水线、跨 kernel 状态传递

最后一级即完整 GDN / KDA,对标 flash-linear-attention 中的 chunkwise 实现。

这样拆分的收益是误差定位能力:任何一级数值对不上,怀疑对象只有这一级新引入的那一个机制,前面几级已经验证过了。本文对应第一级,一个机制都还没引入,因此这里能验证的只有分块恒等式本身。


2. 数学推导:分块恒等式

本级递推不含遗忘项,状态单调累加:

St=St1+ktvt,ot=qtStS_t = S_{t-1} + k_t v_t^\top, \qquad o_t = q_t^\top S_t

展开为显式求和,因果且包含当前 token:

oi=jiqi(kjvj)=ji(qikj)vjo_i = \sum_{j \le i} q_i^\top (k_j v_j^\top) = \sum_{j \le i} (q_i \cdot k_j)\, v_j

这是标准线性注意力。与 softmax attention 的唯一区别是分数未经 softmax,因此求和顺序可自由交换–这是下述分块重写成立的前提。

2.1 将求和拆分为跨块与块内两部分

序列按 BCBC 切分为 NC=N/BCNC = N / BC 个 chunk。拆分的依据不是下标落在哪里,而是因果约束是否需要逐 token 判断

  • 当前块之前的 chunk:其内每一个 jj 对当前块里的每一个 ii 都满足 jij \le i,全部可见。因果约束在这些块上恒真、与 ii 无关,所以整块直接相乘就行。
  • 当前块:内部的 jj 才真正受 jij \le i 约束,只能取 kj,vjk_j, v_jjij \le i 的那一半。

oi=qi(c<cKcVc)跨块:全部可见,无需掩码+jchunk cji(qikj)vj块内:需因果掩码o_i = \underbrace{q_i^\top \Big( \sum_{c' < c} K_{c'}^\top V_{c'} \Big)}_{\text{跨块:全部可见,无需掩码}} + \underbrace{\sum_{\substack{j \in \text{chunk } c \\ j \le i}} (q_i \cdot k_j) v_j}_{\text{块内:需因果掩码}}

关键在第一项:正因为它的因果判断与 ii 无关,括号内的量对整个 chunk cc 才能是同一个矩阵,记作 Scprev=c<cKcVcS_c^{\text{prev}} = \sum_{c' < c} K_{c'}^\top V_{c'},形状 dk×dvd_k \times d_v,含义是处理该 chunk 之前已累积的状态。于是跨块贡献退化为一次矩阵乘 QcScprevQ_c S_c^{\text{prev}},块内贡献是一个 BC×BCBC \times BC 的下三角掩码矩阵乘:

Oc=QcScprev+tril(QcKc)VcO_c = Q_c S_c^{\text{prev}} + \operatorname{tril}(Q_c K_c^\top) V_c

图上半部分是前序 chunk 的注意力方阵——它自身也带因果三角,但这些计算在处理 chunk cc 时已经完成,结果被汇总进一个 dk×dvd_k \times d_v 的状态矩阵 ScprevS_c^{\text{prev}},不必再逐 token 展开。下半部分是本文要算的两项:左边整块打满阴影,表示 QcQ_c 的每一行都能无条件地乘上完整的历史状态,没有掩码;右边只有下三角有阴影,表示块内必须逐 token 判断 jij \le i跨块项的历史长度随序列增长,但它被压进固定形状的 ScprevS_c^{\text{prev}}——这正是线性复杂度的来源;块内三角形的边长恒为 BCBC,与序列长度无关。

这条恒等式是后续各级的公共基础,之后每一级都只是在它的两项上分别插入衰减权重,张量形状与 GEMM 调用次序保持不变。

2.2 掩码可以直接置零的原因

FlashAttention 中掩码必须填 -\infty,因为后续要经过 expe=0e^{-\infty} = 0 才能使被掩位置不产生贡献。线性注意力没有 softmax,掩码位置直接写 0 即可:

1
2
for i, j in T.Parallel(BC, BC):
A[i, j] = T.if_then_else(j <= i, A[i, j], 0.0) # 无 exp,掩码直接清零

这一行是线性注意力与 FA 的第一处分岔,也是「无 softmax」在代码上的全部体现–不需要 online 重标定、不需要 mm\ell 两个行状态、不需要输出阶段的最终除法


3. 四层 PyTorch 参考实现

验证 kernel 之前需要建立可信参考。本文构造三份 PyTorch fp64 实现,从「最直白」逐步过渡到「与 kernel 控制流同构」,相邻两层互相验证:

参考 实现方式 验证目标
A 逐 token 递归,双重 for 循环 递推式定义本身,几乎不可能写错
B 分块向量化,cumsum 求前缀状态 §2.1 分块恒等式的正确性
C 逐块独立重算,模拟 grid 与 T.Pipelined kernel 的控制流与访存次序

参考 C 是连接 PyTorch 与 TileLang 的关键一层–它的循环结构与 kernel 逐句对应,若 kernel 与 C 不符,问题必在 TileLang 语法或内存层级使用上,与数学无关。

3.1 参考 A:逐 token 递归

1
2
3
4
5
6
7
8
9
10
11
def ref_recurrent(Q, K, V):
"""最直白的实现:严格照抄递推式 S += k v^T; o = q S"""
B, N, H, D = Q.shape
O = torch.empty(B, N, H, D, dtype=torch.float64, device=Q.device)
for b in range(B):
for h in range(H):
S = torch.zeros(D, D, dtype=torch.float64, device=Q.device)
for t in range(N):
S += torch.outer(K[b, t, h].double(), V[b, t, h].double())
O[b, t, h] = Q[b, t, h].double() @ S
return O

3.2 参考 B:分块向量化

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
def ref_chunked(Q, K, V, BC):
"""分块向量化:用 cumsum 一次算出所有 chunk 的前缀状态"""
B, N, H, D = Q.shape
NC = N // BC
Qc = Q.permute(0, 2, 1, 3).double().reshape(B, H, NC, BC, D)
Kc = K.permute(0, 2, 1, 3).double().reshape(B, H, NC, BC, D)
Vc = V.permute(0, 2, 1, 3).double().reshape(B, H, NC, BC, D)

outer = torch.einsum("bhcnd,bhcnv->bhcdv", Kc, Vc) # 每块自身的 K^T V
states = torch.cumsum(outer, dim=2) # 含自身的前缀和
states_prev = torch.cat( # 右移一格 → 不含自身
[torch.zeros_like(states[:, :, :1]), states[:, :, :-1]], dim=2)

O_inter = torch.einsum("bhcnd,bhcdv->bhcnv", Qc, states_prev)

A = torch.einsum("bhcnd,bhcmd->bhcnm", Qc, Kc) # 块内分数 [BC, BC]
mask = torch.tril(torch.ones(BC, BC, dtype=torch.bool, device=Q.device))
A = A.masked_fill(~mask, 0.0) # 线性注意力:置 0,非 -inf
O_intra = torch.einsum("bhcnm,bhcmv->bhcnv", A, Vc)

return (O_inter + O_intra).reshape(B, H, N, D).permute(0, 2, 1, 3).contiguous()

states_prev 的右移拼接对应 c<c\sum_{c' < c} 中的严格小于号:cumsum 给出的是含自身的前缀和,右移一格补零才是处理该 chunk 之前的累积状态。这是分块线性注意力最易出错的一行–不右移等价于把当前块的 KVK^\top V 重复计入,块内贡献会被计算两次。该错误的量级实测为相对 L2 误差 7.78×1017.78 \times 10^{-1},属于结构性错误而非精度问题,容易识别。

3.3 参考 C:kernel 控制流镜像

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
def ref_kernel_mimic(Q, K, V, BC):
"""逐块独立重算:外层三重循环模拟 grid,内层 range(bx) 模拟 T.Pipelined(bx)"""
B, N, H, D = Q.shape
NC = N // BC
O = torch.zeros(B, N, H, D, dtype=torch.float64, device=Q.device)
mask = torch.tril(torch.ones(BC, BC, dtype=torch.float64, device=Q.device))

for bz in range(B): # grid.z ← batch
for by in range(H): # grid.y ← head
for bx in range(NC): # grid.x ← Q 块
sl = slice(bx * BC, (bx + 1) * BC)

# ① T.clear(S_f); for c in T.Pipelined(bx): T.gemm(..., transpose_A=True)
S = torch.zeros(D, D, dtype=torch.float64, device=Q.device)
for c in range(bx):
cs = slice(c * BC, (c + 1) * BC)
S += K[bz, cs, by, :].double().T @ V[bz, cs, by, :].double()

# ② T.gemm(Q_s, S_s, acc_o)
Qb = Q[bz, sl, by, :].double()
acc = Qb @ S

# ③ T.gemm(Q_s,K_s,A,transpose_B=True) → 掩码 → T.gemm(A_cast,V_s,acc_o)
Kb = K[bz, sl, by, :].double()
Vb = V[bz, sl, by, :].double()
acc += ((Qb @ Kb.T) * mask) @ Vb

O[bz, sl, by, :] = acc
return O

对照参考 B 与参考 C 可以看出两种前缀状态求法的区别:B 用 cumsum 一次性算出全部 NCNC 个前缀状态、总代价 O(NC)O(NC);C 的每个 bxbx 独立重算、总代价 O(NC2)O(NC^2)kernel 采用的是 C 的策略,原因与代价见 §6。

3.4 四层参考的一致性验证

B=2,H=2,N=12,D=4,BC=4B=2, H=2, N=12, D=4, BC=4,fp64(numpy 复现):

比较 max abs 误差 相对 L2
B 分块向量化 vs A 逐 token 递归 3.55×10153.55 \times 10^{-15} 1.62×10161.62 \times 10^{-16}
C kernel 结构镜像 vs A 逐 token 递归 7.11×10157.11 \times 10^{-15} 1.86×10161.86 \times 10^{-16}
C vs B 4.44×10154.44 \times 10^{-15} 1.78×10161.78 \times 10^{-16}
参考 B 去掉右移(错误实现) 1.97×1011.97 \times 10^{1} 7.78×1017.78 \times 10^{-1}

前三行误差均在 fp64 机器精度量级(ε2.2×1016\varepsilon \approx 2.2 \times 10^{-16}),恒等式与三份实现均无误。第四行是刻意引入的错误,用于确认该验证流程对结构性错误敏感。

参考 D(§6.4.2 的序列内循环形式)另外验证一件事–换 grid 划分、切 DV 后算的还是同一个恒等式B=2,S=256,H=3,DK=64,DV=128,C=64B{=}2, S{=}256, H{=}3, DK{=}64, DV{=}128, C{=}64,fp64:

对象 相对 L2(vs 参考 A)
参考 D,blockDV=32\text{block}_{DV} = 32 5.60×10165.60 \times 10^{-16}
参考 D,blockDV=64\text{block}_{DV} = 64 5.60×10165.60 \times 10^{-16}
参考 D,blockDV=128\text{block}_{DV} = 128(不切) 5.60×10165.60 \times 10^{-16}
参考 D,状态更新提到写回之前 6.98×1016.98 \times 10^{-1}

三个 blockDV\text{block}_{DV} 误差完全相同,这就是「DV 切块零依赖」的直接证据–切不切、切多细,算出来的是比特级相同的东西,因为每个 dv 竖条的浮点累加顺序本就不受其他竖条影响。对照最后一行:仅仅把状态更新从循环尾部提到开头,误差就跳到 70%–右移语义是硬要求。

以上均为 numpy fp64 实测。TileLang kernel 本身需要 CUDA 设备,本文未给出实测数字–待真卡跑通后单独补充,此处不做性能推测。预期的主导误差源是 T.copy(S_f, S_s) 把 f32 状态降至 f16 落 shared(Tensor Core MMA 的输入必须是低精度);状态 SS 是多个块累加的结果,越靠后的 block 累加项越多,误差随 bxbx 单调增长,影响远大于块内 AA 的那次降精度。

3.5 einsum 输出下标的约束

参考 B 中若把 torch.einsum("bhcnd,bhcnv->bhcdv", ...) 误写为 ->bhcdd(意图表达「输出是 d×dd \times d 方阵」),会直接抛出异常:

1
ValueError: einstein sum subscripts string includes output subscript 'd' multiple times

einsum 的输出下标不允许重复–重复下标在输出侧的语义是取对角线。而 KVK^\top V 的两个维度虽然长度均为 DD语义上分别是 key 维与 value 维,必须使用不同字母 dv。同理 bhcnd,bhcdd->bhcnv 也不合法,因为输出的 v 从未在输入中出现。dk=dvd_k = d_v 时长度相同掩盖了语义差异,一旦 KDA 中 dkdvd_k \ne d_v,该疏忽会立即表现为形状错误。


4. TileLang kernel:三步实现

grid 划分与 FA 一致–按 Q 块切分所有权,每个 block 负责输出一个 BC×dBC \times d 的 tile:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
@tilelang.jit(out_idx=[3])
def linattn_chunk(batch, heads, seq_len, dim, blk, num_stages=2,
dtype=T.float16, accum_dtype=T.float32):
BC = blk

@T.prim_func
def main(
Q: T.Tensor([batch, seq_len, heads, dim], dtype),
K: T.Tensor([batch, seq_len, heads, dim], dtype),
V: T.Tensor([batch, seq_len, heads, dim], dtype),
O: T.Tensor([batch, seq_len, heads, dim], dtype),
):
with T.Kernel(T.ceildiv(seq_len, BC), heads, batch, threads=128) as (bx, by, bz):
Q_s = T.alloc_shared([BC, dim], dtype)
K_s = T.alloc_shared([BC, dim], dtype)
V_s = T.alloc_shared([BC, dim], dtype)
S_s = T.alloc_shared([dim, dim], dtype) # 状态降精度落地,喂第二个 gemm
O_s = T.alloc_shared([BC, dim], dtype)

S_f = T.alloc_fragment([dim, dim], accum_dtype) # 跨迭代累加的状态
acc_o = T.alloc_fragment([BC, dim], accum_dtype)
A = T.alloc_fragment([BC, BC], accum_dtype)
A_cast = T.alloc_fragment([BC, BC], dtype)

4.1 步骤①:流式累加前缀状态

对应参考 C 中的 for c in range(bx)

1
2
3
4
5
6
7
T.clear(S_f)
for c in T.Pipelined(bx, num_stages=num_stages): # 循环上界是 runtime 的 bx
T.copy(K[bz, c * BC:(c + 1) * BC, by, :], K_s)
T.copy(V[bz, c * BC:(c + 1) * BC, by, :], V_s)
T.gemm(K_s, V_s, S_f, transpose_A=True) # S_f += K_s^T @ V_s ← TileLang 的 T.gemm 默认累加

T.copy(S_f, S_s) # f32 → f16,本 kernel 主要精度损失点

代码里看不到 +=,但这个循环确实是在累加:T.clear(S_f) 先把 fragment 清零,之后每次 T.gemm 都在 S_f 原地累加,因此循环等价于 Sf=c<bxKcVcS_f = \sum_{c < bx} K_c^\top V_c。累加语义来自 Tensor Core MMA 的基本形式–MMA 指令算的是 d = a @ b + c,累加器既是输入也是输出,T.gemm 默认沿用这一行为;若要覆盖而非累加,需显式传 clear_accum=True

T.Pipelined(bx) 的上界是 block 索引而非编译期常量–不同 block 的循环次数不同,第 0 块一次都不执行(S=0S = 0),最后一块需执行 NC1NC-1 次。TileLang 支持 runtime 上界的流水线,代价是各 block 负载严重不均衡,尾部 block 构成整个 kernel 的关键路径。

这里每个 block 都在重算自己需要的前缀状态,而递推式本身只需一次加法。为什么本文要这么写、以及官方实现如何避开它,见 §6.2 与 §6.4。

4.2 步骤②③:两项贡献求和

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
# ② 跨块:O = Q_c @ S_prev
T.copy(Q[bz, bx * BC:(bx + 1) * BC, by, :], Q_s)
T.clear(acc_o)
T.gemm(Q_s, S_s, acc_o)

# ③ 块内:O += tril(Q_c K_c^T) @ V_c
T.copy(K[bz, bx * BC:(bx + 1) * BC, by, :], K_s) # 重载当前块
T.copy(V[bz, bx * BC:(bx + 1) * BC, by, :], V_s)
T.clear(A)
T.gemm(Q_s, K_s, A, transpose_B=True)
for i, j in T.Parallel(BC, BC):
A[i, j] = T.if_then_else(j <= i, A[i, j], 0.0)
T.copy(A, A_cast) # f32 → f16 喂 MMA
T.gemm(A_cast, V_s, acc_o)

T.copy(acc_o, O_s)
T.copy(O_s, O[bz, bx * BC:(bx + 1) * BC, by, :])

步骤③开头必须重新 T.copy 当前块的 K/V。步骤①的流水线循环结束时,K_s / V_s 中残留的是第 bx1bx-1 块的数据,且在 num_stages > 1 时具体残留哪一块取决于流水线的展开方式,不可假设。复用 shared buffer 节省了显存,但必须显式重载–这是 tile 编程中典型的隐式状态陷阱。参考 C 里没有这个问题,因为 Python 每轮都重新切片,不存在缓冲区复用。

T.copy(A, A_cast) 这次 f32→f16 转换与 FA 中 P~\tilde{P} 喂入第二个 GEMM 前的降精度是同一操作–Tensor Core 的 MMA 输入必须是低精度,仅累加器为 f32。

4.3 kernel 与参考 C 的逐句对应

参考 C(PyTorch) kernel(TileLang) 说明
for bz / by / bx 三重循环 T.Kernel(ceildiv(N,BC), heads, batch) 循环变并行 grid
S = torch.zeros(D, D) T.clear(S_f) 状态初始化,fragment 常驻寄存器
for c in range(bx) T.Pipelined(bx, num_stages) 顺序循环变软件流水线
S += K[cs].T @ V[cs] T.gemm(K_s, V_s, S_f, transpose_A=True) 显式 shared 暂存 + Tensor Core
acc = Qb @ S T.gemm(Q_s, S_s, acc_o) 需先 T.copy(S_f, S_s) 降精度
(Qb @ Kb.T) * mask T.gemm(..., transpose_B=True) + T.Parallel 掩码 掩码从广播乘变逐元素条件
acc += (...) @ Vb T.copy(A, A_cast) + T.gemm(A_cast, V_s, acc_o) 多一次 f32→f16 转换
O[bz, sl, by, :] = acc T.copy(acc_o, O_s) + T.copy(O_s, O[...]) 经 shared 中转写回 HBM

两处 PyTorch 中不存在的操作:f32→f16 显式降精度(Tensor Core 输入约束)与shared buffer 中转(内存层级手动管理)。这两项也正是 kernel 与参考 C 数值差异的全部来源。


5. 与 FlashAttention 的三点结构差异

本文实现复用了 FA 的全部 tile 模式,但有三处必须修改:

维度 FlashAttention 线性注意力 原因
归一化 online softmax,维护 mm\ell 两个行状态,每块重标定 无,掩码直接置 0 exp,求和顺序可交换
循环携带的量 无跨块状态,每个 Q 块独立 S[dk,dv]S[d_k, d_v] 是跨迭代累加的状态 递推式本身带状态
T.gemm policy 必须 FullRow–行归约要求整行在同一 warp 默认 Square 即可 无按行归约

第三点值得展开:FA 需对分数块做 rowmax / rowsum,若一行被切分到多个 warp,归约就需跨 warp 通信,因此必须用 FullRow policy 强制整行不拆。线性注意力的块内矩阵 AA 计算完成后仅做逐元素掩码,无任何跨列归约,warp 划分方式不影响正确性,编译器可自由选择寄存器分布最优方案。

第二点需要说明的是,本文实现并没有真正兑现这一行:逐块独立重算把状态依赖藏起来了,每个 block 都从零重算自己需要的前缀状态,块之间不传递任何东西。这不是偷懒,而是本文这种「序列轴占据 grid.x」的划分下,block 之间无法顺序传递状态。换一种 grid 划分就没有这个限制,见 §6.4。


6. 代价账本:线性复杂度从何而来,如何被丢掉,以及如何拿回

6.1 两条路线的规模差异

先把线性注意力的本质说清楚。同样的输出,有两条算法:

路线 计算方式 规模
softmax / FA 先算分数矩阵 QKQK^\topN×NN \times N),再乘 VV O(N2D)O(N^2 D)
线性注意力 先算状态 S=KVS = K^\top VD×DD \times D),再乘 QQ O(ND2)O(N D^2)

区别在结合律往哪边括:(QK)V(QK^\top)V 要物化一个 N×NN \times N 的中间矩阵,Q(KV)Q(K^\top V) 物化的是 D×DD \times DNN 的次数从 2 降到 1,DD 的次数从 1 升到 2–这才是线性注意力唯一的、也是全部的优势来源,代价是 O(D2)O(D^2) 的固定状态取代了 O(N2)O(N^2) 的自由分数矩阵,表达能力随之受限。

chunkwise 形式是这两条路线的混合:跨块走 SS 路线(每块一次 QcScprevQ_c S_c^{\text{prev}},共 N2D2N \cdot 2D^2),块内走 QKQK^\top 路线(每块一个 BC×BCBC \times BC 分数矩阵,共 N4BCDN \cdot 4 BC D),状态更新本身再花 N2D2N \cdot 2D^2。理想总量:

FLOPsideal=4ND2+4NBCD=4ND(D+BC)\text{FLOPs}_{\text{ideal}} = 4N D^2 + 4N\,BC\,D = 4ND(D + BC)

NN 是一次方。与因果 FA 的 2N2D2N^2 D 相比:

FLOPsidealFLOPscausal FA=2(D+BC)N\frac{\text{FLOPs}_{\text{ideal}}}{\text{FLOPs}_{\text{causal FA}}} = \frac{2(D + BC)}{N}

比值按 1/N1/N 衰减–序列越长优势越大,这是线性注意力值得做的全部理由。

6.2 本文为何没有顺着递推式只加一次

理想账本里的 NCNCKVK^\top V,对应的是顺序递推:

Sc=Sc1+Kc1Vc1S_{c} = S_{c-1} + K_{c-1}^\top V_{c-1}

每个 chunk 只做一次 GEMM,读上一步的结果、加上自己这一块。参考 B 就是这么算的(cumsum 一次扫完),CPU 上顺序执行毫无问题。

但 kernel 做不到这一点,原因在 grid 的划分方式。 §4 把所有权按 Q 块切分,bxbxgrid.x 的索引:所有 NCNC 个 block 由硬件并行调度,执行顺序不确定、彼此之间没有同步点、也没有共享的可写缓冲区。block bxbx 若想读 SbxS_{bx},就得等 block bx1bx-1 算完并把结果落到某处——这两件事在单个 kernel 内都不成立:CUDA 不保证 block 间的执行次序(bx1bx-1 可能还没启动),也没有跨 block 的 barrier 可用。

顺序递推要求的是"前一步已完成",而 grid 提供的是"所有步同时开始"。但这个矛盾是本文自己造出来的–它成立的前提是"序列轴必须占据 grid 的一维"。放弃这个前提,矛盾就不存在了,见 §6.4。

本文仍保留逐块独立重算:它让 kernel 保持单文件闭环、控制流与参考 C 严格对应,适合把分块恒等式本身隔离出来验证。下面先算清这个选择的代价。

6.3 冗余重算把 NN 的次数还了回去

bxbx 块执行 bxbxKVK^\top V,全部 block 合计 c=0NC1c=NC(NC1)/2\sum_{c=0}^{NC-1} c = NC(NC-1)/2 次,而顺序递推共 NCNC 次。冗余系数 (NC1)/2(NC-1)/2,且NN 线性增长–正是这个增长把 NN 的次数从 1 顶回 2:

FLOPs=NC(NC1)22BCD2N2D2BC\text{FLOPs}_① = \frac{NC(NC-1)}{2} \cdot 2\, BC\, D^2 \approx \frac{N^2 D^2}{BC}

于是与因果 FA 的比值退化成常数:

FLOPsFLOPscausal FA=D2BC\frac{\text{FLOPs}_①}{\text{FLOPs}_{\text{causal FA}}} = \frac{D}{2\,BC}

这个比值与 NN 无关,恰恰是失败的判据:它说明本文实现与 FA 同属 O(N2)O(N^2)1/N1/N 的衰减优势被完全抹掉了。实测账本(D=64D = 64BC=64BC = 64):

NN NCNC 冗余系数 本文实现(步骤①) 理想线性注意力 因果 FA 本文/FA 理想/FA
512 8 3.5x 0.015 GFLOP 0.017 GFLOP 0.034 GFLOP 0.438 0.500
2048 32 15.5x 0.260 GFLOP 0.067 GFLOP 0.537 GFLOP 0.484 0.125
8192 128 63.5x 4.261 GFLOP 0.268 GFLOP 8.590 GFLOP 0.496 0.031
16384 256 127.5x 17.113 GFLOP 0.537 GFLOP 34.360 GFLOP 0.498 0.016
65536 1024 511.5x 274.609 GFLOP 2.147 GFLOP 549.756 GFLOP 0.500 0.004

看最后两列的走向:本文/FA 收敛到常数 D/(2BC)=0.5D/(2BC) = 0.5,理想/FA 按 1/N1/N 一路衰减到 0.004。 N=65536N = 65536 时理想实现只需 2.1 GFLOP,本文实现要 274.6 GFLOP–差 128 倍,而这个倍数还会随 NN 继续涨。

BCBC 的取舍随之明确:

BCBC D/(2BC)D/(2BC) 含义
32 1.000 与因果 FA 计算量持平,块过小无收益
64 0.500 默认值,寄存器压力可控
128 0.250 冗余减半,但 AA128×128128 \times 128 f32 fragment,每线程约 128 个寄存器,易溢出
256 0.125 理论最省,实际 shared memory 与寄存器均无法容纳

但要注意这张表只是在常数上打折,O(N2)O(N^2) 的量级不变。增大 BCBC 能线性减少冗余,但 AABC×BCBC \times BC 的 f32 fragment,BC=128BC = 128 时寄存器压力已接近溢出边界。真正的解法不是调 BCBC,见下一节。

6.4 更好的方案:把序列轴从 grid 里拿掉

TileLang 官方 examples/gdn 中的 chunk_delta_h 给出了另一种划分。它的 grid 只有两维:

1
2
3
4
5
6
7
8
with T.Kernel(T.ceildiv(DV, block_DV), B * H, threads=threads) as (bv, bbh):
...
for i_s in T.Pipelined(T.ceildiv(S, block_S), num_stages=num_stages):
# 存上一轮的状态快照,供下游 kernel 使用
T.copy(b_h_shared, h[bb, i_s, bh, 0:DK, bv * block_DV:(bv + 1) * block_DV])
...
T.gemm(K_shared, V_new_shared, b_h_fragment, transpose_A=True) # 状态原地累加
T.copy(b_h_fragment, b_h_shared)

序列不在 grid 里,而是 kernel 内部的一个顺序循环。 b_h_fragment 成为 loop-carried 的寄存器变量,跨 chunk 一路累加–这正是 Sc=γcSc1+KcVcS_c = \gamma_c S_{c-1} + K_c^\top V_c 的直接翻译,每个 chunk 只做一次 KVK^\top V,一次都不重算。§6.2 里那个"block 间无法同步"的矛盾根本不会出现,因为递推的顺序性被限制在单个 block 内部,用寄存器解决了,从来不需要跨 block 通信

并行度从另外两个轴补回来:

来源 数量(B=1,H=32,DV=128,blockDV=32B=1,H=32,DV=128,block_{DV}=32
bv DV 切块 128/32=4128/32 = 4
bbh batch ×\times head 融合 1×32=321 \times 32 = 32

合计 128 个 block,足够填满 SM。

6.4.1 为什么只切 DV、不切 DK

这一点容易给出错误的理由。若只看状态更新那一个 GEMM:

Sc+=KcVc,Kc:[DK,C], Vc:[C,DV]S_c \mathrel{+}= K_c^\top V_c,\qquad K_c^\top: [DK, C],\ V_c: [C, DV]

它的收缩维是 chunk 长度 CCDKDK 是 M 维、DVDV 是 N 维。而切 M 维只是输出 tiling,各 block 写各自的行,同样零依赖–所以「DKDK 是 M 维、切开要跨 block 归约」是站不住的,单看这一步切哪边都行。

不对称性来自下游谁消费这个状态。把一个 chunk 里三个 GEMM 的收缩维列出来:

GEMM 形状 收缩维 DKDK 的角色 DVDV 的角色
状态更新 KcVcK_c^\top V_c [DK,C]×[C,DV][DK,C] \times [C,DV] CC M(自由) N(自由)
跨块读出 QcScprevQ_c S_c^{\text{prev}} [C,DK]×[DK,DV][C,DK] \times [DK,DV] DKDK 收缩 N(自由)
块内项 QcKcQ_c K_c^\top [C,DK]×[DK,C][C,DK] \times [DK,C] DKDK 收缩 不出现

结论一句话:DVDV 在三个 GEMM 里始终是自由维,DKDK 在其中两个里是收缩维。

把两种切法的张量尺寸画出来,区别就是一个词:拼接还是相加

  • 切 DV:block (bv,bbh)(bv, bbh) 独占 S[:,dv]S[:, \text{dv}] 这一竖条,自己跑完整条递推,输出 O[:,dv]=QcSprev[:,dv]+tril(QcKc)Vc[:,dv]O[:, \text{dv}] = Q_c S^{\text{prev}}[:, \text{dv}] + \operatorname{tril}(Q_cK_c^\top)V_c[:, \text{dv}]C×DV/nC \times DV/n 的条带,已是最终值,直接写回。状态和输出两步都只有拼接,零累加、零跨块通信。
  • 切 DK:状态更新那步确实也没事(Si=KiVS^i = K_i^\top VDK/n×DVDK/n \times DV 的行条带,也不重叠),但读出就崩了QQ 被迫跟着切成 C×DK/nC \times DK/n,算出的 o~i=QiSi\tilde{o}^i = Q^i S^iC×DVC \times DV 全宽nno~i\tilde{o}^i 盖在同一块 OO 上,各含 1/n 根收缩轴的贡献,Oc=io~iO_c = \sum_i \tilde{o}^i,少加任何一份结果就错。这就是 split-K:每个 chunk 每条序列都要归约一次 [C,DV][C, DV] 的中间结果,摊销不掉,只能选 fp32 atomic add(非确定性 + 带宽)或另开 workspace 起第二个 kernel。

更麻烦的是块内项:QcKcQ_cK_c^\top 也沿 DKDK 收缩,只持有 dk\text{dk} 子块的 block 根本算不出完整的 [C,C][C,C] 转移矩阵(到了 DeltaNet 那级还要求 (I+tril(diag(β)KK))1(I + \operatorname{tril}(\operatorname{diag}(\beta)K K^\top))^{-1},逻辑上无法切)。

图里“QQ 也要切”那一行值得单独指出:DK 是吸收维,一旦切它,所有沿 DK 寻址的张量(QQKKSS 的行)都被动跟随,而产出反而变成全宽部分和–输入维度变小、输出维度不变,这个尺寸上的不匹配就是归约的必然信号。切 DV 恰好相反:输入切窄一根,输出也跟着窄一根。

还有一个纯工程的理由:状态 fragment 是 [DK,blockDV][DK, \text{block}_{DV}] 的 f32,DK=DV=128DK{=}DV{=}128、128 线程时,不切 DV 要 128×128/128=128128\times128/128 = 128 个寄存器/线程,直接溢出;切成 blockDV=32\text{block}_{DV}{=}32 降到 32 个。DV 切块同时解决了 SM 占用率和寄存器压力两件事,且不引入任何归约。

6.4.2 完整例子:按 (b,h,dv)(b, h, dv) 切 tile,沿 SS 走 Pipelined

把上面的结论落成第一级(无衰减、无删除)的可运行形态。与 §4 相比只改一件事:grid 里的序列轴换成 DV 轴,序列退回 kernel 内部的顺序循环

所有权划分:block (bv,bbh)(bv, bbh) 负责 O[b,:,h, dv 竖条]O[b, :, h, \ \text{dv 竖条}]整条序列的一个 value 通道子集。

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
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
@tilelang.jit(out_idx=[3])
def linattn_seq_in_loop(B, H, S, DK, DV, block_S=64, block_DV=32,
num_stages=2, threads=128,
dtype=T.float16, accum_dtype=T.float32):
C = block_S
NS = T.ceildiv(S, C)

@T.prim_func
def main(
Q: T.Tensor([B, S, H, DK], dtype),
K: T.Tensor([B, S, H, DK], dtype),
V: T.Tensor([B, S, H, DV], dtype),
O: T.Tensor([B, S, H, DV], dtype),
):
# grid 两维:DV 切块 x (batch·head) 融合;序列轴不在这里
with T.Kernel(T.ceildiv(DV, block_DV), B * H, threads=threads) as (bv, bbh):
bb = bbh // H # batch 索引
bh = bbh % H # head 索引
dv0 = bv * block_DV # 本 block 负责的 value 通道起点

Q_s = T.alloc_shared([C, DK], dtype)
K_s = T.alloc_shared([C, DK], dtype)
V_s = T.alloc_shared([C, block_DV], dtype) # 只加载自己那一竖条
S_s = T.alloc_shared([DK, block_DV], dtype) # 状态的 f16 副本,喂 MMA
O_s = T.alloc_shared([C, block_DV], dtype)

S_f = T.alloc_fragment([DK, block_DV], accum_dtype) # loop-carried 状态
acc_o = T.alloc_fragment([C, block_DV], accum_dtype)
A = T.alloc_fragment([C, C], accum_dtype)
A_cast = T.alloc_fragment([C, C], dtype)

T.clear(S_f) # 循环外只清一次:整条序列共用同一个状态

for i_s in T.Pipelined(NS, num_stages=num_stages):
s0 = i_s * C
T.copy(Q[bb, s0:s0 + C, bh, :], Q_s) # 全 DK
T.copy(K[bb, s0:s0 + C, bh, :], K_s) # 全 DK
T.copy(V[bb, s0:s0 + C, bh, dv0:dv0 + block_DV], V_s) # 仅 dv 竖条

# ② 跨块项:用「进入本 chunk 前」的状态,必须先读后更新
T.copy(S_f, S_s) # f32 -> f16
T.gemm(Q_s, S_s, acc_o, clear_accum=True) # [C,DK]x[DK,bDV]

# ③ 块内项:tril(Q K^T) V,沿 DK 收缩,本 block 持有完整 DK
T.gemm(Q_s, K_s, A, transpose_B=True, clear_accum=True)
for i, j in T.Parallel(C, C):
A[i, j] = T.if_then_else(j <= i, A[i, j], 0.0)
T.copy(A, A_cast)
T.gemm(A_cast, V_s, acc_o) # 累加到 ②

T.copy(acc_o, O_s)
T.copy(O_s, O[bb, s0:s0 + C, bh, dv0:dv0 + block_DV])

# ① 状态更新:放在 ②③ 之后 == 递推式里的「右移一格」
T.gemm(K_s, V_s, S_f, transpose_A=True) # [DK,C]x[C,bDV]

return main

四处必须讲清的细节:

  1. 循环体内的顺序就是右移语义。 T.gemm(K_s, V_s, S_f, ...) 必须排在输出写回之后:SfS_f 在读出时代表 Scprev=c<cKcVcS^{\text{prev}}_c = \sum_{c' < c} K_{c'}^\top V_{c'},本 chunk 自己的贡献由 tril\operatorname{tril} 那一项负责。把状态更新提到前面,就变成了 inclusive 前缀,块内项会被重复计入–这是本级唯一的结构性错误点,且数值上表现为「整体偏大」而非 NaN,很容易漏掉。§3.4 里刻意去掉右移的那个错误实现,对应的就是这里的顺序写反。
  2. T.clear(S_f) 在循环外,clear_accum=True 在循环内。 状态要跨迭代累加,所以只能循环外清零一次;而 acc_oA 每个 chunk 都是全新的,进了流水线循环就不能再用 T.clear(清零会被排到流水线的错误阶段),必须靠 clear_accum=True 在 MMA 那一刻覆盖累加器。
  3. num_stages 只能盖住访存,盖不住计算。 S_f -> S_s -> T.gemm -> S_f 构成一条真正的循环依赖,编译器无法把相邻两个 chunk 的计算重叠;流水线的收益全部来自把下一个 chunk 的 Q/K/V 的 HBM→shared 搬运提前发出。所以 num_stages=2 基本够用,继续加只是多占 shared。
  4. Q/K 被 dv 方向重复读。 每个 bvbv 都要读整份 QQKK(各 DKDK 全宽),读放大系数 DV/blockDVDV/\text{block}_{DV}VVOO 则严格切分不重复。这是切 DV 唯一的代价–多读,但不归约。相比之下切 DK 是要归约,这就是取舍的本质区别。

对应的 PyTorch 参考(延续 §3 的写法,可直接跟参考 A/B 对齐):

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
def ref_D_seq_in_loop(Q, K, V, C, block_DV):
B, S, H, DK = Q.shape
DV = V.shape[-1]
O = torch.zeros(B, S, H, DV, dtype=torch.float64)
mask = torch.tril(torch.ones(C, C, dtype=torch.float64))
for b in range(B):
for h in range(H):
for dv0 in range(0, DV, block_DV): # <- grid.x
dv = slice(dv0, dv0 + block_DV)
Sm = torch.zeros(DK, block_DV, dtype=torch.float64)
for i_s in range(S // C): # <- T.Pipelined
sl = slice(i_s * C, (i_s + 1) * C)
Qb, Kb, Vb = Q[b, sl, h, :], K[b, sl, h, :], V[b, sl, h, dv]
O[b, sl, h, dv] = Qb @ Sm + (Qb @ Kb.T * mask) @ Vb
Sm = Sm + Kb.T @ Vb # 右移:更新在写回之后
return O

这份参考与参考 B 的相对 L2 误差在 fp64 下应落在 101610^{-16} 量级:它算的是同一个恒等式,只是把「每块重算前缀」换成了「顺序携带前缀」,数学上完全等价,冗余系数从 (NC1)/2(NC-1)/2 降到 1。

FLOPs 回到 SS 的一次方–每 chunk 每 dv-block 两次 GEMM,求和得 4SDKDV4\,S\,DK\,DV。与本文实现对比(D=64D = 64BC=64BC = 64):

NN 本文步骤① 序列出 grid 倍数
8192 4.261 GFLOP 0.067 GFLOP 63.5x
16384 17.113 GFLOP 0.134 GFLOP 127.5x
32768 68.585 GFLOP 0.268 GFLOP 255.5x
65536 274.609 GFLOP 0.537 GFLOP 511.5x

倍数恰好是冗余系数 (NC1)/2(NC-1)/2,随 NN 线性增长。

6.4.3 融合还是拆成两个 kernel

值得注意的是,§6.4.2 那个 kernel 没有把状态快照写回 HBM–因为它在同一个循环里就把 OO 算完了:block 持有完整的 DKDKQcSprevQ_c S^{\text{prev}}tril(QcKc)Vc\operatorname{tril}(Q_cK_c^\top)V_c 都能就地完成,状态从头到尾只活在寄存器里。

官方 chunk_delta_h 却在循环开头写了 T.copy(b_h_shared, h[...]),把每个 chunk 的状态快照全部落盘。h 的形状是 (B,S/blockS,H,DK,DV)(B, S/block_S, H, DK, DV),按官方 main() 的配置(B=1,S=32768,H=32,DK=DV=128B{=}1, S{=}32768, H{=}32, DK{=}DV{=}128,chunk 64,bf16)达 512 MiB–与 KKVV 输入之和等量。

方案 KVK^\top V 次数 HBM 额外开销 适用
§4 序列进 grid NC(NC1)/2NC(NC-1)/2 隔离验证恒等式
§6.4.2 融合单 kernel NCNC 无门控/无删除的前向、推理
官方两 kernel(chunk_delta_h + chunk_o NCNC hh 快照(示例 512 MiB) 训练(反向要 hh)、DeltaNet 的 UT 变换

拆开的三个真实理由:反向传播需要每个 chunk 的 SprevS^{\text{prev}},重算不如存;DeltaNet 的 WW(I+tril(diag(β)KK))1(I + \operatorname{tril}(\operatorname{diag}(\beta)KK^\top))^{-1} 需要沿 DKDK 收缩的独立阶段,塞不进这个循环;chunk_o 有自己的 tile 划分自由度–读一份现成的快照就能对所有 chunk 完全并行,不再受递推顺序约束。顺序性和并行性被分到两个 kernel 里,各自取所需;融合省 HBM,拆开换灵活性和反向所需的中间量。 第一级用不到后两者,所以融合是更好的起点。

还有两处细节值得对照本文的实现:

  • clear_accum=True 的用途T.gemm(W_shared, b_h_shared, V_new_fragment, clear_accum=True) 显式覆盖而非累加,因为 V_new_fragment 每轮都要重算,不能沿用上一 chunk 的残留。本文 §4.1 靠循环外的 T.clear(S_f) 达到同样目的,两种写法都行;进入流水线循环后就必须用 clear_accum
  • 门控在 log 域相减,从不物化比值。官方实现写作 T.exp2((G_last_local - G_fragment[i_s2, i_v]) * 1.442695),其中 GG 已是 logsigmoid 后的累积和,1.442695=log2e1.442695 = \log_2 eexe^x 转成硬件 exp2GG 单调递减保证 GlastGi0G_{last} - G_i \le 0指数结果恒不大于 1,不存在溢出。本文的第一级没有门控,这里仅作为对照记录。

一句话总结:线性注意力的全部优势是把 O(N2D)O(N^2 D) 换成 O(ND2)O(N D^2),而这个优势能否落地,取决于序列轴放不放进 grid。放进去,block 间无法传递状态,只能各自重算,(NC1)/2(NC-1)/2 的冗余把 NN 的次数顶回 2;不放进去,递推退回单 block 内的顺序循环、用寄存器承载状态,并行度改由 batch/head/DV 提供,代价是把状态快照写回 HBM。本文选前者换取实现的可隔离性,生产实现选后者。


7. 总结

  1. 本级移除了 KDA 的全部衰减因子与删除因子,只保留 SS+KVS \leftarrow S + K^\top V,目的是将 chunkwise 分解恒等式 Oc=QcScprev+tril(QcKc)VcO_c = Q_c S_c^{\text{prev}} + \operatorname{tril}(Q_c K_c^\top) V_c 单独隔离验证。该恒等式是后续各级的公共基础,每级只在它的两项上插入衰减权重,GEMM 的形状与调用次序不变。
  2. 四层 PyTorch 参考构成从数学到 kernel 的完整链条:A 逐 token 递归验证递推式定义,B 分块向量化验证恒等式,C 逐块独立重算镜像 §4 kernel 控制流,D(§6.4.2)按 (b,h,dv)(b,h,dv) 切 tile + 序列内循环镜像生产型 kernel。kernel 与参考的差异仅剩两项–f32→f16 显式降精度与 shared buffer 中转,这也是全部数值差异的来源。
  3. 与 FA 的差异集中在三点:无 online softmax(掩码置 0 而非 -\infty)、循环携带 dk×dvd_k \times d_v 状态、无按行归约(T.gemm policy 用默认 Square 即可)。
  4. 线性注意力的优势与本文实现的退化:优势来自结合律换边,(QK)V(QK^\top)VO(N2D)O(N^2 D) 变成 Q(KV)Q(K^\top V)O(ND2)O(N D^2),理想实现与因果 FA 的比值按 2(D+BC)/N2(D+BC)/N 衰减。但本文的逐块独立重算使冗余系数 (NC1)/2(NC-1)/2NN 线性增长,把 NN 的次数顶回 2–实测比值收敛到常数 D/(2BC)=0.5D/(2BC) = 0.5比值与 NN 无关正是失败的判据,不是中性描述。
  5. 退化的根源是 grid 划分,不是算法:把序列轴放进 grid.x 后,block 间既无执行次序保证也无同步原语,前缀状态只能各自重算。换成 grid =(DV/blockDV, BH)= (DV/block_{DV},\ B \cdot H)序列轴不进 grid,递推退回单 block 内的 T.Pipelined 顺序循环,状态作为 loop-carried fragment 常驻寄存器,每 chunk 只做一次 KVK^\top V。并行度改由 batch/head/DV 提供(示例配置 128 个 block)。N=65536N = 65536 时两者相差 511.5 倍。§6.4.2 给出了完整可运行的融合版本。
  6. 只切 DV、不切 DK 的真正理由不是「KVK^\top V 的 M 维」KVK^\top V 的收缩维是 chunk 长度 CC,DK 和 DV 在这一步都是自由维,切哪边都不需归约。不对称性来自下游:DKDKQcSprevQ_c S^{\text{prev}}QcKcQ_c K_c^\top 两个 GEMM 的收缩维,DVDV 在三个 GEMM 里始终是自由维。切 DK 会把读出变成 split-K(每 chunk 归约一次 [C,DV][C, DV]),并且块内项根本算不出完整的 [C,C][C,C] 转移矩阵;切 DV 只付出Q/K 的读放大,多读而不归约。实测三个 blockDV\text{block}_{DV}(32/64/128)相对 L2 完全相同(5.60×10165.60 \times 10^{-16},见 §3.4)。
  7. 循环体内的顺序就是右移语义:状态更新 T.gemm(K_s, V_s, S_f, transpose_A=True) 必须排在输出写回之后,提前就变成 inclusive 前缀、块内项被重复计入,实测相对 L2 从 101610^{-16} 跳到 6.98×1016.98 \times 10^{-1}(见 §3.4)。还有一条配套规则:T.clear 只能在流水线循环外给跨迭代的状态用,循环内那些每轮重算的累加器必须靠 clear_accum=True
  8. 融合还是拆两个 kernel:本级无门控无删除,单 block 持有完整 DKDK,状态可以全程待在寄存器里,不需要写 hh 快照。官方 chunk_delta_h + chunk_o 拆开的理由是反向传播需要 SprevS^{\text{prev}}、DeltaNet 的 UT 变换需要沿 DKDK 收缩的独立阶段、以及 chunk_o 想要自己的 tile 划分自由度,代价是示例配置下 512 MiB 的 HBM 写回。
  9. 实测数字(均 numpy fp64,见 §3.4):四层参考互验的相对 L2 误差均在 1.65.6×10161.6\text{--}5.6 \times 10^{-16},恒等式无误;刻意去掉右移的错误实现相对 L2 为 7.78×1017.78 \times 10^{-1}(参考 B)与 6.98×1016.98 \times 10^{-1}(参考 D),验证流程对结构性错误敏感。TileLang kernel 本身需 CUDA 设备,实测待补。
  10. 两个实现陷阱:einsum 输出下标不可重复(->bhcdd 会抛异常,key 维与 value 维必须用不同字母);T.Pipelined 之后 shared buffer 的残留内容取决于流水线展开方式,复用前必须显式重载。

参考

本文的分块恒等式验证与 FLOPs 账本在 numpy fp64 上复现;GPU 实测数据待补。

TileLang 实战:FlashAttention 前向 Kernel

FlashAttention 没有发明新的数学。它是 online softmax 递推与两个 GEMM 在同一个 tile 循环里的交织–分数块 SS 算出来就地做 softmax 得到 P~\tilde{P}P~\tilde{P} 就地乘上 VV,中间矩阵全程驻留在 SRAM,HBM 流量从 O(N2)O(N^2) 降到 O(Nd)O(N \cdot d)。本文用 TileLang(v0.1.13)把 FA-2 论文的 Algorithm 1 写成可运行的 kernel:先补齐数学地基,再讲清楚"怎么切",然后逐行拆解主循环,最后做数值验证、测速与调参。只覆盖前向;TileLang 五原语与 T.Pipelined 流水线机制见上一篇《TileLang 编程基本知识点》,本文直接使用其结论。


一、为什么标准 Attention 慢:IO 才是瓶颈

标准 scaled dot-product attention 的流程是 S=QK/dP=softmax(S)O=PVS = QK^\top / \sqrt{d} \to P = \mathrm{softmax}(S) \to O = PV。计算本身没有问题,问题在中间矩阵:SSPP 都是 N×NN \times N,必须写进显存再读回来。N=8192N = 8192 时单个 fp32 矩阵就是 256 MB,长序列下这个开销直接失控。

GPU 内存层级(A100 量级):

层级 容量 带宽 延迟量级
HBM(显存) 40~80 GB ~1.5-2 TB/s ~400-800 cycle
SRAM(每 SM) 164 KB ~19 TB/s ~20-30 cycle
寄存器(每 SM) 256 KB ~19 TB/s 接近 0

SRAM 比 HBM 快 10 倍以上,但只有一百多 KB。attention 的 FLOPs 是 O(N2d)O(N^2 d),标准实现与 FlashAttention 完全相同–快慢差别全部来自中间矩阵走没走 HBM

指标 标准 Attention FlashAttention
中间矩阵内存 O(N2)O(N^2),S/P 落显存 O(N)O(N),S/P̃ 留在 SRAM
HBM 往返 3 次(写 S → 读写 P → 读 P 做 PV) 1 次(读 Q/K/V → 写 O)
FLOPs O(N2d)O(N^2 d) 相同
精确性 精确 精确(非近似)

这就是论文标题里 “IO-aware” 的含义:不省计算,省搬运。

一句话总结:attention 是访存受限的算子,FlashAttention 的全部收益来自让 N×NN \times N 的中间矩阵不落 HBM。


二、数学地基:softmax 的三级台阶

2.1 平移不变性 → 安全 softmax

softmax 对输入加任意常数 cc 不变(分子分母的 ece^c 约掉)。这个自由度是数值安全的救命稻草:fp16 上限只有 65504,s=100s = 100e1001043e^{100} \approx 10^{43} 直接 inf,inf/inf 变 NaN。减去行最大值 mm 后指数上限恰为 0:

softmax(xi)=eximjexjm,m=maxjxj\mathrm{softmax}(x_i) = \frac{e^{x_i - m}}{\sum_j e^{x_j - m}}, \quad m = \max_j x_j

代价是计算顺序被强制:先扫一遍求 mm,再扫一遍求分母与输出–两遍扫描。

2.2 分块场景下,"两遍"成为灾难

SRAM 装不下整行 K/V,只能按块流入。设处理完第 1 块时按基准 m=3m=3 攒好部分和;第 2 块冒出更大分数,基准换成 m=7m=7–旧 exp 全部基准错误。三条路:

方案 做法 代价
存下全部 S 等全局 max 再统一归一化 N×NN \times N 落显存,爆
两遍重扫 pass1 求 mm\ell;pass2 重算加权 K/V 从 HBM 读两次,带宽 ×2
online softmax 边收块边维护"以当前 max 为基准"的部分和 一遍过,数学无损

2.3 online softmax 递推

每个 query 行维护三个状态,初值 m=m = -\infty=0\ell = 0O=0O = 0。新块 SnewS_{new} 到来时:

mnew=max(mold, rowmax(Snew))α=emoldmnew旧成果的打折系数new=αold+rowsum(eSnewmnew)Onew=αOold+eSnewmnewVnew\begin{aligned} m_{new} &= \max(m_{old},\ \mathrm{rowmax}(S_{new})) \\ \alpha &= e^{m_{old} - m_{new}} && \leftarrow \text{旧成果的打折系数} \\ \ell_{new} &= \alpha \cdot \ell_{old} + \mathrm{rowsum}(e^{S_{new} - m_{new}}) \\ O_{new} &= \alpha \cdot O_{old} + e^{S_{new} - m_{new}} \cdot V_{new} \end{aligned}

循环结束后做全程唯一一次归一化:Output=O/\mathrm{Output} = O / \ell

为什么无损:缩放因子连乘时指数项逐项相消,em1m2em2m3=em1m3e^{m_1 - m_2} \cdot e^{m_2 - m_3} = e^{m_1 - m_3},循环结束时每项都被精确换算到最终 max 的基准下。中间任何时刻 OO 都是合法但未归一化的加权和,数值永远有界(当前最大项的 exp 恰为 1)。

数值走一遍(一行 query,三块各来一个分数 1、2、3,正确答案 = softmax(1, 2, 3) 的权重):

1
2
3
4
5
6
7
8
init : m=−∞   ℓ=0      O=0
j=1 : m=1, α=e^(−∞)=0
O = 0×0 + e^0·v₁ = 1.000·v₁ ℓ = 1.000
j=2 : m=2, α=e^(1−2)=0.368
O = 0.368·v₁ + 1·v₂ ℓ = 1×0.368 + 1 = 1.368
j=3 : m=3, α=e^(2−3)=0.368
O = 0.135·v₁ + 0.368·v₂ + 1·v₃ ℓ = 1.503
收尾 : O/ℓ = (0.090, 0.245, 0.665)·v ✓

盯住 v1v_1 的系数:1.000 → 0.368 → 0.135,每来一个更大的 max,历史成果整体打折一次。这就是代码里 acc_o *= scores_scale 那一行的全部含义。

2.4 FlashAttention = online softmax × 两个 GEMM

online softmax(2018,Milakov & Gimelshein)只解决流式归一化,没碰矩阵乘。FlashAttention(2022)的洞察是:attention 恰好是 matmul → softmax → matmul 的三明治,三个部件可以共用同一个 tile 循环:

1
2
3
4
5
for j in 1..Tc:                                  # K/V 方向遍历
S_j = Q_tile @ K_j.T # gemm#1:喂料
m, alpha, P̃_j, ℓ = online_softmax_update(S_j) # 三件套
O_tile = alpha * O_tile + P̃_j @ V_j # gemm#2:消费
Output = O_tile / ℓ # 唯一一次归一化

SSP~\tilde{P} 的生成和消费全部在 SRAM/寄存器内完成,不写入 HBM。

一句话总结:先有 online softmax 这条递推式,才有 FlashAttention 这个 kernel;数学在前,工程在后。


三、切法:block_M 进 grid,block_N 进循环

3.1 都沿 seq 轴切:Q/O 进 grid,K/V 进循环

Q/O 与 K/V 都沿着 seq 轴切块,块大小分别由 block_M 与 block_N 控制,但两者的去向不同:

1
2
3
4
5
6
seq ──->
├── Q₁ ├── Q₂ ├── Q₃ ┤ 按 block_M 切,进 grid:每个 CUDA block 认领一块 Q,
├── O₁ ├── O₂ ├── O₃ ┤ 独立算出对应的 O,块与块之间零依赖
├── K₁ ├── K₂ ├── K₃ ┤ 按 block_N 切,进内层循环:每个 block 沿 K/V 块逐块扫描;
├── V₁ ├── V₂ ├── V₃ ┤ K 与 V 必须按同样的边界切块(P̃ 的列与 V_j 的行须是同一批 key)
└────────────────────┘

为什么这样分?softmax 对分数矩阵做逐行归一化,O 的第 ii 行只由 Q 的第 ii 行决定,与 Q 的其他行无关。因此按行把 Q/O 切成 block_M 大小的块、分给不同的 CUDA block 并行计算,块间不需要任何通信。K/V 则不同:任何一行的 softmax 分母都要对整条 seq 求和,每一块 Q 都需要全部的 K/V。K/V 无法划归某个 block 独占,只能作为内层循环的遍历方向,按 block_N 逐块加载、逐块累积。

对比普通 GEMM 就能看出差别:C=ABC = AB 的输出块 CijC_{ij} 只依赖 AA 的第 ii 个行块与 BB 的第 jj 个列块,M、N 两个方向都可以切块并行;attention 的输出在 seq 方向上对 K/V 有全局依赖(softmax 分母是全行求和),这个方向只能串行遍历。block_M 决定并行度(grid 大小),block_N 决定每个 block 的循环长度

3.2 grid 与布局

张量布局用 BSHD([batch, seq_len, heads, dim]),batch/head 在外层、seq 连续,tile 拷贝才能访存合并。grid 三维:

1
with T.Kernel(T.ceildiv(seq_len, block_M), heads, batch, threads=threads) as (bx, by, bz):

每个 CUDA block 认领一个 Q tile(bx)、一个 head(by)、一个 batch(bz),沿 K/V 方向循环。这正是 FA-2 的循环序:Q 装载一次驻留 SRAM 全程不动。FA-1 是反过来的(外层 K/V、内层 Q),每轮 K/V 块的结果要经 HBM 中转更新各 Q 块,多出大量中间读写–FA-2 把循环反过来之后这条路才彻底堵死。

3.3 块大小的约束

SMEM 需求近似为:

QiBr×d+KjBc×d+VjBc×d+SijBr×BcMSRAM\underbrace{Q_i}_{B_r \times d} + \underbrace{K_j}_{B_c \times d} + \underbrace{V_j}_{B_c \times d} + \underbrace{S_{ij}}_{B_r \times B_c} \le M_{\text{SRAM}}

通常取 Br=BcMSRAM/dB_r = B_c \approx \sqrt{M_{\text{SRAM}} / d}。A100 上 M100M \approx 100 KB、d=128d = 128B128B \approx 128;再考虑 fragment 状态(acc_s、acc_o 等)占的寄存器,128×128 是常见起点,seq 很长或 head_dim 很大时倾向减小 block_N。

3.4 缓冲区清单

缓冲区 位置 形状 角色
Q_shared / K_shared / V_shared SMEM [block_M/N, dim] GEMM 的 A/B 操作数走 SMEM 路径
O_shared SMEM [block_M, dim] 写回前的中转(fragment 布局重排)
acc_s fragment(fp32) [block_M, block_N] S = QK^T 分块累加器
acc_s_cast fragment(fp16) [block_M, block_N] softmax 后的 P̃,cast 给第二个 MMA
acc_o fragment(fp32) [block_M, dim] 输出累加器,全程不归一化
scores_max / _prev / _scale / _sum / ell fragment(fp32) [block_M] online softmax 的逐行状态

一句话总结:所有逐行状态都放 fragment(寄存器),只有 GEMM 操作数和最终输出走 SMEM–这是三级内存抽象在 FA 里的标准分工。


四、主循环逐行拆解

完整代码基于官方 examples/flash_attention/example_mha_fwd_bshd.py(注释为本文所加)。这是仅推理版:不产出 backward 需要的 logsumexp,分母变量直接叫 ell

4.1 三层结构与初始化

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
@autotune(configs=get_configs(), warmup=10, rep=10)
@tilelang.jit(
out_idx=[3], # 第 4 个形参 Output 是输出
pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True},
)
def flashattn(batch, heads, seq_len, dim, is_causal,
block_M=64, block_N=64, num_stages=1, threads=128):
# softmax 的 1/sqrt(dim),预先乘上 log2(e),后面全部用硬件更快的 exp2
scale = (1.0 / dim) ** 0.5 * 1.44269504

@T.prim_func
def main(Q: T.Tensor(shape, dtype), K: T.Tensor(shape, dtype),
V: T.Tensor(shape, dtype), Output: T.Tensor(shape, dtype)):
with T.Kernel(T.ceildiv(seq_len, block_M), heads, batch,
threads=threads) as (bx, by, bz):
... # 缓冲区分配见 3.4

# Q tile 装载一次,全程驻留
T.copy(Q[bz, bx*block_M:(bx+1)*block_M, by, :], Q_shared)
T.fill(acc_o, 0)
T.fill(ell, 0)
T.fill(scores_max, -T.infinity(accum_dtype))

# causal 截断:Q 块最远看到 (bx+1)*block_M - 1,右侧整块跳过
loop_range = (
T.min(T.ceildiv(seq_len, block_N),
T.ceildiv((bx+1)*block_M, block_N))
if is_causal else T.ceildiv(seq_len, block_N)
)

三层结构是 TileLang 的标准混用写法:外层 @autotune 管配置搜索,中间 @tilelang.jit(out_idx=[3]) 声明输出形参,内部 @T.prim_func 管精确签名。

causal 截断值得推一遍:Q 块 bx 的行最远到 (bx+1)blockM1(bx+1) \cdot block_M - 1,key 位置超过它的全被掩码,对应的 K/V 块整块不进循环。N=4096N = 4096、块 128 时对角线右侧一半块直接消失,省一半计算。注意非 causal 时这个 T.min 退化为总块数,掩码换成右边界越界判断(见下节)。

4.2 掩码写进累加器:顺序是精髓

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
for k in T.Pipelined(loop_range, num_stages=num_stages):
T.copy(K[bz, k*block_N:(k+1)*block_N, by, :], K_shared)

# ① 先写掩码,再做 GEMM
if is_causal:
for i, j in T.Parallel(block_M, block_N):
acc_s[i, j] = T.if_then_else(
bx*block_M + i >= k*block_N + j, 0, -T.infinity(acc_s.dtype))
else:
for i, j in T.Parallel(block_M, block_N):
acc_s[i, j] = T.if_then_else(
k*block_N + j >= seq_len, -T.infinity(acc_s.dtype), 0)

# ② S = Q @ K^T(tensor core)
T.gemm(Q_shared, K_shared, acc_s, transpose_B=True,
policy=T.GemmWarpPolicy.FullRow)

掩码写在 GEMM 之前能成立,靠的是 T.gemm 的累加语义(C+=ABC \mathrel{+}= A B,所以纯 GEMM 例子里要先 T.clear):$-\infty + $ 有限值 == -\infty,被掩码的位置在 GEMM 里"存活"下来,之后 exp2(m)=0\mathrm{exp2}(-\infty - m) = 0,对分数和、对输出零贡献。省掉一遍 GEMM 后的掩码 pass。

两个分支的条件不同,处理的边界不同:

  • causal:query 全局位置 ≥ key 全局位置才保留(只许看过去);
  • 非 causal:k*block_N + j >= seq_len 置 −inf,处理的是 seq 不整除 block_N 时最后一块里的"假 key"。

4.3 online softmax 七步

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
# ③ online softmax 更新
T.copy(scores_max, scores_max_prev) # 1. 存历史 max
T.fill(scores_max, -T.infinity(accum_dtype))
T.reduce_max(acc_s, scores_max, dim=1, clear=False) # 2. 本块行 max
for i in T.Parallel(block_M):
scores_max[i] = T.max(scores_max[i], scores_max_prev[i]) # 3. 合并成 m_new
for i in T.Parallel(block_M):
scores_scale[i] = T.exp2( # 4. α = exp2(m_old − m_new)
scores_max_prev[i]*scale - scores_max[i]*scale)
for i, j in T.Parallel(block_M, block_N):
acc_s[i, j] = T.exp2( # 5. P̃ = exp2(S·scale − m·scale)
acc_s[i, j]*scale - scores_max[i]*scale)
T.reduce_sum(acc_s, scores_sum, dim=1) # 6. 本块 rowsum(P̃)
for i in T.Parallel(block_M):
ell[i] = ell[i]*scores_scale[i] + scores_sum[i] # 7. ℓ ← α·ℓ + 本块部分和
T.copy(acc_s, acc_s_cast) # fp32 → fp16,喂第二个 MMA

对应 2.3 节的递推式,逐步可查。两个容易踩的坑:

  • max 与 scale 的次序:先由历史 max 和本块 max 定出 mnewm_{new}由这一对 max 算出 α\alpha。max 是 scale 的来源,不是被更新对象;
  • scores_sumell 别混:前者只是本轮 rowsum(P̃) 的临时量,每次迭代被覆盖;ell 是带折扣链跨轮累积的总分母。归一化除的是 ell,除成某一块的局部和就错了。

4.4 第二个 GEMM 与收尾

1
2
3
4
5
6
7
8
9
10
11
    # ④ 先打折旧成果,再加新项
for i, j in T.Parallel(block_M, dim):
acc_o[i, j] *= scores_scale[i]
T.copy(V[bz, k*block_N:(k+1)*block_N, by, :], V_shared)
T.gemm(acc_s_cast, V_shared, acc_o, policy=T.GemmWarpPolicy.FullRow)

# ⑤ 循环外:唯一一次归一化,然后写回
for i, j in T.Parallel(block_M, dim):
acc_o[i, j] /= ell[i]
T.copy(acc_o, O_shared)
T.copy(O_shared, Output[bz, bx*block_M:(bx+1)*block_M, by, :])

acc_o *= scores_scale 必须在第二个 GEMM 之前:先加新项再打折,新项就被错误地折旧了。归一化只在循环外做一次,中途除会错–分母没定型。

写回走 fragment → O_shared → Output 两跳:fragment 里的数据按 MMA 布局散在各线程寄存器,先落 SMEM 重排布局,再合并写 global。

4.5 工程细节清单

  • exp2 技巧:把 1/d1/\sqrt{d} 预乘 log2e1.4427\log_2 e \approx 1.4427 折进 scale,全部指数运算用 exp2 硬件指令(比 exp 快),FA 实现的标配;
  • P̃ 的 fp16 cast 是全 kernel 唯一降精度点:MMA 只吃低精度操作数,acc_s/acc_o/ell 等累加器与状态全程 fp32;
  • FullRow warp policyT.gemm 把输出 tile 切给 block 内各 warp 的切法。FA 的 gemm#1 后紧跟按行归约(reduce_max/reduce_sum)和按行缩放(acc_o *= scores_scale),FullRow 保证一行完整住在一个 warp 内,这些操作全部退化为 warp 内 shuffle;若用默认 Square,第 i 行的左右两半住在两个 warp 里,归约就得经 SMEM 中转外加 bar.sync–结果仍正确,但每轮迭代多两次往返:
1
2
3
4
5
6
7
8
Square(默认):每 warp 近似方块    FullRow:按 M 横切,每 warp 拿全宽整段行
┌─────────┬─────────┐ ┌──────────────────────┐
│ warp0 │ warp1 │ │ warp0 行 0..31 │
│ 64×64 │ 64×64 │ ├──────────────────────┤
├─────────┼─────────┤ │ warp1 行 32..63 │
│ warp2 │ warp3 │ ├──────────────────────┤
│ 64×64 │ 64×64 │ │ warp2/3 行 64..127│
└─────────┴─────────┘ └──────────────────────┘
  • forward 用 num_stages=1:循环体是两个相互依赖的 GEMM 夹一串 elementwise(S 等 K、P̃ 等 S 和 m),预取重叠窗口小而寄存器压力大,官方权衡后的选择。流水线什么时候真的有用,见上一篇的三因素模型与对照实验;
  • O 不中途归一化是 FA-2 的关键改动之一(见下节)。

一句话总结:主循环五行骨架–搬 K、写掩码、GEMM 喂料、online softmax 三件套、rescale 后 GEMM 消费;每一步都有"顺序不能错"的理由。


五、与 FA-2 论文 Algorithm 1 的对照

这份代码就是 FlashAttention-2 论文 Algorithm 1 的逐行实现。伪代码符号 ↔ 代码变量:

伪代码 内容 代码对应
Br×BcB_r \times B_c 分块大小 block_M × block_N
for i ≤ Tᵣ(外层 Q 块) 外层循环进 grid with T.Kernel(...) as (bx, by, bz)
QiQ_i 进 SRAM 装载后驻留 T.copy(Q[...], Q_shared)
O0=0, 0=0, m0=O^0 = 0,\ \ell^0 = 0,\ m^0 = -\infty 初始化三状态 T.fill(acc_o/ell, 0)T.fill(scores_max, -inf)
for j ≤ T_c(内层 K/V 块) 内层循环 for k in T.Pipelined(loop_range, ...)
Sij=QiKjS_{ij} = Q_i K_j^\top 分数块 T.gemm(Q_shared, K_shared, acc_s, transpose_B=True)
mm 更新、P~\tilde{P}\ell 更新 online softmax 一族 reduce_maxscores_scaleexp2ell 递推
Odiag(eΔm)1O+P~VjO \leftarrow \mathrm{diag}(e^{\Delta m})^{-1} O + \tilde{P} V_j 先打折再加新项 acc_o *= scores_scalegemm(acc_s_cast, V_shared, acc_o)
O/O/\ell(第 12-13 行) 循环外归一化 acc_o /= ell

四处伪代码没写、真实 kernel 必须处理的差异:

  1. 底数:论文 e 底,实现全用 exp2 硬件指令,scale 预乘 log2e\log_2 e
  2. scale 折叠1/d1/\sqrt{d} 不显式乘,折进指数表达式省一遍逐元素乘;
  3. acc_s_cast:P̃ 从 fp32 fragment cast 成 fp16 才能进 MMA;
  4. causal 截断:算法按稠密写,代码用 loop_range 让对角线右侧整块跳过。

顺带把 FA-1 → FA-2 的三个改动列清楚,这份代码全部站在 FA-2 一边:

FA-1(2022) FA-2(2023)
外层循环 K/V 块 Q 块
中间读写 各 Q 块的结果需经 HBM 中转更新 Q 驻留 SRAM,块间零 HBM 往返
O 的归一化 逐块保持已归一化(多两次乘除) 循环外一次性 /= ell
并行度 受 K/V 块数限制 Q 块 × head × batch,并行度更高
warp 分工 均分 K/V 切给不同 warp,减少同步(TileLang 里由 GemmWarpPolicy 表达)

六、验证、测速与调参

写完 kernel 的三个标准动作(数值验证 -> dump 生成码 -> 基准测试),在 FA 上一个不少:

① 数值验证profiler.assert_allclose 直接接受 PyTorch 参考实现做正确性比对,不需要自己写验证框架:

1
2
3
4
5
kernel = flashattn(batch, heads, seq_len, dim, is_causal,
block_M=128, block_N=128, num_stages=1, threads=128)
ref = partial(ref_program, is_causal=is_causal) # einsum 朴素注意力 + tril 掩码
profiler = kernel.get_profiler()
profiler.assert_allclose(ref, rtol=0.01, atol=0.01)

容差 0.01 是 fp16 输出的合理范围;更严的判定方法(与 fp32 参考比、大误差元素计数)见上一篇第九节。

② dump 生成码kernel.get_kernel_source() 打印完整 CUDA 源码。对 FA 值得确认三件事:T.copy 降成了 cp.async 还是 TMA、FullRow 下 warp 怎么切行、exp2 是否真的成了硬件指令。

③ 测速profiler.do_bench(warmup=500) 自动 warmup 多次取统计。TFlops 按 2BHN2d×22 \cdot B \cdot H \cdot N^2 \cdot d \times 2 个 matmul 计算,causal 乘 0.5:

1
2
3
4
flops = 2.0 * batch * heads * seq_len * seq_len * dim   # 单个 matmul
total_flops = 2 * flops * (0.5 if is_causal else 1.0)
latency = profiler.do_bench(warmup=500)
print(f"{total_flops / latency * 1e-9:.2f} TFlops")

调参交给 @autotune:把 block_M/block_N/num_stages/threads 写成带默认值的参数,搜索空间在 get_configs() 里列全(笛卡尔积),机器自己选。注意首跑很慢(每组都要编译 + 测速),调通阶段先裁成单组。


七、总结

  1. 数学在前:softmax 平移不变性 → 安全 softmax → 分块下的换基准问题 → online softmax 递推。FlashAttention 只是这条递推式与两个 GEMM 的交织,精确、非近似;
  2. 切法是结构:Q/O 与 K/V 都沿 seq 轴切块,block_M 控制的 Q 块进 grid(决定并行度),block_N 控制的 K/V 块进内层循环(决定遍历长度);softmax 的逐行归约使 seq 方向对 K/V 产生全局依赖,只能作为循环方向,这是 FA 与 GEMM 的本质结构差异;
  3. 实现靠 TileLang 五原语T.Kernel 定切法、T.copy 管搬运、T.gemm 喂两个 matmul、T.Pipelined 管 K/V 流水(FA 前向官方权衡后用 stages=1)、T.Parallel + reduce_* 承载 online softmax 三件套。掩码写进累加器、exp2 折叠、FullRow 行所有权、O 延迟归一化,四个细节决定了这份代码"像论文"还是"只是能跑"。

可迁移的启示:读 FA 代码的正确顺序是先读递推式再读循环–所有 tile 级 kernel 都是"一条数学递推 + 一个 tile 循环",TileLang 只是把后者写到了 30 行的量级。


参考

  • FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(Dao et al., NeurIPS 2022, arXiv:2205.14135)
  • FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning(Dao, 2023, arXiv:2307.08691)
  • Online normalizer calculation for softmax(Milakov & Gimelshein, 2018, arXiv:1805.02867)
  • TileLang GitHub(v0.1.13):examples/flash_attention/example_mha_fwd_bshd.py 为本文代码底本
  • 本站:《TileLang 编程基本知识点

本文基于 TileLang v0.1.13 官方示例与 FlashAttention-2 论文整理,代码注释为本文所加。

TileLang 的定位可以用一句话概括:把 CUDA kernel 里「怎么切 tile、数据放哪级存储、什么时候预取」这几件事变成显式的 Python 语句,其余的寄存器分配、指令选择、同步插入交给编译器。它不是又一个 Triton–Triton 隐藏 shared memory,TileLang 让你直接写 T.alloc_shared。这个差别决定了两者能碰到的性能天花板不同。

本文整理 TileLang 的编程基本知识点:它是什么、编程模型的 5 个原语、循环原语与 T.Pipelined 的工作原理、同一份代码在不同 GPU 架构上如何 lowering、四个进阶主题(Split-K / warp specialization / autotune / Blackwell 两条路径),以及写完 kernel 之后的三个标准动作(验证 / dump 源码 / bench)。


一、TileLang 是什么

项目 现状(2026-08)
版本 v0.1.13(2026-08-03 发布)
GitHub tile-ai/tilelang,7.2k stars
底层 TVM(IR 已迁移到 TIRX)
Python ≥ 3.10
主力后端 CUDA(SM70~SM120)
其他后端 ROCm/HIP、Apple Metal、LLVM CPU(实验)、CuTe DSL(实验)、WebGPU(实验)
生态后端 华为 Ascend、沐曦 MACA、摩尔线程 MUSA(独立仓库维护)

出身是学术项目:主要由 LeiWang1999、chengyupku、nox-410 在北大杨智教授指导下开发,部分工作在 MSRA 实习期间完成。2025-01 开源。

值得注意的是上游模型厂在用它:TileLang 仓库的 examples/ 里有 deepseek_mladeepseek_v32deepseek_v4deepseek_mhc 四个目录,DeepSeek 系列的 MLA / 稀疏注意力 / mHC 融合 kernel 都有 TileLang 参考实现。这意味着读 TileLang examples 等于读一份最新算子的可执行论文附录–这是它相比 Triton 的一个实际优势。

一句话总结:TileLang = Pythonic 语法 + 显式 tile/memory 层级控制 + TVM 编译基础设施,目标是「写起来像 Triton,控制力接近 CUTLASS」。


二、编程模型:5 个原语撑起全部

TileLang 的 API 面很窄,这是刻意的。一个完整 kernel 基本只用这 5 类原语:

1
2
3
4
5
6
7
8
T.Kernel(grid_x, grid_y, threads=N)     ← ① 定义 grid / block,拿到 block index

├── T.alloc_shared(shape, dtype) ← ② 显式声明 shared memory buffer
├── T.alloc_fragment(shape, dtype) ← ②' 显式声明 register fragment(累加器)

└── for k in T.Pipelined(N, num_stages=3): ← ③ 软件流水(自动双/三缓冲)
T.copy(global_tile, shared) ← ④ 数据搬运(按架构 lower 到 cp.async / TMA)
T.gemm(A_s, B_s, C_frag) ← ⑤ tile 级 MMA(映射到 Tensor Core)

完整的 FP16 GEMM + ReLU(基于官方 quickstart):

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
32
33
34
35
36
37
38
39
40
41
42
43
44
import torch
import tilelang
import tilelang.language as T


@tilelang.jit
def matmul(A, B, block_M: int, block_N: int, block_K: int):
M, N, K = T.const("M, N, K")
dtype = T.float16
accum_dtype = T.float32
A: T.Tensor((M, K), dtype)
B: T.Tensor((K, N), dtype)
C = T.empty((M, N), dtype)

# grid: (N 方向块数, M 方向块数),每 block 128 线程
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
A_shared = T.alloc_shared((block_M, block_K), dtype) # shared tile
B_shared = T.alloc_shared((block_K, block_N), dtype)
C_local = T.alloc_fragment((block_M, block_N), accum_dtype) # 寄存器累加器,fp32

T.clear(C_local)

# 三级软件流水:搬 k+1 块的同时算 k 块
for ko in T.Pipelined(T.ceildiv(K, block_K), num_stages=3):
T.copy(A[by * block_M, ko * block_K], A_shared) # global -> shared
T.copy(B[ko * block_K, bx * block_N], B_shared)
T.gemm(A_shared, B_shared, C_local) # tensor core,fp32 累加

# epilogue:ReLU
for i, j in T.Parallel(block_M, block_N):
C_local[i, j] = T.max(C_local[i, j], 0)

T.copy(C_local, C[by * block_M, bx * block_N]) # fragment -> global

return C


M = N = K = 1024
# 用静态 shape 编译出可执行 kernel
matmul_kernel = matmul.compile(M=M, N=N, K=K, block_M=128, block_N=128, block_K=32)

a = torch.randn(M, K, device="cuda", dtype=torch.float16)
b = torch.randn(K, N, device="cuda", dtype=torch.float16)
c = matmul_kernel(a, b)

30 行写完一个带 fused epilogue 的 Tensor Core GEMM。几个关键点:

  1. T.const("M, N, K") 声明符号化 shape.compile(...) 按实际 shape 特化出可执行 kernel(也可以直接调用让 @tilelang.jit 在首次调用时惰性编译)。
  2. 累加器 dtype 和存储 dtype 分离C_local 是 fp32 fragment,写回时才降到 fp16。这是数值稳定的标准配方。
  3. epilogue 用 T.Parallel 表达,不需要另起 kernel–省一次 HBM 往返。
  4. 没有一行同步代码__syncthreads() 由编译器根据 T.Pipelined 的依赖关系自动插入。

一句话总结:TileLang 的 API 面窄到只有 5 类原语,但这 5 类恰好覆盖了 GPU kernel 性能的全部决定因素–tile 怎么切、数据放哪级存储、什么时候预取。


三、循环原语:从 T.serial 到 T.Pipelined

GEMM 例子里已经出现过两种循环(T.PipelinedT.Parallel)。TileLang 的循环构造一共就这四个,按暴露给编译器的并行度递进,先顺次过一遍:

原语 语义 典型用途
T.serial 普通 for 循环,迭代之间有依赖 递推、边界处理
T.unroll 要求编译器完全展开 小循环,省掉分支和循环开销
T.Parallel 嵌套并行循环,所有迭代互相独立 elementwise、epilogue
T.Pipelined 软件流水,生产者-消费者跨迭代重叠 GEMM / attention 主循环
1
2
3
4
5
6
7
8
9
10
11
for i in T.serial(N):                        # 串行:下一轮依赖上一轮的结果
...

for k in T.unroll(K_TILE): # 展开:编译期摊平
acc += a[k] * b[k]

for i, j in T.Parallel(M, N): # 并行:迭代独立,映射到线程
C[i, j] = A[i, j] + B[i, j]

for ko in T.Pipelined(iters, num_stages=3): # 流水:copy 与 compute 时间重叠
...

补充几点:

  • T.serial 支持三参数形式 T.serial(0, N, 2)(起点、终点、步长);
  • T.Parallel 可以加 coalesced_width= 提示控制访存合并宽度,loop_layout= 挂 fragment layout 标注;
  • 另有一个高级构造 T.Persistent,表达 persistent thread-block 风格的循环(5.1 提到的 stream-K 变体就靠它);
  • Python 原生的 if/elsewhilebreak/continue 都可用,条件是 TIR 表达式即可;潜在的越界访问由 LegalizeSafeMemoryAccess pass 自动加 guard(见第七节)。

前三个原语都好理解,真正值得单独一节展开的是最后一个。

T.Pipelined 做了什么

上面例子里最「魔法」的一行是 for ko in T.Pipelined(...)num_stages=3 不是「循环展开 3 次」,而是建立 3 级软件流水:

1
2
3
4
5
6
7
时间 ->
iter 0: [copy k=0]
iter 1: [copy k=1] [gemm k=0] ← 进入稳态:搬运与计算重叠
iter 2: [copy k=2] [gemm k=1]
...
iter N-1: [gemm k=N-2] ← epilogue:只剩计算
└─ HBM 延迟被后续 iter 的计算隐藏

手写 CUDA 要实现同样效果,需要自己管理 stage 数组下标、cp.async 的 commit/wait group、以及每级之间的 barrier。TileLang 把这压缩成一个参数。具体地,编译器在这一个循环上做了四件事:

  1. 循环重写:把源代码里的单层循环拆成 prologue(预取前 N-1 轮)/ 稳态 body(搬第 k+N-1 块的同时算第 k 块)/ epilogue(算完尾部) 三段。稳态时 copy 和 compute 在时间上重叠。
  2. 共享内存多缓冲num_stages=3 意味着 A_shared / B_shared 会被自动复制成 3 份(三缓冲),生产者写第 k+2 块、消费者读第 k 块,互不冲突。你在源码里写的是一份 buffer,编译器做的是 buffer 乘法。
  3. 异步拷贝插入:流水线里的 T.copy 会 lower 成异步拷贝(Ampere+ 上是 cp.async,Hopper+ 上是 TMA 的 cp.async.bulk),并自动配好 commit_group / wait_group 的配对。
  4. 同步插入:编译器的 PipelinePlanning / InjectSoftwarePipeline / InjectTmaBarrier 等一系列 pass 负责推导生产者-消费者依赖,在正确的地方插入 __syncthreads()(或 mbarrier)。这就是为什么源代码里一行同步都没有。

注意 T.copy 本身的语义是同步的–语句结束后 dst 就可读,如果 lower 到了异步指令,编译器会补上 wait 保证这一点。想手动控制异步,用 T.async_copy(不自动插 wait,需要自己写 T.ptx_wait_group)。

手动标注 stage / order

常规 GEMM 形态的流水线,num_stages=N 就够了,编译器自己推断谁是生产者谁是消费者。当循环体顺序不寻常(比如想让「下一轮的 copy」排在「这一轮的 compute」之前发射)时,可以显式标注:

1
2
3
4
5
6
7
for ko in T.Pipelined(
num_tiles,
stage=[0, 1], # copy 是 stage 0,gemm 是 stage 1
order=[1, 0], # 发射顺序上 gemm 先、copy 后
):
T.copy(A[ko * BK], A_shared)
T.gemm(A_shared, B_shared, C_local)

规则:

  • stage / order 与循环体内的可调度语句(copy、gemm、reduction、store、同步)按源码顺序一一对齐;
  • 流水线深度由 max(stage) + 1 推断,此时不要再传 num_stages
  • 循环体里的标量别名(base = ko * BK 这类 Bind 语句)不占标注位–它们没有副作用,编译器会在每个消费者处按需重放;
  • 编译器会校验依赖:生产者的 stage 必须不晚于消费者,同 stage 内 order 必须生产者在消费者前。

一句话总结T.Pipelined(num_stages=N) = 循环三分重写 + shared memory N 重缓冲 + 异步拷贝 + 自动同步,这是 GEMM/attention 类 kernel 能打到带宽/算力天花板的全部前提。


四、同一份代码,不同架构的 lowering

TileLang 源码里只有 T.copyT.gemm 这两个「意图」,落到哪条指令由 target 决定。这是它和直接写 CUDA 最大的分工差异–你描述数据流和 tile 结构,编译器按架构选指令

架构 异步拷贝(流水线内的 T.copy MMA 指令(T.gemm
SM70~75(V100/T4) SIMT ld.global + st.shared mma.sync(fp16)
SM80~89(A100/3090/4090/Ada) cp.async 多级流水 mma.sync
SM90a(H100/H200) TMA(cp.async.bulk.tensor)+ mbarrier wgmma(warp group MMA)
SM100a(B100/B200) TMA + mbarrier tcgen05.mma(TMEM 累加)
SM120(RTX 50 / RTX PRO,消费级 Blackwell) 没有 TMA,走 cp.async(LDGSTS) 普通 mma(第五代 tensor core)

这里有个容易踩的坑:不能按 SM 版本号大小推测能力。SM120 数字上大于 SM90,但它没有 TMA、没有 tcgen05,异步拷贝走的是 Ampere 时代引入的 cp.async。在 SM120 卡上写完全相同的 T.Pipelined 代码,编译器会自动退回 cp.async 风格的流水–语义不变,只是底层搬运指令不同。这正是 T.Pipelined 这个抽象的价值:你表达的是「我要 N 级预取流水」这个意图,而不是「我要发 cp.async.bulk.tensor」这个指令。

target 通过三种方式指定:

1
2
3
4
5
6
7
kernel = tilelang.compile(func, target={"kind": "cuda", "arch": "sm_90"})

@tilelang.jit(target="cuda") # 或裸字符串
def factory(...): ...

# 或环境变量(适合整台机器固定 GPU 型号的场景)
# export TILELANG_DEFAULT_TARGET='{"kind": "cuda", "arch": "sm_90"}'

auto(默认)按 CUDA → HIP → Metal 顺序探测。arch 直接对应 NVCC 的 -arch=sm_XX;需要一份代码出多个 SASS 时用 code 列表(fatbin)。跨厂商同理:HIP 配 mcpu="gfx90a",Metal / LLVM CPU / WebGPU 各有对应 kind。

判断「我写的代码在这张卡上到底变成了什么」,最直接的办法还是后面第六节的 dump 源码–看生成的 CUDA 里是 cp.async、TMA descriptor 还是 wgmma,比读文档更可靠。


五、进阶主题

5.1 Split-K:grid 第三维 + atomic_add

M、N 小而 K 很大时,M×N 的 tile 数不够填满 SM,就把 K 切开分给多个 block 并行累加。TileLang 里不需要新原语–T.Kernel 的第三个 grid 维就是 split 因子,最后用 T.atomic_add 归约(来自 examples/gemm_splitk):

1
2
3
4
5
6
7
8
9
10
11
12
splitK = K // split_k

with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), split_k, threads=128) as (bx, by, bz):
...
T.clear(C_local)
for ko in T.Pipelined(T.ceildiv(splitK, block_K), num_stages=0):
T.copy(A[by * block_M, bz * splitK + ko * block_K], A_shared) # K 维偏移带上 bz
T.copy(B[bz * splitK + ko * block_K, bx * block_N], B_shared)
T.gemm(A_shared, B_shared, C_local)

for i, j in T.Parallel(block_M, block_N):
T.atomic_add(C[by * block_M + i, bx * block_N + j], C_local[i, j])

要点:bz 只出现在 K 维索引里;累加结束整体做一次 atomic_add(而不是每个元素多次原子写);输出 C 必须先清零;num_stages=0 是官方示例关掉了自动流水(生产代码里该开的还是要开)。注意 atomic_add 的归约顺序不定,数值不可复现–对复现性有要求的场合要改成两阶段确定性归约(partial 先写回 workspace 再单独 reduce),本博客 DeepSeek-V4 mHC Pre-Block 融合 Kernel 详解 里有完整分析。同一目录下还有 stream-K 变体(gemm_streamk,把尾部 wave 按 K 拆给 peer block 再 fixup),是 persistent kernel 的入门样本。

5.2 warp specialization:T.ws + mbarrier

Hopper 之后,生产者和消费者可以拆成不同 warp group,各自跑各自的「循环」,靠 mbarrier 握手。TileLang 用 T.ws(i) 划分角色、T.alloc_barrier 建握手信号(来自 examples/warp_specialize):

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
with T.Kernel(..., threads=256) as (bx, by):
A_shared = T.alloc_shared((block_M, block_K), dtype)
...
data_is_ready = T.alloc_barrier(arrive_count=128)
compute_is_done = T.alloc_barrier(arrive_count=128)

with T.ws(1): # 消费者 warp group
T.clear(C_local)

for ko in T.Pipelined(T.ceildiv(K, block_K), num_stages=0): # 手动流水,不走自动重写
with T.ws(0): # 生产者:等上一轮算完 → TMA 搬数 → 通知就绪
T.barrier_wait(compute_is_done, (ko + 1) % 2)
T.tma_copy(A[by * block_M, ko * block_K], A_shared, barrier=data_is_ready)
T.tma_copy(B[ko * block_K, bx * block_N], B_shared, barrier=data_is_ready)
T.barrier_arrive(data_is_ready)
with T.ws(1): # 消费者:等数就绪 → gemm → 通知算完
T.barrier_wait(data_is_ready, ko % 2)
T.gemm(A_shared, B_shared, C_local)
T.barrier_arrive(compute_is_done)

with T.ws(1):
T.copy(C_local, C[by * block_M, bx * block_N])

要点:

  • num_stages=0 关掉自动流水–因为流水逻辑已经由两个 warp group 的 barrier 协议手动表达了,双缓冲体现在 barrier_wait% 2 相位翻转上;
  • T.tma_copy(..., barrier=...) 是显式 TMA 入口,比 T.copy 更低一层;
  • 这是 TileLang 里「控制力接近 CUTLASS」的具体形态:mbarrier 的 arrive_count、相位奇偶全部由你负责。examples 目录里同族还有 barrierpipe / softpipe 等多种流水协议写法可以对照。

5.3 autotune:tile 参数交给搜索

block_M / block_N / block_K / num_stages / threads 这组参数的理论最优值依赖具体 GPU 和问题规模,手调不现实。TileLang 内置 autotuner,用法是把可调参数写成带默认值的函数参数,再套一层装饰器:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
def matmul_configs(M, N, K):
return [
dict(block_M=BM, block_N=BN, block_K=BK, num_stages=S, threads=TH)
for BM in [64, 128]
for BN in [64, 128]
for BK in [32, 64]
for S in [2, 3]
for TH in [128, 256]
]

@tilelang.autotune(configs=matmul_configs, warmup=25, rep=100, timeout=60)
@tilelang.jit(out_idx=[-1])
def matmul(M: int, N: int, K: int,
block_M: int = 128, block_N: int = 128, block_K: int = 32,
threads: int = 128, num_stages: int = 3, ...):
...

with set_autotune_inputs(a, b, c): # 固定输入,保证各 config 可比
tuned = matmul(M, N, K) # 编译+验证+benchmark 全部 config,返回最优

值得知道的工程细节:候选 kernel 并行编译、逐个 benchmark;每个 config 会先过正确性检查(ref_prog 或默认的 torch 对比,容差 rtol/atol 默认 1e-2);结果缓存在 ~/.tilelang/cache/autotuner,缓存 key 包含 TileLang 版本 + 函数源码 + config 列表,改代码自动失效。调稳定后建议把最优 config 烘焙成函数默认值写回源码,autotune 只作为开发期工具。

5.4 Blackwell 两条路径

Blackwell 不是一个统一架构,写跟架构相关的 kernel 时必须分开看:

SM100a(B100/B200,数据中心) SM120(RTX PRO / RTX 50,消费级)
异步拷贝 TMA + mbarrier cp.async(LDGSTS),无 TMA
MMA 指令 tcgen05 + TMEM 普通 mma
two-SM(2-CTA)kernel
NVFP4 block-scaled T.mma_gemm_blockscaled T.mma_gemm_blockscaled(2026-07-30 加入)
对应 example blockscaled_gemm_sm100gemm_tcgen05 gemm_sm120

SM100a-- tcgen05 路径。第五代 Tensor Core 的 MMA 指令 tcgen05.mma 从 shared memory 直读操作数、累加到独立的 Tensor Memory(TMEM),还支持两个 CTA 配对发射。对应到 TileLang 是一组新原语:T.alloc_tmem 分配 TMEM 累加器,T.tcgen05_gemm 发射 MMA(不带隐式等待),T.alloc_barrier + T.mbarrier_wait_parity(mbar, k % 2) 手动做相位同步。这条路径目前是实验性 preview:同步协议要自己写,官方 README 明说「manual implementation required」。较新版本提供了半自动入口 T.gemm(..., mbar=...)–发射后自动插入匹配的 mbarrier_wait_parity,并把 fence 插入交给 InjectTcgen05Fence pass。examples/gemm_tcgen05/ 下有从裸 tcgen05_gemm 到 warp-specialized persistent、再到 2-CTA stream-K 的完整梯度。

SM120-- 传统路径。没有 tcgen05、TMEM 和 TMA,走的仍是 SM80 风格的 mma.sync + cp.async 流水线,TileLang 现有代码基本直接可用。换句话说,为 Hopper 写的 kernel 迁到 RTX 50 通常只是换个 arch,迁到 B200 才需要考虑 TMEM 那套新原语。两边唯一真正共享的新能力是 NVFP4 block-scaled MMA(T.mma_gemm_blockscaled,SM120 路径 2026-07-30 加入,可对照 Colfax 那篇 NVFP4 Blockscaled GEMM on RTX Pro Blackwell (sm12x)),但底层指令并不相同。

编译 target 上,数据中心 Blackwell 需要 fatbin 时可以 {"kind": "cuda", "arch": "sm_100f", "code": ["sm_100a", "sm_103a"]} 一份代码出多个 SASS。


六、写完之后的三个动作

quickstart 的官方流程把「kernel 写完之后」固化成了三步,建议形成肌肉记忆:

① 验证正确性–先于一切性能讨论,且以 fp32 参考为准

1
2
ref32 = torch.relu(a.float() @ b.float())
torch.testing.assert_close(c.float(), ref32, rtol=1e-2, atol=0.05)

为什么要跟 fp32 参考比而不是直接 torch.relu(a @ b):kernel 内部用 fp32 累加,a @ b 则是 fp16 累加,拿后者做参考的话误差来源两边不一致–你分不清到底是自己 kernel 错了还是累加精度差异。用 a.float() @ b.float() 做基准,差异才能归因到 kernel 本身。K=1024 累加下 atol=0.05 是合理范围。更复杂的 kernel(attention、MoE)可以保留一个 PyTorch 参考实现专门做这件事;TileLang 的 autotuner 也用同样思路验证每个候选 config。

② dump 生成源码–看编译器到底做了什么:

1
2
cuda_source = matmul_kernel.get_kernel_source()
print(cuda_source)

返回的是最终生成的完整 CUDA 源码。这一步回答所有「lowering 疑问」:T.copy 变成了 cp.async 还是 TMA descriptor?T.gemmmma.sync 还是 wgmma?多缓冲分配了多少 shared memory?__syncthreads() 插在了哪?调性能之前先读一遍生成代码,能省掉大量盲猜。换个 num_stagesblock_K 再 dump 一次、diff 两份源码,比读任何文档都直观。

③ benchmark–拿可信的延迟数字:

1
2
profiler = matmul_kernel.get_profiler(tensor_supply_type=tilelang.TensorSupplyType.Normal)
latency = profiler.do_bench() # ms

get_profiler 会自动生成输入、warmup、多次重复取统计,不用自己写 torch.cuda.synchronize() + time.time() 那套容易测错的东西。注意 TensorSupplyType.Normal 指定用正态分布造输入–对 GEMM 无影响,但对带 exp 的 attention kernel 会影响数值路径,别用全 0。没有 JITKernel 对象时也可以直接 from tilelang.profiler import do_bench 包一个 callable。有了 ① 的正确性和 ③ 的基线数字,后面任何改动(换 tile 尺寸、加 swizzle、上 warp specialization)都是可度量的。


七、学习路径与调试工具

环境与选卡

1
2
pip install tilelang
python -c "import tilelang; print(tilelang.__version__)"

需要 Python ≥ 3.10。想要最新特性走 nightly:pip install tilelang --find-links https://tile-ai.github.io/whl/nightly

没有 GPU 怎么办:TileLang 的 CUDA 后端需要真卡才能编译执行,Apple Silicon 可以走 metal target 但例子覆盖有限。学 TileLang 属于典型的短时高强度用卡–跑几小时 examples 就停,不需要长期占资源,RunPod 按小时租一张卡是最省事的路子。

选卡要看想学什么(原因见第四节):

学习目标 需要的卡
基础 tile / T.Pipelined / autotune 任意 SM80+,一张 4090 就够
TMA、T.tma_copy、warp specialization 必须 SM90a(H100/H200)
tcgen05 MMA、TMEM、two-SM kernel 必须 SM100a(B100/B200)
SM120 NVFP4 block-scaled RTX PRO 6000 / RTX 50 系

别拿 SM120 的卡去学 TMA–消费级 Blackwell 没有这个硬件单元。

调试工具箱

TileLang 在调试工具链上比 Triton 强,别浪费:

  • T.print(buffer, msg=...):kernel 内部打印 shared/fragment buffer,TileLang 自动只从单个线程打印避免刷屏;配合 if i == 0: 谓词用。
  • T.device_assert(cond, msg):device 侧断言,CUDA target 上生效,排查越界和 NaN 比注释掉半段代码快得多。
  • Pass Visualizer(2026-07 加入):结构树浏览器,看每个编译 pass 对 IR 做了什么。
  • IR Lower Trace(2026-07 加入):逐 pass dump IR,定位「我写的 layout 到哪一步被改掉了」。
  • TileLang LSP(2026-08 开源):VSCode 里显示 buffer 的 shape / dtype / scope / 推断 layout 的 inlay hint,写 kernel 时不用反复回头查 shape,第一天就该装上。
  • layout 可视化:把 fragment layout 画出来,检查 bank conflict。
  • get_kernel_source():见第六节,最强的「调试器」其实是读生成的 CUDA。
  • 缓存目录 ~/.tilelang/cache/:autotuner 产物、编译出的 .so / cubin 都在这;改了代码行为不对时先怀疑缓存,TILELANG_DISABLE_CACHE=1 一键排除。
  • 越界防护:LegalizeSafeMemoryAccess pass 会在可能越界的访问处自动插 guard(证明安全则自动消除),所以边界处理很多情况下不用手写 if–但自定义边界逻辑仍建议显式写。

推荐顺序

阶段 材料 目标
TileLang Puzzles 10 题 建立 tile 思维,比读文档有效得多
quickstart.py + elementwise 摸清 5 个原语
语言基础文档 补齐 layout / 内存作用域概念
examples/gemm swizzle、autotune、架构特化
examples/flash_attention online softmax + 双 GEMM 融合
examples/gemm_fp8 / blockscaled_gemm_sm100 量化 GEMM,per-block scale
examples/gemm_splitk / warp_specialize 本文第五节的出处
examples/deepseek_mla / deepseek_v4 / deepseek_mhc 真实生产算子

Puzzles 优先这点要强调:TileLang 的文档偏参考手册风格,直接读容易只学到 API 名字、学不到「为什么这么切」。10 个难度递增的 puzzle 会强迫你自己想清楚 tile 划分。第 ⑧ 步的收益也不在语法,而在于看懂「一个生产级 LLM 算子是如何被拆成 tile 数据流的」。

一句话总结:语法半天就能过完,TileLang 的学习成本在「tile 级思维」–而这只能靠 Puzzles 和 examples 里那些真实 kernel 攒出来。


参考

线性注意力与 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)

本文的实验都跑在单卡上,8B 级别的 target model + draft model 一张 80GB 卡就够。想复现这套投机解码对比的话,不必自己攒机器——RunPod 上按小时租一张 H100/H200 即可,nsys 采集需要的 --cap-add=SYS_ADMIN 权限它的容器实例也放开了。

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 源码: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
  • 复现环境:RunPod 单卡实例(按小时计费,适合这类短时评测)

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 尚未开源。