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

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 可行性的前提。当一个数值技巧成为架构设计的必要条件时,它就不再是实现细节。

参考