TileLang 实战:KDA 从零到一–标量衰减
上一篇实现了不带任何遗忘机制的 chunked 线性注意力,状态单调累加。本篇引入第一个衰减因子:把递推式改为 S ← γ S + K ⊤ V S \leftarrow \gamma S + K^\top V S ← γ S + K ⊤ V ,γ \gamma γ 是一个标量常数。
改动看起来只是多乘一个数,但它决定了后续三级的全部数值策略。本文先引入 GDN 论文用来统一递归形式与并行形式的累积衰减积 Λ j = ∏ i ≤ j α i \Lambda_j = \prod_{i \le j} \alpha_i Λ j = ∏ i ≤ j α i ,它把任意区间的衰减化归为两个前缀量之比。论文给出的矩阵并行形式是 O = ( ( Q K ⊤ ) ⊙ Γ ) V O = ((QK^\top) \odot \Gamma)V O = (( Q K ⊤ ) ⊙ Γ ) V ,其中 Γ i j = Λ i / Λ j \Gamma_{ij} = \Lambda_i/\Lambda_j Γ ij = Λ i / Λ j –这是一个比值,因果约束下恒不大于 1,数值安全 。但它存在一个看似有利的因式分解 Λ i / Λ j = Λ i ⋅ ( 1 / Λ j ) \Lambda_i/\Lambda_j = \Lambda_i \cdot (1/\Lambda_j) Λ i / Λ j = Λ i ⋅ ( 1/ Λ j ) ,能把两步合成单次 GEMM 并省去 B C 2 BC^2 B C 2 权重表。本文实测给出结论:这个分解会把中间量推到 1 / Λ B C 1/\Lambda_{BC} 1/ Λ B C ,γ ≤ 0.8 \gamma \le 0.8 γ ≤ 0.8 时 fp16 下直接产生 NaN 。论文的比值形式必须原样保留。这个约束在第三级 Λ \Lambda Λ 从标量升级为逐通道向量后会变成 KDA 实现中最棘手的问题。
上一篇见《TileLang 实战:KDA 从零到一–Chunked 线性注意力 》,本文沿用其三层参考实现的验证框架与符号约定。
1. 递推式与分块重写
本级递推式在上一级基础上给状态加一个遗忘系数 γ ∈ ( 0 , 1 ] \gamma \in (0, 1] γ ∈ ( 0 , 1 ] :
S t = γ S t − 1 + k t v t ⊤ , o t = q t ⊤ S t S_t = \gamma S_{t-1} + k_t v_t^\top, \qquad o_t = q_t^\top S_t
S t = γ S t − 1 + k t v t ⊤ , o t = q t ⊤ S t
展开成显式求和,注意每个 k j v j ⊤ k_j v_j^\top k j v j ⊤ 被后续每一步各乘一次 γ \gamma γ ,从 j j j 到 t t t 共乘 t − j t - j t − j 次:
S t = ∑ j ≤ t γ t − j k j v j ⊤ , o i = ∑ j ≤ i γ i − j ( q i ⋅ k j ) v j S_t = \sum_{j \le t} \gamma^{\,t-j} k_j v_j^\top, \qquad
o_i = \sum_{j \le i} \gamma^{\,i-j} (q_i \cdot k_j)\, v_j
S t = j ≤ t ∑ γ t − j k j v j ⊤ , o i = j ≤ i ∑ γ i − j ( q i ⋅ k j ) v j
对比上一级的 o i = ∑ j ≤ i ( q i ⋅ k j ) v j o_i = \sum_{j \le i} (q_i \cdot k_j) v_j o i = ∑ j ≤ i ( q i ⋅ k j ) v j ,唯一变化是每一项多了权重 γ i − j \gamma^{i-j} γ i − j –距离越远权重越小,这就是「衰减」的含义。γ = 1 \gamma = 1 γ = 1 时退化为上一级。
1.1 三处需要插入衰减权重的位置
沿用上一级的分块框架,序列按 B C BC B C 切分为 N C NC N C 个 chunk。把 j ≤ i j \le i j ≤ i 的求和拆成跨块与块内两部分,衰减权重会分别落到三个位置。设 token i i i 在其所属 chunk 内的局部下标为 i ′ = i m o d B C i' = i \bmod BC i ′ = i mod B C :
位置一–块内下三角 。同块内 j ′ ≤ i ′ j' \le i' j ′ ≤ i ′ ,相对距离就是局部下标之差:
A i ′ j ′ intra = { γ i ′ − j ′ j ′ ≤ i ′ 0 j ′ > i ′ A^{\text{intra}}_{i'j'} = \begin{cases} \gamma^{\,i'-j'} & j' \le i' \\ 0 & j' > i' \end{cases}
A i ′ j ′ intra = { γ i ′ − j ′ 0 j ′ ≤ i ′ j ′ > i ′
上一级这里是 0/1 掩码,本级变成指数下三角。注意权重全部落在 ( 0 , 1 ] (0, 1] ( 0 , 1 ] 区间 –对角线是 γ 0 = 1 \gamma^0 = 1 γ 0 = 1 ,左下角最小值是 γ B C − 1 \gamma^{BC-1} γ B C − 1 。
位置二–每块写入状态时的块尾对齐 。跨块状态需要定义一个统一的时间基准,取所属 chunk 的末尾 。chunk c c c 内第 m ′ m' m ′ 个 token 的贡献衰减到该块末尾要乘 γ B C − 1 − m ′ \gamma^{BC-1-m'} γ B C − 1 − m ′ :
Δ c = ∑ m ′ = 0 B C − 1 γ B C − 1 − m ′ k m ′ v m ′ ⊤ = ( K c ⊙ γ B C − 1 − m ′ ) ⊤ V c \Delta_c = \sum_{m'=0}^{BC-1} \gamma^{\,BC-1-m'} k_{m'} v_{m'}^\top = (K_c \odot \gamma^{\,BC-1-m'})^\top V_c
Δ c = m ′ = 0 ∑ B C − 1 γ B C − 1 − m ′ k m ′ v m ′ ⊤ = ( K c ⊙ γ B C − 1 − m ′ ) ⊤ V c
位置三–跨块状态递推与 query 侧缩放 。相邻块之间隔了整块 B C BC B C 步,因此前缀状态的递推是:
S c prev = γ B C S c − 1 prev + Δ c − 1 S_c^{\text{prev}} = \gamma^{BC} S_{c-1}^{\text{prev}} + \Delta_{c-1}
S c prev = γ B C S c − 1 prev + Δ c − 1
而 token i ′ i' i ′ 读取这个状态时,它距离上一块末尾还有 i ′ + 1 i' + 1 i ′ + 1 步,所以 query 要乘 γ i ′ + 1 \gamma^{\,i'+1} γ i ′ + 1 :
O c = ( Q c ⊙ γ i ′ + 1 ) S c prev + A intra V c O_c = (Q_c \odot \gamma^{\,i'+1}) S_c^{\text{prev}} + A^{\text{intra}} V_c
O c = ( Q c ⊙ γ i ′ + 1 ) S c prev + A intra V c
三个位置的权重汇总:
位置
权重
取值范围
作用对象
块内下三角
γ i ′ − j ′ \gamma^{\,i'-j'} γ i ′ − j ′
[ γ B C − 1 , 1 ] [\gamma^{BC-1},\ 1] [ γ B C − 1 , 1 ]
B C × B C BC \times BC B C × B C 分数矩阵,逐元素
写入状态
γ B C − 1 − m ′ \gamma^{\,BC-1-m'} γ B C − 1 − m ′
[ 1 , γ B C − 1 ] [1,\ \gamma^{BC-1}] [ 1 , γ B C − 1 ]
K c K_c K c 的行,逐行缩放
跨块递推
γ B C \gamma^{BC} γ B C
标量
整个状态矩阵
读取状态
γ i ′ + 1 \gamma^{\,i'+1} γ i ′ + 1
[ γ , γ B C ] [\gamma,\ \gamma^{BC}] [ γ , γ B C ]
Q c Q_c Q c 的行,逐行缩放
四个权重全部 ≤ 1 \le 1 ≤ 1 ,这一点是本级数值安全的根本原因,也是下一节讨论的分歧点。
一句话总结 :标量衰减不改变分块恒等式的结构,只是在三个位置插入指数权重;GEMM 的形状与调用次序与上一级完全一致。
2. 累积衰减积:从连乘到矩阵并行形式
上一节的推导是直接展开求和得到的,但衰减机制有一个更本质的表述方式,GDN 论文用它统一了递归形式与并行形式。理解这个表述是本级乃至后续三级的关键。
2.1 累积衰减积的定义
把标量衰减推广为依赖数据的 α t ∈ ( 0 , 1 ) \alpha_t \in (0, 1) α t ∈ ( 0 , 1 ) –每个时刻的遗忘强度由输入决定,而非固定常数(本级的常数 γ \gamma γ 是 α t ≡ γ \alpha_t \equiv \gamma α t ≡ γ 的特例)。定义累积衰减积 :
Λ j = ∏ i = 1 j α i \Lambda_j = \prod_{i=1}^{j} \alpha_i
Λ j = i = 1 ∏ j α i
Λ j \Lambda_j Λ j 的含义是从序列起点衰减到第 j j j 步的总折扣。有了它,递推式的展开可以写得非常紧凑。展开 S t = α t S t − 1 + k t v t ⊤ S_t = \alpha_t S_{t-1} + k_t v_t^\top S t = α t S t − 1 + k t v t ⊤ :
S t = ∑ j ≤ t ( ∏ i = j + 1 t α i ) k j v j ⊤ = ∑ j ≤ t Λ t Λ j k j v j ⊤ S_t = \sum_{j \le t} \Big( \prod_{i=j+1}^{t} \alpha_i \Big) k_j v_j^\top
= \sum_{j \le t} \frac{\Lambda_t}{\Lambda_j} k_j v_j^\top
S t = j ≤ t ∑ ( i = j + 1 ∏ t α i ) k j v j ⊤ = j ≤ t ∑ Λ j Λ t k j v j ⊤
中间那个连乘 ∏ i = j + 1 t α i \prod_{i=j+1}^{t} \alpha_i ∏ i = j + 1 t α i 正好是两个累积积的比值 Λ t / Λ j \Lambda_t / \Lambda_j Λ t / Λ j –这是累积积定义的全部价值:把「从 j j j 到 t t t 的区间连乘」化归为「两个前缀量之比」 ,于是任意区间的衰减都可以由一个长度 N N N 的前缀数组 O ( 1 ) O(1) O ( 1 ) 查得,不必对每个 ( j , t ) (j, t) ( j , t ) 对重新连乘。
2.2 两种等价形式
代入 o t = q t ⊤ S t o_t = q_t^\top S_t o t = q t ⊤ S t ,同一个结果可以写成两种形式:
向量形式(vector form) –逐时刻递归,对应推理阶段:
S t = α t S t − 1 + k t v t ⊤ , o t = q t ⊤ S t S_t = \alpha_t S_{t-1} + k_t v_t^\top, \qquad o_t = q_t^\top S_t
S t = α t S t − 1 + k t v t ⊤ , o t = q t ⊤ S t
矩阵并行形式(matrix parallel form) –整块一次算出,对应训练与 prefill:
O = ( ( Q K ⊤ ) ⊙ Γ ) V , Γ i j = { Λ i Λ j i ≥ j 0 i < j O = \big( (Q K^\top) \odot \Gamma \big) V, \qquad
\Gamma_{ij} = \begin{cases} \dfrac{\Lambda_i}{\Lambda_j} & i \ge j \\[2mm] 0 & i < j \end{cases}
O = ( ( Q K ⊤ ) ⊙ Γ ) V , Γ ij = ⎩ ⎨ ⎧ Λ j Λ i 0 i ≥ j i < j
Γ \Gamma Γ 是一个衰减感矩阵(decay-aware causal mask) –把上一级的 0/1 因果掩码换成了衰减比值。验证第 ( i , j ) (i,j) ( i , j ) 元素:
[ ( Q K ⊤ ) ⊙ Γ ] i j = Λ i Λ j ( q i ⋅ k j ) ( j ≤ i ) \big[(Q K^\top) \odot \Gamma\big]_{ij} = \frac{\Lambda_i}{\Lambda_j} (q_i \cdot k_j) \quad (j \le i)
[ ( Q K ⊤ ) ⊙ Γ ] ij = Λ j Λ i ( q i ⋅ k j ) ( j ≤ i )
与 §2.1 展开式逐项一致。这个形式的价值是把逐 token 的递归变成一次稠密 GEMM 加一次逐元素乘 ,完全并行,这正是「parallel within each chunk」的含义。递归形式与并行形式的这种等价在 Mamba2 中被称为状态空间对偶性(state space duality, SSD) 。
2.3 衰减矩阵的数值范围:为何比值形式是安全的
这里有一个容易看错的关键细节:Γ i j = Λ i / Λ j \Gamma_{ij} = \Lambda_i / \Lambda_j Γ ij = Λ i / Λ j 是一个比值 ,而且因为因果约束 i ≥ j i \ge j i ≥ j 、Λ \Lambda Λ 单调递减,所以:
0 < Γ i j = Λ i Λ j = ∏ r = j + 1 i α r ≤ 1 0 < \Gamma_{ij} = \frac{\Lambda_i}{\Lambda_j} = \prod_{r=j+1}^{i} \alpha_r \le 1
0 < Γ ij = Λ j Λ i = r = j + 1 ∏ i α r ≤ 1
论文的形式是数值安全的 –它先算 Q K ⊤ Q K^\top Q K ⊤ 再逐元素乘 Γ \Gamma Γ ,全程不需要物化任何大于 1 的量。实测确认(B C = 8 BC = 8 B C = 8 ,α t ∼ U ( 0.85 , 0.999 ) \alpha_t \sim \mathcal{U}(0.85,\ 0.999) α t ∼ U ( 0.85 , 0.999 ) ,fp64):
验证项
结果
论文式 (1) vs 逐 token 递归
max abs 误差 1.78 × 10 − 15 1.78 \times 10^{-15} 1.78 × 1 0 − 15
状态更新 S → \overrightarrow{S} S vs 递归末态
max abs 误差 8.88 × 10 − 16 8.88 \times 10^{-16} 8.88 × 1 0 − 16
Γ i j \Gamma_{ij} Γ ij 取值范围(i ≥ j i \ge j i ≥ j )
[ 0.628 , 1.000 ] [0.628,\ 1.000] [ 0.628 , 1.000 ] ,全部 ≤ 1 \le 1 ≤ 1
论文还给了一套简洁的箭头记号来表达这三个方向的衰减,与 §1.1 推导的三个位置一一对应:
论文记号
定义
含义
对应 §1.1
q r ← = Λ r q r \overleftarrow{q^r} = \Lambda_r\, q^r q r = Λ r q r
衰减到 chunk 首 位置
query 侧缩放
位置三(读取状态)
k r → = Λ C Λ r k r \overrightarrow{k^r} = \dfrac{\Lambda_C}{\Lambda_r} k^r k r = Λ r Λ C k r
衰减到 chunk 末 位置
key 侧缩放
位置二(块尾对齐)
S → = Λ C S \overrightarrow{S} = \Lambda_C\, S S = Λ C S
整块衰减
状态递推
位置三(跳块递推)
注意 k r → \overrightarrow{k^r} k r 里的 Λ C / Λ r \Lambda_C / \Lambda_r Λ C / Λ r 同样是比值且 ≤ 1 \le 1 ≤ 1 (因为 r ≤ C r \le C r ≤ C )。论文从头到尾没有单独物化过 1 / Λ j 1/\Lambda_j 1/ Λ j ,所有衰减因子都以比值形式出现。这是一个值得学习的记号设计–箭头方向直接编码了「衰减到哪个基准点」,而基准点选得当(首或末)就能保证指数非正。
常数 γ \gamma γ 是它的特例:Λ j = γ j \Lambda_j = \gamma^j Λ j = γ j ,于是 Γ i j = γ i − j \Gamma_{ij} = \gamma^{i-j} Γ ij = γ i − j ,回到 §1.1 的位置一。实测验证(B C = 4 BC = 4 B C = 4 ,fp64):Λ i / Λ j \Lambda_i/\Lambda_j Λ i / Λ j 与 γ i − j \gamma^{i-j} γ i − j 的最大偏差 2.22 × 10 − 16 2.22 \times 10^{-16} 2.22 × 1 0 − 16 ,即机器精度。
两种形式的 fp64 数值一致性(data-dependent α t ∼ U ( 0.85 , 0.999 ) \alpha_t \sim \mathcal{U}(0.85,\ 0.999) α t ∼ U ( 0.85 , 0.999 ) ,B = 2 , H = 2 , N = 12 , D = 4 , B C = 4 B=2, H=2, N=12, D=4, BC=4 B = 2 , H = 2 , N = 12 , D = 4 , B C = 4 ):
比较
max abs 误差
相对 L2
矩阵并行形式 vs 向量形式(递归)
3.55 × 10 − 15 3.55 \times 10^{-15} 3.55 × 1 0 − 15
2.03 × 10 − 16 2.03 \times 10^{-16} 2.03 × 1 0 − 16
不外提形式 vs 向量形式(递归)
3.55 × 10 − 15 3.55 \times 10^{-15} 3.55 × 1 0 − 15
1.89 × 10 − 16 1.89 \times 10^{-16} 1.89 × 1 0 − 16
常数 α t = 0.9 \alpha_t = 0.9 α t = 0.9 退化检验
2.67 × 10 − 15 2.67 \times 10^{-15} 2.67 × 1 0 − 15
1.83 × 10 − 16 1.83 \times 10^{-16} 1.83 × 1 0 − 16
2.4 一个容易走错的变形:把比值拆成乘积
既然论文形式是安全的,为何还要讨论数值问题?因为 Γ i j = Λ i / Λ j \Gamma_{ij} = \Lambda_i/\Lambda_j Γ ij = Λ i / Λ j 存在一个看似有利的因式分解:
Λ i Λ j = Λ i ⋅ 1 Λ j ⟹ ( Q K ⊤ ) ⊙ Γ = Tril [ ( Q ⊙ Λ ) ( K ⊘ Λ ) ⊤ ] \frac{\Lambda_i}{\Lambda_j} = \Lambda_i \cdot \frac{1}{\Lambda_j}
\quad\Longrightarrow\quad
(Q K^\top) \odot \Gamma = \operatorname{Tril}\big[ (Q \odot \Lambda)(K \oslash \Lambda)^\top \big]
Λ j Λ i = Λ i ⋅ Λ j 1 ⟹ ( Q K ⊤ ) ⊙ Γ = Tril [ ( Q ⊙ Λ ) ( K ⊘ Λ ) ⊤ ]
左侧需要先做 B C × B C BC \times BC B C × B C 的 GEMM、再逐元素乘一张 B C 2 BC^2 B C 2 的权重表;右侧把衰减推到两侧的逐行缩放上,只需一次 GEMM、不需权重表。看起来是纯改进。
但右侧必须显式物化 1 / Λ j 1/\Lambda_j 1/ Λ j ,而这个量不小于 1 且随 j j j 指数增长。 左侧的 Γ i j \Gamma_{ij} Γ ij 恒在 ( 0 , 1 ] (0,1] ( 0 , 1 ] ,右侧的中间量却能到 1 / Λ B C 1/\Lambda_{BC} 1/ Λ B C 。两者在实数域完全相等,在有限精度下完全不同。
data-dependent α t \alpha_t α t 下 1 / Λ j 1/\Lambda_j 1/ Λ j 的实测范围(2000 组随机采样取最大值):
α t \alpha_t α t 采样区间
B C = 64 BC=64 B C = 64 最大 1 / Λ 1/\Lambda 1/Λ
B C = 128 BC=128 B C = 128 最大 1 / Λ 1/\Lambda 1/Λ
fp16 溢出
[ 0.99 , 0.999 ] [0.99,\ 0.999] [ 0.99 , 0.999 ]
1.55 1.55 1.55
2.23 2.23 2.23
否
[ 0.95 , 0.999 ] [0.95,\ 0.999] [ 0.95 , 0.999 ]
7.74 7.74 7.74
47.6 47.6 47.6
否
[ 0.90 , 0.990 ] [0.90,\ 0.990] [ 0.90 , 0.990 ]
84.4 84.4 84.4
4.18 × 10 3 4.18 \times 10^{3} 4.18 × 1 0 3
否
[ 0.80 , 0.950 ] [0.80,\ 0.950] [ 0.80 , 0.950 ]
2.66 × 10 4 2.66 \times 10^{4} 2.66 × 1 0 4
3.40 × 10 8 3.40 \times 10^{8} 3.40 × 1 0 8
B C = 128 BC=128 B C = 128 溢出
[ 0.50 , 0.900 ] [0.50,\ 0.900] [ 0.50 , 0.900 ]
1.59 × 10 12 1.59 \times 10^{12} 1.59 × 1 0 12
6.23 × 10 23 6.23 \times 10^{23} 6.23 × 1 0 23
均溢出
对照同一组 α \alpha α 下两种写法的中间量范围(B C = 64 BC = 64 B C = 64 ):
α t \alpha_t α t 区间
因式分解后 max ( 1 / Λ j ) \max(1/\Lambda_j) max ( 1/ Λ j )
论文形式 max Γ i j \max \Gamma_{ij} max Γ ij
论文形式最小非零
[ 0.99 , 0.999 ] [0.99,\ 0.999] [ 0.99 , 0.999 ]
1.49 1.49 1.49
1.000 1.000 1.000
6.75 × 10 − 1 6.75 \times 10^{-1} 6.75 × 1 0 − 1
[ 0.90 , 0.990 ] [0.90,\ 0.990] [ 0.90 , 0.990 ]
33.7 33.7 33.7
1.000 1.000 1.000
3.26 × 10 − 2 3.26 \times 10^{-2} 3.26 × 1 0 − 2
[ 0.80 , 0.950 ] [0.80,\ 0.950] [ 0.80 , 0.950 ]
9.34 × 10 3 9.34 \times 10^{3} 9.34 × 1 0 3
1.000 1.000 1.000
1.33 × 10 − 4 1.33 \times 10^{-4} 1.33 × 1 0 − 4
论文形式的中间量上界恒为 1,与 α \alpha α 的分布无关 ;因式分解后的上界随 α \alpha α 变小而指数恶化。§3 给出 fp16 下的实测失效点。
2.5 为什么累积积要在 log 域计算
即使不做外提,Λ j \Lambda_j Λ j 本身也不宜用 cumprod 直接计算–连乘会下溢。实测(fp32 / fp16 直接连乘对比 fp64 log 域 cumsum):
α t \alpha_t α t 区间
B C BC B C
fp32 cumprod
fp16 cumprod
fp64 log-cumsum
[ 0.95 , 0.999 ] [0.95,\ 0.999] [ 0.95 , 0.999 ]
512
1.93 × 10 − 6 1.93 \times 10^{-6} 1.93 × 1 0 − 6
2.03 × 10 − 6 2.03 \times 10^{-6} 2.03 × 1 0 − 6
1.93 × 10 − 6 1.93 \times 10^{-6} 1.93 × 1 0 − 6
[ 0.90 , 0.990 ] [0.90,\ 0.990] [ 0.90 , 0.990 ]
512
1.69 × 10 − 13 1.69 \times 10^{-13} 1.69 × 1 0 − 13
2.38 × 10 − 7 2.38 \times 10^{-7} 2.38 × 1 0 − 7
1.69 × 10 − 13 1.69 \times 10^{-13} 1.69 × 1 0 − 13
[ 0.50 , 0.900 ] [0.50,\ 0.900] [ 0.50 , 0.900 ]
128
1.58 × 10 − 21 1.58 \times 10^{-21} 1.58 × 1 0 − 21
0 \mathbf{0} 0
1.58 × 10 − 21 1.58 \times 10^{-21} 1.58 × 1 0 − 21
[ 0.50 , 0.900 ] [0.50,\ 0.900] [ 0.50 , 0.900 ]
512
1.40 × 10 − 45 1.40 \times 10^{-45} 1.40 × 1 0 − 45
0 \mathbf{0} 0
1.81 × 10 − 84 1.81 \times 10^{-84} 1.81 × 1 0 − 84
fp16 连乘在 B C = 128 BC = 128 B C = 128 、α ∈ [ 0.5 , 0.9 ] \alpha \in [0.5,\ 0.9] α ∈ [ 0.5 , 0.9 ] 时已完全下溢为 0;fp32 在 B C = 512 BC = 512 B C = 512 时也进入非正规数区间(1.40 × 10 − 45 1.40 \times 10^{-45} 1.40 × 1 0 − 45 已是 fp32 最小非正规数量级)。log 域做加法则不受影响–log Λ j = ∑ i ≤ j log α i \log \Lambda_j = \sum_{i \le j} \log \alpha_i log Λ j = ∑ i ≤ j log α i 是线性增长的负数,表示范围绰绰有余。
因此实现上的标准做法是:门控值全程以 log α t \log \alpha_t log α t 的形式存储,累积积用 cumsum 而非 cumprod,需要衰减因子时用 exp2 还原 。这解释了 §5.1 为什么采用 exp2(log2(γ)·(i-j)) 的写法–本级 γ \gamma γ 是常数,这么写只是为了用上硬件指令;下一级 g t = log α t g_t = \log \alpha_t g t = log α t 本身就是网络输出,log 域是它的原生形式,exp2 从优化手段变成结构必需。
一句话总结 :累积积 Λ j = ∏ i ≤ j α i \Lambda_j = \prod_{i \le j} \alpha_i Λ j = ∏ i ≤ j α i 把区间连乘化归为前缀量之比,这是矩阵并行形式成立的代数基础;但比值一旦被因式分解成 Λ i ⋅ ( 1 / Λ j ) \Lambda_i \cdot (1/\Lambda_j) Λ i ⋅ ( 1/ Λ j ) ,就必须物化 1 / Λ j 1/\Lambda_j 1/ Λ j 这个不小于 1 且指数增长的量。
3. 定量分析:因式分解在 fp16 下的失效点
§2.4 的两种写法在实数域完全等价,但落到有限精度上行为分开。本节把问题收窄到常数 γ \gamma γ 的情形做定量分析。
块内下三角 γ i ′ − j ′ \gamma^{i'-j'} γ i ′ − j ′ 存在一个看似有利的代数变形–指数可以拆开:
γ i ′ − j ′ = γ i ′ ⋅ γ − j ′ \gamma^{\,i'-j'} = \gamma^{\,i'} \cdot \gamma^{-j'}
γ i ′ − j ′ = γ i ′ ⋅ γ − j ′
于是块内那一项可以写成:
A intra = tril ( ( Q c ⊙ γ i ′ ) ( K c ⊙ γ − j ′ ) ⊤ ) A^{\text{intra}} = \operatorname{tril}\big( (Q_c \odot \gamma^{\,i'}) (K_c \odot \gamma^{-j'})^\top \big)
A intra = tril ( ( Q c ⊙ γ i ′ ) ( K c ⊙ γ − j ′ ) ⊤ )
两种实现方式的差别:
方案
做法
块内额外开销
权重数值范围
比值形式(论文)
先 GEMM 得分数,再逐元素乘 Γ \Gamma Γ
一次 B C × B C BC \times BC B C × B C 逐元素乘 + 一张 B C 2 BC^2 B C 2 权重表
( 0 , 1 ] (0, 1] ( 0 , 1 ]
因式分解
先分别缩放 Q Q Q 、K K K 的行,再单次 GEMM
两次 B C × D BC \times D B C × D 逐行缩放,无 B C 2 BC^2 B C 2 开销
γ − j ′ \gamma^{-j'} γ − j ′ 最大 γ − ( B C − 1 ) \gamma^{-(BC-1)} γ − ( B C − 1 )
因式分解在 FLOPs 与寄存器占用上都更优:省掉一张 B C × B C BC \times BC B C × B C 的权重表,逐元素乘的规模从 B C 2 BC^2 B C 2 降到 2 B C D 2 BC D 2 B C D 。B C = 64 BC = 64 B C = 64 、D = 64 D = 64 D = 64 时前者是 4096 次乘法,后者 8192 次–乘法次数反而增加,但权重表不必常驻寄存器,这在 B C = 128 BC = 128 B C = 128 时是实质性的压力缓解。
问题在数值范围。γ − j ′ \gamma^{-j'} γ − j ′ 是大于 1 的量 ,且随 j ′ j' j ′ 指数增长:
γ \gamma γ
γ − 31 \gamma^{-31} γ − 31 (BC=32)
γ − 63 \gamma^{-63} γ − 63 (BC=64)
γ − 127 \gamma^{-127} γ − 127 (BC=128)
fp16 安全上限
fp32 安全上限
0.99
1.37 1.37 1.37
1.88 1.88 1.88
3.58 3.58 3.58
B C < 1103 BC < 1103 B C < 1103
B C < 8827 BC < 8827 B C < 8827
0.95
4.90 4.90 4.90
25.3 25.3 25.3
675 675 675
B C < 216 BC < 216 B C < 216
B C < 1729 BC < 1729 B C < 1729
0.90
26.2 26.2 26.2
763 763 763
6.47 × 10 5 6.47 \times 10^{5} 6.47 × 1 0 5
B C < 105 BC < 105 B C < 105
B C < 842 BC < 842 B C < 842
0.80
1.01 × 10 3 1.01 \times 10^{3} 1.01 × 1 0 3
1.27 × 10 6 1.27 \times 10^{6} 1.27 × 1 0 6
2.03 × 10 12 2.03 \times 10^{12} 2.03 × 1 0 12
B C < 49 BC < 49 B C < 49
B C < 397 BC < 397 B C < 397
0.50
2.15 × 10 9 2.15 \times 10^{9} 2.15 × 1 0 9
9.22 × 10 18 9.22 \times 10^{18} 9.22 × 1 0 18
1.70 × 10 38 1.70 \times 10^{38} 1.70 × 1 0 38
B C < 15 BC < 15 B C < 15
B C < 127 BC < 127 B C < 127
fp16 上限是 65504。γ = 0.8 \gamma = 0.8 γ = 0.8 、B C = 64 BC = 64 B C = 64 时 γ − 63 = 1.27 × 10 6 \gamma^{-63} = 1.27 \times 10^6 γ − 63 = 1.27 × 1 0 6 已经溢出;γ = 0.5 \gamma = 0.5 γ = 0.5 时 9.22 × 10 18 9.22 \times 10^{18} 9.22 × 1 0 18 连 fp32 都接近极限。
3.1 实测:单块块内计算的精度对比
单块块内计算的精度对比(B C = 64 BC = 64 B C = 64 ,D = 64 D = 64 D = 64 ,fp16 存储 + fp32 累加,20 组随机输入取中位数,参考值为 fp64 精确计算):
γ \gamma γ
比值形式 相对 L2
因式分解 相对 L2
γ − ( B C − 1 ) \gamma^{-(BC-1)} γ − ( B C − 1 )
0.99
4.39 × 10 − 4 4.39 \times 10^{-4} 4.39 × 1 0 − 4
4.15 × 10 − 4 4.15 \times 10^{-4} 4.15 × 1 0 − 4
1.88 1.88 1.88
0.95
4.57 × 10 − 4 4.57 \times 10^{-4} 4.57 × 1 0 − 4
4.14 × 10 − 4 4.14 \times 10^{-4} 4.14 × 1 0 − 4
25.3 25.3 25.3
0.90
4.47 × 10 − 4 4.47 \times 10^{-4} 4.47 × 1 0 − 4
4.15 × 10 − 4 4.15 \times 10^{-4} 4.15 × 1 0 − 4
763 763 763
0.80
4.60 × 10 − 4 4.60 \times 10^{-4} 4.60 × 1 0 − 4
NaN
1.27 × 10 6 1.27 \times 10^{6} 1.27 × 1 0 6
0.50
4.14 × 10 − 4 4.14 \times 10^{-4} 4.14 × 1 0 − 4
NaN
9.22 × 10 18 9.22 \times 10^{18} 9.22 × 1 0 18
结论清晰:
γ ≥ 0.9 \gamma \ge 0.9 γ ≥ 0.9 时两者精度相当,因式分解甚至略优(少一次逐元素乘引入的舍入);
γ ≤ 0.8 \gamma \le 0.8 γ ≤ 0.8 时因式分解在 fp16 下彻底失效 ,产生 NaN 而非精度下降–γ − j ′ \gamma^{-j'} γ − j ′ 溢出成 inf,随后 inf 乘 0 得 NaN;
比值形式的误差在全部 γ \gamma γ 取值下稳定在 4.1 – 4.6 × 10 − 4 4.1\text{--}4.6 \times 10^{-4} 4.1 – 4.6 × 1 0 − 4 ,与 γ \gamma γ 无关。
比值形式的误差之所以稳定,是因为 Γ i j \Gamma_{ij} Γ ij 恒在 ( 0 , 1 ] (0,1] ( 0 , 1 ] –这是一个与 γ \gamma γ 无关的数值保证 。实测各 γ \gamma γ 下衰减矩阵的取值:
γ \gamma γ
B C BC B C
最大值
最小非零值
下溢为 0 的比例
0.99
128
1.000
2.79 × 10 − 1 2.79 \times 10^{-1} 2.79 × 1 0 − 1
0.0%
0.90
128
1.000
1.55 × 10 − 6 1.55 \times 10^{-6} 1.55 × 1 0 − 6
0.0%
0.50
128
1.000
5.88 × 10 − 39 5.88 \times 10^{-39} 5.88 × 1 0 − 39
0.0%
即使 γ = 0.5 \gamma = 0.5 γ = 0.5 、B C = 128 BC = 128 B C = 128 ,最小权重 5.88 × 10 − 39 5.88 \times 10^{-39} 5.88 × 1 0 − 39 在 fp32 下仍是正规数,无下溢。下溢比溢出安全得多 :权重下溢为 0 意味着「这个远距离贡献可以忽略」,语义上正确;而溢出为 inf 会污染整行输出。
3.2 比值形式自身的下溢边界
比值形式也不是完全没有约束。Γ \Gamma Γ 本身若用 fp16 存储,γ B C − 1 \gamma^{BC-1} γ B C − 1 可能低于 fp16 最小正规数 6.10 × 10 − 5 6.10 \times 10^{-5} 6.10 × 1 0 − 5 :
γ \gamma γ
B C = 64 BC=64 B C = 64
fp16 状态
B C = 128 BC=128 B C = 128
fp16 状态
0.99
5.31 × 10 − 1 5.31 \times 10^{-1} 5.31 × 1 0 − 1
正常
2.79 × 10 − 1 2.79 \times 10^{-1} 2.79 × 1 0 − 1
正常
0.95
3.95 × 10 − 2 3.95 \times 10^{-2} 3.95 × 1 0 − 2
正常
1.48 × 10 − 3 1.48 \times 10^{-3} 1.48 × 1 0 − 3
正常
0.90
1.31 × 10 − 3 1.31 \times 10^{-3} 1.31 × 1 0 − 3
正常
1.55 × 10 − 6 1.55 \times 10^{-6} 1.55 × 1 0 − 6
非正规数
0.80
7.85 × 10 − 7 7.85 \times 10^{-7} 7.85 × 1 0 − 7
非正规数
4.93 × 10 − 13 4.93 \times 10^{-13} 4.93 × 1 0 − 13
下溢为 0
0.50
1.08 × 10 − 19 1.08 \times 10^{-19} 1.08 × 1 0 − 19
下溢为 0
5.88 × 10 − 39 5.88 \times 10^{-39} 5.88 × 1 0 − 39
下溢为 0
处理方式很简单:权重表用 f32 fragment 保存,只在喂入 MMA 前把乘完的分数矩阵降到 f16 。分数矩阵本身量级正常,降精度无损。这也是下面 kernel 采用的做法。
一句话总结 :因式分解能省一次 B C 2 BC^2 B C 2 逐元素乘,但把中间量范围从 ( 0 , 1 ] (0,1] ( 0 , 1 ] 推到 [ 1 , γ − ( B C − 1 ) ] [1, \gamma^{-(BC-1)}] [ 1 , γ − ( B C − 1 ) ] ,γ ≤ 0.8 \gamma \le 0.8 γ ≤ 0.8 时 fp16 直接 NaN–这个 FLOPs 优化不值得,论文的比值形式应原样保留。
4. 四层参考实现
沿用上一篇的三层框架,本级增加一层专门验证因式分解写法:
参考
实现方式
验证目标
A
逐 token 递归
递推式定义 S ← γ S + k v ⊤ S \leftarrow \gamma S + k v^\top S ← γ S + k v ⊤
B
分块向量化,三处衰减权重
§1.1 三个权重位置的推导
C
逐块独立重算,模拟 grid
kernel 控制流
D
因式分解版块内计算
§2.4 与 §3 两种写法的代数等价性
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.1 的位置一、二、三""" 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) decay_tri = torch.where(i[:, None ] >= i[None , :], g ** (i[:, None ] - i[None , :]), torch.zeros((), dtype=torch.float64, device=Q.device)) w_state = g ** (BC - 1 - i) w_query = g ** (i + 1 ) contrib = torch.einsum("bhcmd,bhcmv->bhcdv" , Kc * w_state[:, None ], Vc) 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 是无权重的前缀和,而本级递推带系数 γ B C \gamma^{BC} γ B C 。这里保留显式循环;若要向量化,需改用加权前缀和 的写法:把 Δ c \Delta_c Δ c 先除以 γ c B C \gamma^{cBC} γ c B C 再 cumsum、最后乘回,代价是又引入 γ − c B C \gamma^{-cBC} γ − c B C 这个溢出源,c c c 大时比 §3 的块内因式分解更危险。这是同一个取舍在跨块层面的重演 –它的本质仍是 §2.4 那个 1 / Λ 1/\Lambda 1/Λ 物化问题,只是尺度从 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_state = 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) 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_state[:, 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 , B C = 4 , γ = 0.9 B=2, H=2, N=12, D=4, BC=4, \gamma=0.9 B = 2 , H = 2 , N = 12 , D = 4 , B C = 4 , γ = 0.9 ,fp64(numpy 复现):
比较
max abs 误差
相对 L2
B 分块向量化 vs A 逐 token 递归
3.55 × 10 − 15 3.55 \times 10^{-15} 3.55 × 1 0 − 15
1.82 × 10 − 16 1.82 \times 10^{-16} 1.82 × 1 0 − 16
C kernel 结构镜像 vs A 逐 token 递归
2.67 × 10 − 15 2.67 \times 10^{-15} 2.67 × 1 0 − 15
1.83 × 10 − 16 1.83 \times 10^{-16} 1.83 × 1 0 − 16
D 因式分解 vs A 逐 token 递归
3.55 × 10 − 15 3.55 \times 10^{-15} 3.55 × 1 0 − 15
1.98 × 10 − 16 1.98 \times 10^{-16} 1.98 × 1 0 − 16
B(γ = 1.0 \gamma = 1.0 γ = 1.0 )vs 上一级参考 A
3.55 × 10 − 15 3.55 \times 10^{-15} 3.55 × 1 0 − 15
1.62 × 10 − 16 1.62 \times 10^{-16} 1.62 × 1 0 − 16
前三行确认四份实现数学等价,误差均在 fp64 机器精度量级。第四行是退化检验 :令 γ = 1 \gamma = 1 γ = 1 应当精确回到上一级的无衰减实现,这条验证能同时捕获三处权重中任何一处的指数写错–例如把 γ B C − 1 − m ′ \gamma^{BC-1-m'} γ B C − 1 − m ′ 误写成 γ B C − m ′ \gamma^{BC-m'} γ B C − m ′ ,γ = 1 \gamma = 1 γ = 1 时两者都是 1,退化检验通过但 γ = 0.9 \gamma = 0.9 γ = 0.9 时参考 B 与 A 不符。两条验证必须都做。
5. TileLang kernel 的改动
grid 划分、shared/fragment 分配与上一级一致,新增一个 f32 的权重表:
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 @tilelang.jit(out_idx=[3 ] ) def linattn_decay_chunk (batch, heads, seq_len, dim, blk, gamma, 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) 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) Dtri = T.alloc_fragment([BC, BC], accum_dtype)
5.1 预计算衰减下三角
1 2 3 4 5 6 for i, j in T.Parallel(BC, BC): Dtri[i, j] = T.if_then_else( j <= i, T.exp2(T.log2(T.Cast(accum_dtype, gamma)) * T.Cast(accum_dtype, i - j)), 0.0 )
指数用 exp2(log2(γ) · (i-j)) 而非 pow(γ, i-j)。原因是 exp2 与 log2 都有单指令硬件实现,而 pow 通常展开成多条指令;log 2 γ \log_2 \gamma log 2 γ 是循环不变量,编译器会提到循环外。这个改写在下一级会变成必需 –如 §2.5 所述,逐 token 门控的 g t = log α t g_t = \log \alpha_t g t = log α t 本身就存在 log 域,直接就是 exp2 的输入,不需要再取对数。
5.2 状态累加加入整块衰减
1 2 3 4 5 6 7 8 9 10 11 12 13 T.clear(S_f) for c in T.Pipelined(bx, num_stages=num_stages): T.copy(K[bz, c * BC:(c + 1 ) * BC, by, :], K_s) T.copy(V[bz, c * BC:(c + 1 ) * BC, by, :], V_s) for i, j in T.Parallel(dim, dim): S_f[i, j] *= gamma_pow_bc for m, d in T.Parallel(BC, dim): K_s[m, d] *= w_state[m] T.gemm(K_s, V_s, S_f, transpose_A=True ) T.copy(S_f, S_s)
这里有一个次序约束:S_f *= γ^BC 必须在 T.gemm 之前 。递推式是 S c = γ B C S c − 1 + Δ c − 1 S_c = \gamma^{BC} S_{c-1} + \Delta_{c-1} S c = γ B C S c − 1 + Δ c − 1 ,先衰减旧状态再累加新贡献;若顺序颠倒,本块贡献会被多衰减一次。参考 C 里写作 S = g**BC * S + Kd.T @ V,一行内表达了这个次序,翻译成 kernel 的两条语句时必须保持。
K_s 被原地缩放,因此步骤③重载当前块 K/V 的必要性比上一级更强–上一级只是「残留内容不确定」,本级是「残留内容已被乘过权重」,复用会导致权重叠加两次。
5.3 两项贡献与 query 侧缩放
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 T.copy(Q[bz, bx * BC:(bx + 1 ) * BC, by, :], Q_s) for i, d in T.Parallel(BC, dim): Q_s[i, d] *= w_query[i] T.clear(acc_o) T.gemm(Q_s, S_s, acc_o) T.copy(Q[bz, bx * BC:(bx + 1 ) * BC, by, :], Q_s) 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] *= 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[bz, bx * BC:(bx + 1 ) * BC, by, :])
步骤③开头必须重新载入 Q ,因为步骤②已把 Q_s 原地乘上了 γ i ′ + 1 \gamma^{i'+1} γ i ′ + 1 ,而块内那一项用的是未缩放的 Q c Q_c Q c 。这是本级新增的一处陷阱:上一级 Q_s 在两步之间可以直接复用。
原地缩放省了一个 buffer,代价是必须重载;若寄存器与 shared memory 有余量,另开一个 Q_scaled 更安全。这个取舍在 B C = 128 BC = 128 B C = 128 时倾向于原地缩放。
5.4 与上一级 kernel 的改动汇总
位置
上一级
本级
新增开销
权重表
无
Dtri[BC, BC] f32 fragment
B C 2 BC^2 B C 2 个 f32 寄存器
状态累加
T.gemm 直接累加
先 S_f *= γ^BC 再 gemm
D 2 D^2 D 2 次乘法 / 轮
K 写入状态
原样
逐行乘 γ B C − 1 − m ′ \gamma^{BC-1-m'} γ B C − 1 − m ′
B C × D BC \times D B C × D 次乘法 / 轮
Q 读取状态
原样
逐行乘 γ i ′ + 1 \gamma^{i'+1} γ i ′ + 1
B C × D BC \times D B C × D 次乘法
块内掩码
if_then_else 置 0
逐元素乘 Dtri
B C 2 BC^2 B C 2 次乘法
Q 复用
步骤②③共用
步骤③必须重载
一次 shared 写入
Dtri 占 B C 2 BC^2 B C 2 个 f32 寄存器是本级最主要的资源代价。B C = 64 BC = 64 B C = 64 时是 4096 个 f32、128 线程下每线程 32 个;B C = 128 BC = 128 B C = 128 时升到 16384 个、每线程 128 个–与 A 加起来已经接近寄存器上限。这是 B C = 128 BC = 128 B C = 128 在本级比上一级更难用的原因 ,也是因式分解唯一真正有吸引力的地方(它不需要 B C 2 BC^2 B C 2 权重表)。但 §3.1 的 NaN 结论表明这个吸引力不成立。
6. 衰减带来的新维度:有效记忆长度
γ \gamma γ 引入了一个上一级不存在的语义参数–状态的遗忘速度。定义有效记忆长度 为权重衰减到 10 − 3 10^{-3} 1 0 − 3 所需的 token 数,即 γ L = 10 − 3 \gamma^{L} = 10^{-3} γ L = 1 0 − 3 :
L = ln 10 − 3 ln γ L = \frac{\ln 10^{-3}}{\ln \gamma}
L = ln γ ln 1 0 − 3
γ \gamma γ
γ 64 \gamma^{64} γ 64
有效记忆长度
0.999
9.38 × 10 − 1 9.38 \times 10^{-1} 9.38 × 1 0 − 1
6904 token
0.99
5.26 × 10 − 1 5.26 \times 10^{-1} 5.26 × 1 0 − 1
687 token
0.95
3.75 × 10 − 2 3.75 \times 10^{-2} 3.75 × 1 0 − 2
135 token
0.90
1.18 × 10 − 3 1.18 \times 10^{-3} 1.18 × 1 0 − 3
66 token
0.50
5.42 × 10 − 20 5.42 \times 10^{-20} 5.42 × 1 0 − 20
10 token
这张表解释了为什么实际模型里的门控值普遍接近 1:γ = 0.9 \gamma = 0.9 γ = 0.9 的有效记忆只有 66 token,连一个 chunk(B C = 64 BC = 64 B C = 64 )都刚刚覆盖,长程依赖完全丢失。而 γ ≥ 0.99 \gamma \ge 0.99 γ ≥ 0.99 恰好落在 §2.4 与 §3 表格中因式分解仍然安全的区间 –这解释了为什么部分实现敢做这个分解:它们隐含假设了门控接近 1。
这个假设在标量门控下大致成立,但在第三级逐通道门控下会失效 。逐通道意味着每个 key 通道有独立的 γ d \gamma_d γ d ,训练中总会有部分通道学到较小的值以实现快速遗忘。此时「所有通道都接近 1」不再成立,因式分解的溢出从例外变成常态。Kimi K3 的处理方式是给门控加下界,这部分留到第三级展开。
一句话总结 :γ \gamma γ 的安全区间(接近 1)与有用区间(提供实际遗忘能力)方向相反,标量门控下两者尚可兼顾,逐通道门控下这个矛盾会暴露。
7. 数值验证
验证层次沿用上一级结构,新增两项:
四层参考互验 (fp64):A/B/C/D 两两对照,误差应在 10 − 15 10^{-15} 1 0 − 15 量级;
退化检验 :γ = 1 \gamma = 1 γ = 1 时参考 B 应精确回到上一级实现–这条能捕获三处权重的指数偏移错误;
fp16 失效点复现 :γ ≤ 0.8 \gamma \le 0.8 γ ≤ 0.8 时因式分解写法应产生 NaN,确认 §3.1 结论;
kernel vs 参考 C :fp16 输入 + f32 累加,阈值取相对 L2 < 2 × 10 − 2 < 2 \times 10^{-2} < 2 × 1 0 − 2 ;
延迟对比走 CUDA event 中位数。
第 1、2、3 层已在 numpy fp64 上完成(见 §2.3、§3.1 与 §4.4 表格)。第 4、5 层需要 CUDA 设备,待真卡跑通后单独补充实测数据 ,此处不做性能推测。
本级 kernel 的预期主导误差源与上一级相同–T.copy(S_f, S_s) 的 f32→f16 降精度。但衰减带来一个有利变化:γ B C \gamma^{BC} γ B C 每轮把旧状态压缩,早期块的累积误差也随之衰减,因此误差不再像上一级那样随 b x bx b x 单调增长,而是趋于一个稳态。γ = 0.9 \gamma = 0.9 γ = 0.9 、B C = 64 BC = 64 B C = 64 时 γ B C = 1.18 × 10 − 3 \gamma^{BC} = 1.18 \times 10^{-3} γ B C = 1.18 × 1 0 − 3 ,约三轮之后早期误差已不可见。衰减机制顺带改善了数值稳定性 ,这是一个反直觉但合理的副作用。
8. 总结
标量衰减不改变分块恒等式的结构 ,只在三个位置插入指数权重:块内下三角 γ i ′ − j ′ \gamma^{i'-j'} γ i ′ − j ′ 、写入状态时的块尾对齐 γ B C − 1 − m ′ \gamma^{BC-1-m'} γ B C − 1 − m ′ 、跨块递推 γ B C \gamma^{BC} γ B C 与 query 侧 γ i ′ + 1 \gamma^{i'+1} γ i ′ + 1 。四个权重全部 ≤ 1 \le 1 ≤ 1 ,这是本级数值安全的根本原因。
累积衰减积 Λ j = ∏ i ≤ j α i \Lambda_j = \prod_{i \le j} \alpha_i Λ j = ∏ i ≤ j α i 是统一两种形式的代数工具 。它把区间连乘 ∏ i = j + 1 t α i \prod_{i=j+1}^{t} \alpha_i ∏ i = j + 1 t α i 化归为前缀量之比 Λ t / Λ j \Lambda_t / \Lambda_j Λ t / Λ j ,于是递归的向量形式与并行的矩阵形式 O = ( ( Q K ⊤ ) ⊙ Γ ) V O = ((QK^\top) \odot \Gamma)V O = (( Q K ⊤ ) ⊙ Γ ) V 可以相互转换(即 Mamba2 所谓的状态空间对偶性,实测相对 L2 2.03 × 10 − 16 2.03 \times 10^{-16} 2.03 × 1 0 − 16 )。常数 γ \gamma γ 是 Λ j = γ j \Lambda_j = \gamma^j Λ j = γ j 的特例,此时 Γ i j = γ i − j \Gamma_{ij} = \gamma^{i-j} Γ ij = γ i − j 。
论文的衰减矩阵 Γ i j = Λ i / Λ j \Gamma_{ij} = \Lambda_i/\Lambda_j Γ ij = Λ i / Λ j 本身是数值安全的 –因果约束 i ≥ j i \ge j i ≥ j 使它等于 ∏ r = j + 1 i α r ≤ 1 \prod_{r=j+1}^{i}\alpha_r \le 1 ∏ r = j + 1 i α r ≤ 1 ,全程不需要物化大于 1 的量(实测 Γ i j ∈ [ 0.628 , 1.000 ] \Gamma_{ij} \in [0.628,\ 1.000] Γ ij ∈ [ 0.628 , 1.000 ] )。论文的箭头记号 q ← \overleftarrow{q} q 、k → \overrightarrow{k} k 、S → \overrightarrow{S} S 把「衰减到首 / 末位置」直接编码在箭头方向上,基准点选得当就能保证所有指数非正。
真正的陷阱是把比值因式分解成 Λ i ⋅ ( 1 / Λ j ) \Lambda_i \cdot (1/\Lambda_j) Λ i ⋅ ( 1/ Λ j ) 。这个分解能把两步合成单次 GEMM 并免去 B C 2 BC^2 B C 2 权重表,但要求物化 1 / Λ j 1/\Lambda_j 1/ Λ j ,把中间量从 ( 0 , 1 ] (0,1] ( 0 , 1 ] 推到 [ 1 , 1 / Λ B C ] [1,\ 1/\Lambda_{BC}] [ 1 , 1/ Λ B C ] 。实测(B C = 64 BC=64 B C = 64 ,fp16 存储加 fp32 累加):γ ≥ 0.9 \gamma \ge 0.9 γ ≥ 0.9 时两者精度相当(4.1 – 4.6 × 10 − 4 4.1\text{--}4.6 \times 10^{-4} 4.1 – 4.6 × 1 0 − 4 ),γ ≤ 0.8 \gamma \le 0.8 γ ≤ 0.8 时分解写法产生 NaN –γ = 0.8 \gamma = 0.8 γ = 0.8 时 γ − 63 = 1.27 × 10 6 \gamma^{-63} = 1.27 \times 10^{6} γ − 63 = 1.27 × 1 0 6 ,已远超 fp16 上限 65504。data-dependent α t ∈ [ 0.8 , 0.95 ] \alpha_t \in [0.8,\ 0.95] α t ∈ [ 0.8 , 0.95 ] 、B C = 128 BC = 128 B C = 128 时 1 / Λ 1/\Lambda 1/Λ 最大达 3.40 × 10 8 3.40 \times 10^{8} 3.40 × 1 0 8 ,同样溢出。结论是保留论文的比值形式,不做分解。
累积积本身也必须在 log 域计算 。fp16 直接 cumprod 在 B C = 128 BC = 128 B C = 128 、α ∈ [ 0.5 , 0.9 ] \alpha \in [0.5,\ 0.9] α ∈ [ 0.5 , 0.9 ] 时已完全下溢为 0,fp32 在 B C = 512 BC = 512 B C = 512 时进入非正规数区间;log Λ j = ∑ i ≤ j log α i \log \Lambda_j = \sum_{i \le j} \log \alpha_i log Λ j = ∑ i ≤ j log α i 是线性增长的负数,表示范围安全。这就是门控全程存 log 值、用 exp2 还原的原因。
退化检验是本级新增的关键验证手段 :令 γ = 1 \gamma = 1 γ = 1 应精确回到上一级实现(实测相对 L2 1.62 × 10 − 16 1.62 \times 10^{-16} 1.62 × 1 0 − 16 )。但它无法单独定案–三处权重中任何一处的指数偏移在 γ = 1 \gamma = 1 γ = 1 时都不可见,必须同时做 γ = 0.9 \gamma = 0.9 γ = 0.9 的四层参考互验(实测 1.82 – 2.03 × 10 − 16 1.82\text{--}2.03 \times 10^{-16} 1.82 – 2.03 × 1 0 − 16 )。
两处新增实现陷阱 :S_f *= γ^BC 必须在 T.gemm 之前,否则本块贡献被多衰减一次;步骤②把 Q_s 原地乘了 γ i ′ + 1 \gamma^{i'+1} γ i ′ + 1 ,步骤③用的是未缩放 Q c Q_c Q c ,必须重新载入。
有效记忆长度暴露了一个内在矛盾 :γ \gamma γ 的数值安全区间(接近 1)与实际有用区间(提供遗忘能力)方向相反。γ = 0.9 \gamma = 0.9 γ = 0.9 的有效记忆仅 66 token,连一个 chunk 都刚覆盖;而 γ ≥ 0.99 \gamma \ge 0.99 γ ≥ 0.99 才落在因式分解仍然安全的区间。标量门控下两者尚可兼顾。
下一篇把 γ \gamma γ 升级为逐 token 门控 g t g_t g t :衰减不再是编译期常量,需要在 log 域做 cumsum 得到 §2.1 那个累积积 Λ j \Lambda_j Λ j ,权重表从预计算变成运行时构造,exp2 从优化手段变成必需。此时 §6 的矛盾会真正暴露–每个通道独立的门控使「所有 γ \gamma γ 接近 1」的假设失效,Λ \Lambda Λ 从标量变成向量,1 / Λ 1/\Lambda 1/Λ 的溢出从例外变成常态。
可迁移的启示 :代数上等价的两种写法,数值行为可以完全不同。Λ i / Λ j \Lambda_i / \Lambda_j Λ i / Λ j 与 Λ i ⋅ ( 1 / Λ j ) \Lambda_i \cdot (1/\Lambda_j) Λ i ⋅ ( 1/ Λ j ) 在实数域相等,但前者(j ≤ i j \le i j ≤ i 时)恒不大于 1、后者上界随块长指数增长。论文把衰减写成比值而非乘积并不是记号习惯,而是保证数值安全的必要安排 –那套 ⋅ ← \overleftarrow{\cdot} ⋅ / ⋅ → \overrightarrow{\cdot} ⋅ 箭头记号就是在提醒读者每个衰减都有明确的基准点。实现时不要对论文公式做看似等价的代数改写,先问改写后的中间量落在什么范围。
参考 :
Gated Delta Networks(GDN):arXiv:2412.06464,ICLR 2025。§2.1 给出累积衰减积 γ j = ∏ i = 1 j α i \gamma_j = \prod_{i=1}^{j}\alpha_i γ j = ∏ i = 1 j α i 、向量 / 矩阵并行两种形式与衰减感掩码 Γ i j = γ i / γ j \Gamma_{ij} = \gamma_i/\gamma_j Γ ij = γ i / γ j ;式 (1)(2) 给出 chunkwise 形式与 ⋅ ← \overleftarrow{\cdot} ⋅ / ⋅ → \overrightarrow{\cdot} ⋅ 箭头记号
Kimi Linear / KDA 论文:arXiv:2510.26692
flash-linear-attention:fla/ops/simple_gla/(标量衰减的参考实现)
Mamba2 / GLA:标量与门控衰减的 chunkwise 形式
本站《TileLang 实战:KDA 从零到一–Chunked 线性注意力 》《KDA 的来龙去脉 》《TileLang 编程基本知识点 》
本文的四层参考互验、退化检验与 fp16 失效点在 numpy fp64/fp16 上复现;GPU 实测数据待补。