TileLang 实战:KDA 从零到一–Kimi Delta Attention
前三篇分别实现了无遗忘的 chunked 线性注意力、标量衰减、以及带删除项的 Gated DeltaNet。本篇是系列收尾:把 GDN 的标量 门 α t \alpha_t α t 换成逐通道向量 门 a t ∈ R d k \bm{a}_t \in \mathbb{R}^{d_k} a t ∈ R d k 。
S t = S t − 1 Diag ( a t ) ( I − β t k t k t ⊺ ) + β t v t k t ⊺ \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}
S t = S t − 1 Diag ( a t ) ( I − β t k t k t ⊺ ) + β t v t k t ⊺
改动只有一处:α t \alpha_t α t 变成了 Diag ( a t ) \operatorname{Diag}(\bm{a}_t) Diag ( a t ) 。但这一处把前三篇积累的所有便利拆掉了–衰减不再能从矩阵里外提,Γ \Gamma Γ 从 C × C C\times C C × C 变成 C × C × d k C \times C \times d_k C × C × d k ,而累积衰减积的通道离散度会直接把 fp16 打穿。
本篇聚焦实现 。KDA 的数学推导(递推式、逐通道 WY 表示、UT 变换、下界衰减与满秩门控的动机)见《KDA 的来龙去脉 》§3–§4,这里不重复;本文只做一件事:把那些公式落成能跑的 kernel,并量化每一处数值边界 。
前三篇见 Chunked 线性注意力 、标量衰减 、Gated DeltaNet 。
0. 符号与三处变化
沿用前三篇:β t \beta_t β t 写入强度、C C C 块长、[ t ] [t] [ t ] 块序号、r ∈ [ 1 , C ] r \in [1,C] r ∈ [ 1 , C ] 块内位置、S ⊺ ∈ R d k × d v \mathbf{S}^\intercal \in \mathbb{R}^{d_k \times d_v} S ⊺ ∈ R d k × d v 为 kernel 存储布局。
本篇的核心量改为累积 log 衰减 (沿用《来龙去脉》的记号):
γ i = ∑ s = 1 i g s ∈ R d k , g s = log a s < 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 = s = 1 ∑ i g s ∈ R d k , g s = log a s < 0 (逐分量)
注意 γ i \bm{\gamma}_i γ i 是向量 –这是与前两篇最本质的区别。前三篇的 γ r \gamma^r γ r 是标量,Γ i j = γ i / γ j \Gamma_{ij} = \gamma_i/\gamma_j Γ ij = γ i / γ j 是一张 C × C C \times C C × C 的表;本篇 γ i − γ j \bm{\gamma}_i - \bm{\gamma}_j γ i − γ j 是 d k d_k d k 维向量,"衰减矩阵"概念上是 C × C × d k C \times C \times d_k C × C × d k ,不可能物化 。
三处结构性变化:
GDN(标量门)
KDA(逐通道门)
累积积
γ r \gamma^r γ r 标量,cumsum 后 C C C 个数
γ r ∈ R d k \bm{\gamma}_r \in \mathbb{R}^{d_k} γ r ∈ R d k ,cumsum 后 C × d k C \times d_k C × d k
衰减掩码
Γ ∈ R C × C \Gamma \in \mathbb{R}^{C\times C} Γ ∈ R C × C ,可物化
概念上 C × C × d k C\times C\times d_k C × C × d k ,必须融进 GEMM
KK 矩阵
K K ⊺ ⊙ Γ \mathbf{K}\mathbf{K}^\intercal \odot \Gamma K K ⊺ ⊙ Γ ,衰减可外提
M c i = ( k c ⊙ e γ c − γ i ) ⊺ k i M_{ci} = (\bm{k}_c \odot e^{\bm{\gamma}_c - \bm{\gamma}_i})^\intercal\bm{k}_i M c i = ( k c ⊙ e γ c − γ i ) ⊺ k i ,衰减长在内部
第三行是全篇的技术核心。
1. 衰减长在内部:问题与出路
1.1 朴素做法的代价
M c i = ( k c ⊙ e γ c − γ i ) ⊺ k i M_{ci} = (\bm{k}_c \odot e^{\bm{\gamma}_c - \bm{\gamma}_i})^\intercal\bm{k}_i M c i = ( k c ⊙ e γ c − γ i ) ⊺ k i 里的指数依赖 ( c , i , d ) (c, i, d) ( c , i , d ) 三个下标。直接算就是 C 2 / 2 C^2/2 C 2 /2 次长度 d k d_k d k 的加权内积,每次都要现算 d k d_k d k 个 exp 。
这不是常数因子问题–它把一次 GEMM 变成了 C 2 / 2 C^2/2 C 2 /2 个独立的向量运算,完全用不上 Tensor Core。C = 64 C = 64 C = 64 、d k = 128 d_k = 128 d k = 128 时是 2016 次加权内积、约 26 万次 exp2。
1.2 出路:指数可分离
关键观察是指数可以按通道拆开 :
e γ c − γ i = e γ c ⊙ e − γ i e^{\bm{\gamma}_c - \bm{\gamma}_i} = e^{\bm{\gamma}_c} \odot e^{-\bm{\gamma}_i}
e γ c − γ i = e γ c ⊙ e − γ i
于是:
M c i = ∑ d ( k c [ d ] e γ c [ d ] ) ⏟ K ~ + [ c , d ] ( k i [ d ] e − γ i [ d ] ) ⏟ K ~ − [ i , d ] = ( K ~ + K ~ − ⊺ ) c i M_{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}
M c i = d ∑ K + [ c , d ] ( k c [ d ] e γ c [ d ] ) K − [ i , d ] ( k i [ d ] e − γ i [ d ] ) = ( K + K − ⊺ ) c i
一次 GEMM 解决 ,前置两个 C × d k C \times d_k C × d k 的逐元素加权。实测拆分与朴素计算的差异 1.94 × 10 − 16 1.94 \times 10^{-16} 1.94 × 1 0 − 16 ,代数上完全等价。
1.3 代价:e − γ i e^{-\bm{\gamma}_i} e − γ i 是大于 1 的量
第二篇讲过一个教训:把比值 γ i / γ j \gamma_i/\gamma_j γ i / γ j 因式分解成 γ i ⋅ ( 1 / γ j ) \gamma_i \cdot (1/\gamma_j) γ i ⋅ ( 1/ γ j ) 会物化一个指数增长的量,fp16 下 α ≤ 0.8 \alpha \le 0.8 α ≤ 0.8 即 NaN。这里是同一个陷阱的逐通道版本 –e − γ i [ d ] ≥ 1 e^{-\gamma_i[d]} \ge 1 e − γ i [ d ] ≥ 1 ,且随 i i i 与通道衰减强度指数增长。
但本篇的处境和第二篇不同:那里有替代方案(保留比值形式),这里没有 。逐通道衰减无法外提,不拆就用不上 Tensor Core。所以问题从"要不要拆"变成了"怎样让拆分在数值上安全 "。
实测 e − γ e^{-\bm{\gamma}} e − γ 的最大值(C = 64 C = 64 C = 64 、d k = 128 d_k = 128 d k = 128 ):
门控下界
max e − γ \max e^{-\bm{\gamma}} max e − γ
fp16
fp32
0.999
1.04 1.04 1.04
OK
OK
0.99
1.48 1.48 1.48
OK
OK
0.95
7.14 7.14 7.14
OK
OK
0.9
54.95 54.95 54.95
OK
OK
0.8
4.20 × 10 3 4.20 \times 10^{3} 4.20 × 1 0 3
OK
OK
0.5
3.04 × 10 10 3.04 \times 10^{10} 3.04 × 1 0 10
溢出
OK
门控下界在 0.9 以上时,e − γ e^{-\bm{\gamma}} e − γ 连 fp16 都装得下。 这就把 KDA 论文里"下界衰减(lower-bounded decay)"这个设计从模型层面的技巧,变成了 kernel 能否用 Tensor Core 的前提条件。
2. 下界衰减:不是精度调优,是可行性前提
2.1 无下界时 fp16 全线归零
e − γ e^{-\bm{\gamma}} e − γ 会溢出,另一头 e γ e^{\bm{\gamma}} e γ 会下溢。而逐通道门让后者严重得多–总有一些通道学到很小的 a a a ,它们的 γ \gamma γ 累积得最快。
实测 min e γ C \min e^{\bm{\gamma}_C} min e γ C 与 fp16 下归零的通道数(d k = 128 d_k = 128 d k = 128 ):
a \bm{a} a 采样区间
C C C
min e γ C \min e^{\bm{\gamma}_C} min e γ C
fp16 归零通道
[ 0.9 , 0.999 ] [0.9,\ 0.999] [ 0.9 , 0.999 ]
64
1.70 × 10 − 2 1.70\times10^{-2} 1.70 × 1 0 − 2
0 / 128
[ 0.9 , 0.999 ] [0.9,\ 0.999] [ 0.9 , 0.999 ]
128
4.84 × 10 − 4 4.84\times10^{-4} 4.84 × 1 0 − 4
0 / 128
[ 0.5 , 0.999 ] [0.5,\ 0.999] [ 0.5 , 0.999 ]
64
2.20 × 10 − 11 2.20\times10^{-11} 2.20 × 1 0 − 11
118 / 128
[ 0.5 , 0.999 ] [0.5,\ 0.999] [ 0.5 , 0.999 ]
128
1.59 × 10 − 20 1.59\times10^{-20} 1.59 × 1 0 − 20
128 / 128
[ 0.1 , 0.999 ] [0.1,\ 0.999] [ 0.1 , 0.999 ]
64
7.39 × 10 − 28 7.39\times10^{-28} 7.39 × 1 0 − 28
128 / 128
[ 0.01 , 0.999 ] [0.01,\ 0.999] [ 0.01 , 0.999 ]
128
6.44 × 10 − 65 6.44\times10^{-65} 6.44 × 1 0 − 65
128 / 128
C = 128 C = 128 C = 128 、下界 0.5 时全部 128 个通道在 fp16 下归零 –状态被彻底清空,kernel 输出恒为块内项,跨块信息完全丢失。
加下界后:
下界
min e γ C \min e^{\bm{\gamma}_C} min e γ C
fp16 归零
无
1.84 × 10 − 66 1.84\times10^{-66} 1.84 × 1 0 − 66
128 / 128
0.5
9.07 × 10 − 32 9.07\times10^{-32} 9.07 × 1 0 − 32
128 / 128
0.9
1.72 × 10 − 6 1.72\times10^{-6} 1.72 × 1 0 − 6
0 / 128
0.95
1.43 × 10 − 3 1.43\times10^{-3} 1.43 × 1 0 − 3
0 / 128
下界 0.9 是分界线。 这一个数字同时解决了两头:e γ e^{\bm{\gamma}} e γ 不下溢、e − γ e^{-\bm{\gamma}} e − γ 不溢出。
2.2 通道离散度:逐通道门特有的问题
标量门下 γ C \gamma^C γ C 是一个数;逐通道门下它是 d k d_k d k 个数,而它们的极差 决定了同一个 chunk 内不同通道的数值尺度差异:
a \bm{a} a 区间
C C C
γ C \bm{\gamma}_C γ C 极差(log 域)
e 极差 e^{\text{极差}} e 极差
[ 0.9 , 0.999 ] [0.9,\ 0.999] [ 0.9 , 0.999 ]
64
1.28
3.58 3.58 3.58
[ 0.9 , 0.999 ] [0.9,\ 0.999] [ 0.9 , 0.999 ]
128
1.74
5.69 5.69 5.69
[ 0.5 , 0.999 ] [0.5,\ 0.999] [ 0.5 , 0.999 ]
64
8.17
3.53 × 10 3 3.53\times10^{3} 3.53 × 1 0 3
[ 0.5 , 0.999 ] [0.5,\ 0.999] [ 0.5 , 0.999 ]
128
11.25
7.69 × 10 4 7.69\times10^{4} 7.69 × 1 0 4
[ 0.01 , 0.999 ] [0.01,\ 0.999] [ 0.01 , 0.999 ]
128
47.46
4.10 × 10 20 4.10\times10^{20} 4.10 × 1 0 20
极差 47 47 47 意味着同一个 fragment 里最强和最弱通道的数值相差 20 个数量级 –任何浮点格式都无法同时表示。这是逐通道门独有的病:标量门只需要担心整体的尺度漂移,逐通道门还要担心通道之间的尺度撕裂。
下界 0.9 把极差压到 1.74(e 1.74 = 5.7 e^{1.74} = 5.7 e 1.74 = 5.7 ),完全可控。
2.3 sub-chunk:第二道保险
即使有下界,C = 128 C = 128 C = 128 时 max e − γ = 1.79 × 10 3 \max e^{-\bm{\gamma}} = 1.79\times10^3 max e − γ = 1.79 × 1 0 3 (下界 0.9)–fp16 装得下但余量不多。把块内再切成 sub-chunk,每个 sub-chunk 内部重新起算 γ \bm{\gamma} γ :
划分
max e − γ \max e^{-\bm{\gamma}} max e − γ
C = 128 C = 128 C = 128 整块
1.79 × 10 3 1.79\times10^{3} 1.79 × 1 0 3
sub-chunk = 64 = 64 = 64
5.34 × 10 1 5.34\times10^{1} 5.34 × 1 0 1
sub-chunk = 32 = 32 = 32
9.02 9.02 9.02
sub-chunk = 16 = 16 = 16
3.41 3.41 3.41
指数跨度只取决于 sub-chunk 长度,与总块长无关。这与第二篇"累积积按 chunk 重置"是同一个道理 ,只是又下降了一层:chunk 重置控制 γ \bm{\gamma} γ 本身,sub-chunk 重置控制 e − γ e^{-\bm{\gamma}} e − γ 的动态范围。
代价是 sub-chunk 之间需要额外的状态传递,块内变成两层循环。C = 64 C = 64 C = 64 + 下界 0.9 时 max e − γ = 55 \max e^{-\bm{\gamma}} = 55 max e − γ = 55 ,不需要 sub-chunk ;C = 128 C = 128 C = 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] 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 ( a t ) \operatorname{Diag}(\bm{a}_t) Diag ( a t ) 的位置很关键:它夹在 S t − 1 \mathbf{S}_{t-1} S t − 1 与 Householder 之间。写成 Diag ( a ) S \operatorname{Diag}(\bm{a})\mathbf{S} Diag ( a ) 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 ) gC = gam[-1 ] 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 ) Tm = np.linalg.inv(np.eye(C) + L) Ah = bt[:, None ] * Tm W = Ah @ (Kc * np.exp(gam)) U = Ah @ Vc Vt = U - W @ S 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} M 与 A q k \mathbf{A}^{qk} A q k –慢,但与公式逐字对应,作为基准可信。§3.3 的 kernel 镜像才用 GEMM 化写法。
L c i = β i M c i L_{ci} = \beta_i M_{ci} L c i = β i M c i 与 A ^ = diag ( β ) T \hat{\mathbf{A}} = \operatorname{diag}(\beta)\mathbf{T} A ^ = diag ( β ) T 这一对下标必须配套。实测另一种等价写法是 L c i = β c M c i L_{ci} = \beta_c M_{ci} L c i = β c M c i 配 A ^ = T diag ( β ) \hat{\mathbf{A}} = \mathbf{T}\operatorname{diag}(\beta) A ^ = T diag ( β ) ,两者都对,混用则错。
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): dv = slice (bv * block_DV, (bv + 1 ) * block_DV) for bbh in range (B * H): b, h = bbh // H, bbh % H S = np.zeros((DK, block_DV)) for c in range (N // C): 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) 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 ) 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 不同的实现选择:
M \mathbf{M} M 与 A q k \mathbf{A}^{qk} A q k 走 GEMM :K ~ + K ~ − ⊺ \widetilde{\mathbf{K}}^{+}\widetilde{\mathbf{K}}^{-\intercal} K + K − ⊺ 。A q k \mathbf{A}^{qk} A q k 含对角(j ≤ c j \le c j ≤ c ),M \mathbf{M} M 不含(i < c i < c i < c )。
W \mathbf{W} W 与 U \mathbf{U} U 拼成一次前向替换 :两者共用同一个 ( I + L ) (\mathbf{I}+\mathbf{L}) ( I + L ) ,拼接后只解一遍,省一半串行开销。
β \beta β 必须在解完三角系统之后乘 。A ^ = diag ( β ) T \hat{\mathbf{A}} = \operatorname{diag}(\beta)\mathbf{T} A ^ = diag ( β ) T 展开是「先 T \mathbf{T} T 作用、再逐行乘 β \beta β 」;若把 β \beta β 提前乘进右端项,算的就是 T diag ( β ) \mathbf{T}\operatorname{diag}(\beta) T diag ( β ) –那是另一种配对(需搭配 L c i = β c M c i L_{ci} = \beta_c M_{ci} L c i = β c M c i ),混用则错。这个错误在 β \beta β 全部相等时完全看不出来 ,我第一次写就踩了:β \beta β 随机时相对 L2 达 1.3 × 10 − 1 1.3\times10^{-1} 1.3 × 1 0 − 1 。
状态更新用逐行乘代替 Diag \operatorname{Diag} Diag :Ep[-1][:, None] * S 就是 Diag ( e γ C ) S \operatorname{Diag}(e^{\bm{\gamma}_C})\mathbf{S} Diag ( e γ C ) S ,不物化对角矩阵。
3.4 一致性验证
B = 2 , H = 2 , N = 12 , d k = d v = 4 , C = 4 B=2, H=2, N=12, d_k=d_v=4, C=4 B = 2 , H = 2 , N = 12 , d k = d v = 4 , C = 4 ,a ∼ U ( 0.90 , 0.999 ) d k \bm{a} \sim \mathcal{U}(0.90, 0.999)^{d_k} a ∼ U ( 0.90 , 0.999 ) d k 、β ∼ U ( 0.1 , 0.9 ) \beta \sim \mathcal{U}(0.1, 0.9) β ∼ U ( 0.1 , 0.9 ) ,q , k \bm{q},\bm{k} q , k 已 L2 归一化,fp64:
比较
max abs 误差
相对 L2
B chunkwise vs A 逐 token
3.89 × 10 − 16 3.89 \times 10^{-16} 3.89 × 1 0 − 16
2.46 × 10 − 16 2.46 \times 10^{-16} 2.46 × 1 0 − 16
C kernel 镜像(GEMM 化)vs A
4.44 × 10 − 16 4.44 \times 10^{-16} 4.44 × 1 0 − 16
2.57 × 10 − 16 2.57 \times 10^{-16} 2.57 × 1 0 − 16
指数分离 GEMM vs 朴素加权内积
1.11 × 10 − 16 1.11 \times 10^{-16} 1.11 × 1 0 − 16
–
L = β i M L = \beta_i M L = β i M + diag ( β ) T \operatorname{diag}(\beta)\mathbf{T} diag ( β ) T
1.11 × 10 − 16 1.11 \times 10^{-16} 1.11 × 1 0 − 16
–
L = β c M L = \beta_c M L = β c M + T diag ( β ) \mathbf{T}\operatorname{diag}(\beta) T diag ( β )
1.11 × 10 − 16 1.11 \times 10^{-16} 1.11 × 1 0 − 16
–
三个退化检验 (本篇比 GDN 多一个):
退化
应回到
检验的部分
实测
g ≡ log α ⋅ 1 \bm{g} \equiv \log\alpha \cdot \bm{1} g ≡ log α ⋅ 1 (所有通道同值)
GDN
逐通道的指数分离
4.44 × 10 − 16 4.44\times10^{-16} 4.44 × 1 0 − 16
β → 0 \beta \to 0 β → 0
逐通道纯衰减
整个 UT 与三角系统
1.21 × 10 − 27 1.21\times10^{-27} 1.21 × 1 0 − 27
g → 0 \bm{g} \to \bm{0} g → 0 且 β \beta β 保留
纯 DeltaNet
所有衰减权重
4.44 × 10 − 16 4.44\times10^{-16} 4.44 × 1 0 − 16
三个 block D V \text{block}_{DV} block D V (1 / 2 / 4)的相对 L2 分别为 2.48 2.48 2.48 / 2.57 2.57 2.57 / 2.57 × 10 − 16 2.57 \times 10^{-16} 2.57 × 1 0 − 16 –切 DV 零依赖,与第一篇的结论一致。
第一个是本篇独有且最重要的:把逐通道门退化成标量门,必须精确回到第三篇的结果 。这一步能抓住"指数分离时把通道维和位置维搞混"这类错误–而那类错误在通道值本来就相同时会隐身。
4. TileLang kernel
4.0 寄存器账
d k = d v = 128 d_k = d_v = 128 d k = d v = 128 、C = 64 C = 64 C = 64 、block D V = 32 \text{block}_{DV} = 32 block D V = 32 、128 线程:
fragment
GDN
KDA
说明
S_f [ d k , bDV ] [d_k,\text{bDV}] [ d k , bDV ]
32
32
不变
acc_o [ C , bDV ] [C,\text{bDV}] [ C , bDV ]
16
16
不变
Gam [ C , C ] [C,C] [ C , C ]
32
0
逐通道无法物化,取消
A [ C , C ] [C,C] [ C , C ]
32
32
A q k \mathbf{A}^{qk} A q k
Amat [ C , C ] [C,C] [ C , C ]
32
32
三角系统系数 L \mathbf{L} L
Ep/Em [ C , d k ] [C,d_k] [ C , d k ]
–
2×64 = 128
新增:e ± γ e^{\pm\bm{\gamma}} e ± γ
W [ C , d k ] [C,d_k] [ C , d k ]
–
64
新增
Delta/Vt [ C , bDV ] [C,\text{bDV}] [ C , bDV ]
16
16
合计
≈ 160
≈ 320
上限 255
320 超限了。 逐通道门带来的 [ C , d k ] [C, d_k] [ C , d k ] 量级 fragment(e ± γ e^{\pm\bm{\gamma}} e ± γ 、W \mathbf{W} W )比 C 2 C^2 C 2 表更吃寄存器–C × d k = 64 × 128 C \times d_k = 64\times128 C × d k = 64 × 128 是 C 2 = 64 2 C^2 = 64^2 C 2 = 6 4 2 的两倍。
三条出路:
e − γ e^{-\bm{\gamma}} e − γ 不常驻 :它只在构造 M \mathbf{M} M 、A q k \mathbf{A}^{qk} A q k 时用,用完即弃,可以放 shared memory。省 64 reg。
W \mathbf{W} W 走 shared :它是 U − W S \mathbf{U} - \mathbf{W}\mathbf{S} U − WS 的中间量,不参与后续 GEMM 的累加器。省 64 reg。
C C C 降到 32 :所有 C C C 相关量减半。
组合 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 ), 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) Ep_s = T.alloc_shared([C, DK], accum_dtype) Em_s = T.alloc_shared([C, DK], accum_dtype) Kp_s = T.alloc_shared([C, DK], dtype) Km_s = T.alloc_shared([C, DK], dtype) 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) 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) be_f = T.alloc_fragment([C], accum_dtype)
累积衰减是逐通道的前缀和 –d k d_k d 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) 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 ) Em_s[r, d] = T.exp2(-gam[r, d] * 1.4426950408889634 ) 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:那里前缀和是 C C C 步串行、每步 1 个数;这里是 C C C 步串行、每步 d k d_k d k 个通道并行 。串行长度不变,但并行度从 1 涨到 128–逐通道门在这一处反而更适合 GPU。
1.4426950408889634 是 log 2 e \log_2 e log 2 e ,把 e x e^x e x 转成硬件 exp2(与第三篇一致)。
4.2 指数分离:两次 GEMM 取代双重循环
1 2 3 4 5 6 7 8 9 10 11 12 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 ) for r, d in T.Parallel(C, DK): Q_s[r, d] = T.Cast(dtype, Q_s[r, d] * Ep_s[r, d]) 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。 朴素写法需要 C 2 / 2 C^2/2 C 2 /2 次长度 d k d_k d k 的加权内积(C = 64 , d k = 128 C=64, d_k=128 C = 64 , d k = 128 时约 26 万次 exp);分离后 e ± γ e^{\pm\bm{\gamma}} e ± γ 各算一次(C × d k = 8192 C \times d_k = 8192 C × d k = 8192 次 exp2),剩下交给 Tensor Core。
注意两个掩码的边界不同:Lmat 取 j < i j < i j < i (严格下三角,对角会重复计入 β r \beta_r β r ),Aqk 取 j ≤ i j \le i j ≤ 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 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] 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] 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] T.copy(S_f, S_s) T.gemm(W_s, S_s, Vt, clear_accum=True ) 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} W 与 U \mathbf{U} U 共用同一个 ( I + L ) (\mathbf{I}+\mathbf{L}) ( I + L ) ,拼在一次前向替换里只付一遍串行代价 。这一步的并行度是 d k + block D V = 160 d_k + \text{block}_{DV} = 160 d k + block D V = 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 ) T.copy(Aqk, Aq_c) T.gemm(Aq_c, Vt_s, acc_o) T.copy(acc_o, O_s) T.copy(O_s, O[bb, s0:s0+C, bh, dv0:dv0+block_DV]) 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]) 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],每个 d k d_k d k 行乘各自的衰减
状态的不同行按不同速率淡出–这就是"逐通道"在 kernel 层面的全部含义。
e^{\bm{\gamma}_C - \bm{\gamma}_r} 用 Ep_s[C-1,d] * Em_s[r,d] 算,即 e γ C ⋅ e − γ r e^{\gamma_C}\cdot e^{-\gamma_r} e γ C ⋅ e − γ r 。这里复用了已有的两张表,不必重算 exp2;但它也是 e − γ e^{-\bm{\gamma}} e − γ 参与的第三处,进一步说明为什么下界不可省。
4.5 七处次序约束
比 GDN 多两处:
状态更新排在输出写回之后 (右移语义,四篇一致)。
S_f *= Ep_s[C-1,:] 在 T.gemm 之前 。
Q_s 被原地乘 e γ e^{\bm{\gamma}} e γ 后不可再用于块内项 –但本篇的块内项恰好也用 q ⊙ e γ \bm{q}\odot e^{\bm{\gamma}} q ⊙ e γ (A q k \mathbf{A}^{qk} A q k 的定义里就带),所以不需要重载 Q 。这是与 GDN 的一处反差,容易照抄出错。
Lmat 只取严格下三角 (j < i j < i j < i ),Aqk 含对角(j ≤ i j \le i j ≤ i )。
K_s 必须保持原始值 :K p \mathbf{Kp} Kp 、K m \mathbf{Km} Km 、右端项、状态更新四处都从 K_s 派生,任何一处原地修改都会污染后续。本篇的做法是始终写入独立的 Kp_s/Km_s,K_s 只读 –比 GDN 的"污染两次再重载"更清晰。
β \beta β 在前向替换之后乘,不能提前混进右端项 。A ^ = diag ( β ) T \hat{\mathbf{A}} = \operatorname{diag}(\beta)\mathbf{T} A ^ = diag ( β ) T 与 T diag ( β ) \mathbf{T}\operatorname{diag}(\beta) T diag ( β ) 是两种不同配对,各自要搭配 L c i = β i M c i L_{ci} = \beta_i M_{ci} L c i = β i M c i 或 β c M c i \beta_c M_{ci} β c M c i 。β \beta β 全相等时这个错误完全隐身。
T.clear(S_f) 在循环外,clear_accum=True 在循环内 。
第 3、5、6 条都是"照抄上一篇会错"的地方,其中第 6 条我实际踩过。
5. 四篇对照
线性注意力
标量衰减
GDN
KDA
衰减
无
α \alpha α 常数
α t \alpha_t α t 标量
a t ∈ R d k \bm{a}_t \in \mathbb{R}^{d_k} a t ∈ R d k
删除
无
无
I − β k k ⊺ \mathbf{I}-\beta\bm{k}\bm{k}^\intercal I − β k k ⊺
同 GDN
衰减掩码
M \mathbf{M} M (0/1)
Γ \Gamma Γ (C 2 C^2 C 2 表)
Γ \Gamma Γ (C 2 C^2 C 2 表)
无表,融进 GEMM
累积积
–
编译期常量
C C C 步串行,宽度 1
C C C 步串行,宽度 d k d_k d k
前向替换并行度
–
–
bDV = 32 \text{bDV} = 32 bDV = 32
d k + bDV = 160 d_k + \text{bDV} = 160 d k + bDV = 160
主要 fragment
C 2 C^2 C 2
2 C 2 2C^2 2 C 2
3 C 2 3C^2 3 C 2
2 C 2 + 3 C d k 2C^2 + 3Cd_k 2 C 2 + 3 C d k
可用 C C C
128
128
64
64(需 shared 卸载)
数值命门
无
1 / γ 1/\gamma 1/ γ 溢出
三角系统条件数
e ± γ e^{\pm\bm{\gamma}} e ± γ 双向 + 通道离散
关键前提
–
累积积按 chunk 重置
k \bm{k} k L2 归一化
门控下界 ≥ 0.9
四篇的箭头量语义完全一致 :q ← \overleftarrow{\bm{q}} q 乘 e γ r e^{\bm{\gamma}_r} e γ r 、k → \overrightarrow{\bm{k}} k 乘 e γ C − γ r e^{\bm{\gamma}_C-\bm{\gamma}_r} e γ C − γ r 、S → \overrightarrow{\mathbf{S}} S 乘 e γ C e^{\bm{\gamma}_C} e γ C 。从第一篇的"无衰减"到本篇的"逐通道",衰减插入的三个位置从未改变 ,改变的只是每个位置乘的是标量还是向量。
6. 总结
逐通道门只改一个符号,却拆掉了前三篇所有便利 。α t → Diag ( a t ) \alpha_t \to \operatorname{Diag}(\bm{a}_t) α t → Diag ( a t ) 让累积积从标量变向量,Γ \Gamma Γ 从 C × C C\times C C × C 表变成概念上的 C × C × d k C\times C\times d_k C × C × d k –不可能物化 ,必须融进 GEMM。
衰减长在 KK 矩阵内部,这是 KDA chunkwise 的核心难点 。M c i = ( k c ⊙ e γ c − γ i ) ⊺ k i M_{ci} = (\bm{k}_c \odot e^{\bm{\gamma}_c-\bm{\gamma}_i})^\intercal\bm{k}_i M c i = ( k c ⊙ e γ c − γ i ) ⊺ k i 的指数依赖三个下标,无法像标量门那样外提成逐元素乘。
出路是指数可分离 :e γ c − γ i = e γ c ⊙ e − γ i e^{\bm{\gamma}_c-\bm{\gamma}_i} = e^{\bm{\gamma}_c}\odot e^{-\bm{\gamma}_i} e γ c − γ i = e γ c ⊙ e − γ i ,于是 M = K ~ + K ~ − ⊺ \mathbf{M} = \widetilde{\mathbf{K}}^{+}\widetilde{\mathbf{K}}^{-\intercal} M = K + K − ⊺ 一次 GEMM 解决(实测等价,误差 1.94 × 10 − 16 1.94\times10^{-16} 1.94 × 1 0 − 16 )。朴素写法要 C 2 / 2 C^2/2 C 2 /2 次加权内积、约 26 万次 exp;分离后只需 C d k = 8192 C d_k = 8192 C d k = 8192 次 exp2 加两次 GEMM。
代价是必须物化 e − γ e^{-\bm{\gamma}} e − γ –第二篇批判过的 1 / γ 1/\gamma 1/ γ 陷阱的逐通道版本。但这次没有替代方案:不拆就用不上 Tensor Core。问题从"要不要拆"变成"如何让拆分数值安全"。
门控下界是 kernel 可行性的前提,不是精度调优 。实测 C = 128 C=128 C = 128 、无下界时 fp16 下 e γ C e^{\bm{\gamma}_C} e γ C 128/128 通道全部归零 ,状态被彻底清空。下界 0.9 时归零通道数为 0,同时 max e − γ = 55 \max e^{-\bm{\gamma}} = 55 max e − γ = 55 不溢出–一个数字同时管住了下溢与溢出两头 。
通道离散度是逐通道门独有的病 。a ∈ [ 0.01 , 0.999 ] \bm{a}\in[0.01,0.999] a ∈ [ 0.01 , 0.999 ] 、C = 128 C=128 C = 128 时 γ C \bm{\gamma}_C γ C 的通道极差达 47(log 域),即同一 fragment 内最强与最弱通道相差 20 个数量级,任何浮点格式都无法同时表示。下界 0.9 把极差压到 1.74。
sub-chunk 是第二道保险 。e − γ e^{-\bm{\gamma}} e − γ 的动态范围只取决于 sub-chunk 长度:C = 128 C=128 C = 128 整块 1.79 × 10 3 1.79\times10^3 1.79 × 1 0 3 ,切 32 后降到 9.02。这与第二篇"累积积按 chunk 重置"同理,只是又降一层。C = 64 C=64 C = 64 + 下界 0.9 时不需要。
逐通道门在两处反而更适合 GPU 。累积前缀和:GDN 是 C C C 步串行 × 宽度 1,KDA 是 C C C 步串行 × 宽度 d k d_k d k ,串行长度不变而并行度从 1 涨到 128。前向替换:W \mathbf{W} W 与 U \mathbf{U} U 共用 ( I + L ) (\mathbf{I}+\mathbf{L}) ( I + L ) 拼成一次求解,并行度 d k + bDV = 160 d_k + \text{bDV} = 160 d k + bDV = 160 ,远高于 GDN 的 32。
寄存器压力换了主角 。GDN 的瓶颈是三张 C 2 C^2 C 2 表;KDA 是三张 C × d k C \times d_k C × d k (e ± γ e^{\pm\bm{\gamma}} e ± γ 、W \mathbf{W} W ),64 × 128 64\times128 64 × 128 是 64 2 64^2 6 4 2 的两倍。朴素分配约 320 reg 超限,把 e − γ e^{-\bm{\gamma}} e − γ 与 W \mathbf{W} W 卸载到 shared memory 后降到约 192。
三处"照抄上一篇会错" :块内项的 A q k \mathbf{A}^{qk} A q k 定义里本就带 q ⊙ e γ \bm{q}\odot e^{\bm{\gamma}} q ⊙ e γ ,不需要像 GDN 那样重载 Q ;K_s 应始终只读、派生量写独立 buffer,而非 GDN 的"污染再重载";β \beta β 必须在解完三角系统之后乘 –diag ( β ) T \operatorname{diag}(\beta)\mathbf{T} diag ( β ) T 与 T diag ( β ) \mathbf{T}\operatorname{diag}(\beta) T diag ( β ) 需搭配不同的 L L L 下标,写反时 β \beta β 全相等则隐身、β \beta β 随机则相对 L2 达 1.3 × 10 − 1 1.3\times10^{-1} 1.3 × 1 0 − 1 (我第一次写就踩了这个)。
三个退化检验,比 GDN 多一个 。g \bm{g} g 各通道同值 → 回到 GDN(检验指数分离是否搞混通道维与位置维)、β → 0 \beta\to0 β → 0 → 逐通道纯衰减、g → 0 \bm{g}\to\bm{0} g → 0 → 纯 DeltaNet。第一个最重要且本篇独有。
衰减的三处落点四篇未变 。q ← \overleftarrow{\bm{q}} q 、k → \overrightarrow{\bm{k}} k 、S → \overrightarrow{\mathbf{S}} S 从第一篇到第四篇完全一致,变的只是每处乘标量还是乘向量。这是 GDN 那套箭头记号真正的价值–它把"衰减"隔离成了一个独立于门控形态的层面。
可迁移的启示 :把一个标量参数升级成向量,代价从来不在参数量。它会连带改变哪些量能被预计算、哪些表能被物化、哪些运算能进 Tensor Core 。KDA 的 α t → Diag ( a t ) \alpha_t \to \operatorname{Diag}(\bm{a}_t) α t → Diag ( a t ) 只多了 d k d_k d k 个数,却让衰减掩码从"一张可复用的表"变成"必须融进 GEMM 的隐式结构",并且逼出了一个模型层面的约束(门控下界)作为 kernel 可行性的前提。当一个数值技巧成为架构设计的必要条件时,它就不再是实现细节。
参考