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

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

上一篇实现了不带任何遗忘机制的 chunked 线性注意力,状态单调累加。本篇引入第一个衰减因子:把递推式改为 SγS+KVS \leftarrow \gamma S + K^\top Vγ\gamma 是一个标量常数。

改动看起来只是多乘一个数,但它决定了后续三级的全部数值策略。本文先引入 GDN 论文用来统一递归形式与并行形式的累积衰减积 Λj=ijαi\Lambda_j = \prod_{i \le j} \alpha_i,它把任意区间的衰减化归为两个前缀量之比。论文给出的矩阵并行形式是 O=((QK)Γ)VO = ((QK^\top) \odot \Gamma)V,其中 Γij=Λi/Λj\Gamma_{ij} = \Lambda_i/\Lambda_j这是一个比值,因果约束下恒不大于 1,数值安全。但它存在一个看似有利的因式分解 Λi/Λj=Λi(1/Λj)\Lambda_i/\Lambda_j = \Lambda_i \cdot (1/\Lambda_j),能把两步合成单次 GEMM 并省去 BC2BC^2 权重表。本文实测给出结论:这个分解会把中间量推到 1/ΛBC1/\Lambda_{BC}γ0.8\gamma \le 0.8 时 fp16 下直接产生 NaN。论文的比值形式必须原样保留。这个约束在第三级 Λ\Lambda 从标量升级为逐通道向量后会变成 KDA 实现中最棘手的问题。

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


1. 递推式与分块重写

本级递推式在上一级基础上给状态加一个遗忘系数 γ(0,1]\gamma \in (0, 1]

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

展开成显式求和,注意每个 kjvjk_j v_j^\top 被后续每一步各乘一次 γ\gamma,从 jjtt 共乘 tjt - j 次:

St=jtγtjkjvj,oi=jiγij(qikj)vjS_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

对比上一级的 oi=ji(qikj)vjo_i = \sum_{j \le i} (q_i \cdot k_j) v_j,唯一变化是每一项多了权重 γij\gamma^{i-j}–距离越远权重越小,这就是「衰减」的含义。γ=1\gamma = 1 时退化为上一级。

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

沿用上一级的分块框架,序列按 BCBC 切分为 NCNC 个 chunk。把 jij \le i 的求和拆成跨块与块内两部分,衰减权重会分别落到三个位置。设 token ii 在其所属 chunk 内的局部下标为 i=imodBCi' = i \bmod BC

位置一–块内下三角。同块内 jij' \le i',相对距离就是局部下标之差:

Aijintra={γijji0j>iA^{\text{intra}}_{i'j'} = \begin{cases} \gamma^{\,i'-j'} & j' \le i' \\ 0 & j' > i' \end{cases}

上一级这里是 0/1 掩码,本级变成指数下三角。注意权重全部落在 (0,1](0, 1] 区间–对角线是 γ0=1\gamma^0 = 1,左下角最小值是 γBC1\gamma^{BC-1}

位置二–每块写入状态时的块尾对齐。跨块状态需要定义一个统一的时间基准,取所属 chunk 的末尾。chunk cc 内第 mm' 个 token 的贡献衰减到该块末尾要乘 γBC1m\gamma^{BC-1-m'}

Δc=m=0BC1γBC1mkmvm=(KcγBC1m)Vc\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

位置三–跨块状态递推与 query 侧缩放。相邻块之间隔了整块 BCBC 步,因此前缀状态的递推是:

Scprev=γBCSc1prev+Δc1S_c^{\text{prev}} = \gamma^{BC} S_{c-1}^{\text{prev}} + \Delta_{c-1}

而 token ii' 读取这个状态时,它距离上一块末尾还有 i+1i' + 1 步,所以 query 要乘 γi+1\gamma^{\,i'+1}

Oc=(Qcγi+1)Scprev+AintraVcO_c = (Q_c \odot \gamma^{\,i'+1}) S_c^{\text{prev}} + A^{\text{intra}} V_c

三个位置的权重汇总:

位置 权重 取值范围 作用对象
块内下三角 γij\gamma^{\,i'-j'} [γBC1, 1][\gamma^{BC-1},\ 1] BC×BCBC \times BC 分数矩阵,逐元素
写入状态 γBC1m\gamma^{\,BC-1-m'} [1, γBC1][1,\ \gamma^{BC-1}] KcK_c 的行,逐行缩放
跨块递推 γBC\gamma^{BC} 标量 整个状态矩阵
读取状态 γi+1\gamma^{\,i'+1} [γ, γBC][\gamma,\ \gamma^{BC}] QcQ_c 的行,逐行缩放

四个权重全部 1\le 1,这一点是本级数值安全的根本原因,也是下一节讨论的分歧点。

一句话总结:标量衰减不改变分块恒等式的结构,只是在三个位置插入指数权重;GEMM 的形状与调用次序与上一级完全一致。


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

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

2.1 累积衰减积的定义

把标量衰减推广为依赖数据的 αt(0,1)\alpha_t \in (0, 1)–每个时刻的遗忘强度由输入决定,而非固定常数(本级的常数 γ\gammaαtγ\alpha_t \equiv \gamma 的特例)。定义累积衰减积

Λj=i=1jαi\Lambda_j = \prod_{i=1}^{j} \alpha_i

Λj\Lambda_j 的含义是从序列起点衰减到第 jj 步的总折扣。有了它,递推式的展开可以写得非常紧凑。展开 St=αtSt1+ktvtS_t = \alpha_t S_{t-1} + k_t v_t^\top

St=jt(i=j+1tαi)kjvj=jtΛtΛjkjvjS_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

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

2.2 两种等价形式

代入 ot=qtSto_t = q_t^\top S_t,同一个结果可以写成两种形式:

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

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

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

O=((QK)Γ)V,Γij={ΛiΛjij0i<jO = \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}

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

[(QK)Γ]ij=ΛiΛj(qikj)(ji)\big[(Q K^\top) \odot \Gamma\big]_{ij} = \frac{\Lambda_i}{\Lambda_j} (q_i \cdot k_j) \quad (j \le i)

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

2.3 衰减矩阵的数值范围:为何比值形式是安全的

这里有一个容易看错的关键细节:Γij=Λi/Λj\Gamma_{ij} = \Lambda_i / \Lambda_j 是一个比值,而且因为因果约束 iji \ge jΛ\Lambda 单调递减,所以:

0<Γij=ΛiΛj=r=j+1iαr10 < \Gamma_{ij} = \frac{\Lambda_i}{\Lambda_j} = \prod_{r=j+1}^{i} \alpha_r \le 1

论文的形式是数值安全的–它先算 QKQ K^\top 再逐元素乘 Γ\Gamma,全程不需要物化任何大于 1 的量。实测确认(BC=8BC = 8αtU(0.85, 0.999)\alpha_t \sim \mathcal{U}(0.85,\ 0.999),fp64):

验证项 结果
论文式 (1) vs 逐 token 递归 max abs 误差 1.78×10151.78 \times 10^{-15}
状态更新 S\overrightarrow{S} vs 递归末态 max abs 误差 8.88×10168.88 \times 10^{-16}
Γij\Gamma_{ij} 取值范围(iji \ge j [0.628, 1.000][0.628,\ 1.000]全部 1\le 1

论文还给了一套简洁的箭头记号来表达这三个方向的衰减,与 §1.1 推导的三个位置一一对应:

论文记号 定义 含义 对应 §1.1
qr=Λrqr\overleftarrow{q^r} = \Lambda_r\, q^r 衰减到 chunk 位置 query 侧缩放 位置三(读取状态)
kr=ΛCΛrkr\overrightarrow{k^r} = \dfrac{\Lambda_C}{\Lambda_r} k^r 衰减到 chunk 位置 key 侧缩放 位置二(块尾对齐)
S=ΛCS\overrightarrow{S} = \Lambda_C\, S 整块衰减 状态递推 位置三(跳块递推)

注意 kr\overrightarrow{k^r} 里的 ΛC/Λr\Lambda_C / \Lambda_r 同样是比值且 1\le 1(因为 rCr \le C)。论文从头到尾没有单独物化过 1/Λj1/\Lambda_j,所有衰减因子都以比值形式出现。这是一个值得学习的记号设计–箭头方向直接编码了「衰减到哪个基准点」,而基准点选得当(首或末)就能保证指数非正。

常数 γ\gamma 是它的特例:Λj=γj\Lambda_j = \gamma^j,于是 Γij=γij\Gamma_{ij} = \gamma^{i-j},回到 §1.1 的位置一。实测验证(BC=4BC = 4,fp64):Λi/Λj\Lambda_i/\Lambda_jγij\gamma^{i-j} 的最大偏差 2.22×10162.22 \times 10^{-16},即机器精度。

两种形式的 fp64 数值一致性(data-dependent αtU(0.85, 0.999)\alpha_t \sim \mathcal{U}(0.85,\ 0.999)B=2,H=2,N=12,D=4,BC=4B=2, H=2, N=12, D=4, BC=4):

比较 max abs 误差 相对 L2
矩阵并行形式 vs 向量形式(递归) 3.55×10153.55 \times 10^{-15} 2.03×10162.03 \times 10^{-16}
不外提形式 vs 向量形式(递归) 3.55×10153.55 \times 10^{-15} 1.89×10161.89 \times 10^{-16}
常数 αt=0.9\alpha_t = 0.9 退化检验 2.67×10152.67 \times 10^{-15} 1.83×10161.83 \times 10^{-16}

2.4 一个容易走错的变形:把比值拆成乘积

既然论文形式是安全的,为何还要讨论数值问题?因为 Γij=Λi/Λj\Gamma_{ij} = \Lambda_i/\Lambda_j 存在一个看似有利的因式分解:

ΛiΛj=Λi1Λj(QK)Γ=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]

左侧需要先做 BC×BCBC \times BC 的 GEMM、再逐元素乘一张 BC2BC^2 的权重表;右侧把衰减推到两侧的逐行缩放上,只需一次 GEMM、不需权重表。看起来是纯改进。

但右侧必须显式物化 1/Λj1/\Lambda_j,而这个量不小于 1 且随 jj 指数增长。 左侧的 Γij\Gamma_{ij} 恒在 (0,1](0,1],右侧的中间量却能到 1/ΛBC1/\Lambda_{BC}。两者在实数域完全相等,在有限精度下完全不同。

data-dependent αt\alpha_t1/Λj1/\Lambda_j 的实测范围(2000 组随机采样取最大值):

αt\alpha_t 采样区间 BC=64BC=64 最大 1/Λ1/\Lambda BC=128BC=128 最大 1/Λ1/\Lambda fp16 溢出
[0.99, 0.999][0.99,\ 0.999] 1.551.55 2.232.23
[0.95, 0.999][0.95,\ 0.999] 7.747.74 47.647.6
[0.90, 0.990][0.90,\ 0.990] 84.484.4 4.18×1034.18 \times 10^{3}
[0.80, 0.950][0.80,\ 0.950] 2.66×1042.66 \times 10^{4} 3.40×1083.40 \times 10^{8} BC=128BC=128 溢出
[0.50, 0.900][0.50,\ 0.900] 1.59×10121.59 \times 10^{12} 6.23×10236.23 \times 10^{23} 均溢出

对照同一组 α\alpha 下两种写法的中间量范围(BC=64BC = 64):

αt\alpha_t 区间 因式分解后 max(1/Λj)\max(1/\Lambda_j) 论文形式 maxΓij\max \Gamma_{ij} 论文形式最小非零
[0.99, 0.999][0.99,\ 0.999] 1.491.49 1.0001.000 6.75×1016.75 \times 10^{-1}
[0.90, 0.990][0.90,\ 0.990] 33.733.7 1.0001.000 3.26×1023.26 \times 10^{-2}
[0.80, 0.950][0.80,\ 0.950] 9.34×1039.34 \times 10^{3} 1.0001.000 1.33×1041.33 \times 10^{-4}

论文形式的中间量上界恒为 1,与 α\alpha 的分布无关;因式分解后的上界随 α\alpha 变小而指数恶化。§3 给出 fp16 下的实测失效点。

2.5 为什么累积积要在 log 域计算

即使不做外提,Λj\Lambda_j 本身也不宜用 cumprod 直接计算–连乘会下溢。实测(fp32 / fp16 直接连乘对比 fp64 log 域 cumsum):

αt\alpha_t 区间 BCBC fp32 cumprod fp16 cumprod fp64 log-cumsum
[0.95, 0.999][0.95,\ 0.999] 512 1.93×1061.93 \times 10^{-6} 2.03×1062.03 \times 10^{-6} 1.93×1061.93 \times 10^{-6}
[0.90, 0.990][0.90,\ 0.990] 512 1.69×10131.69 \times 10^{-13} 2.38×1072.38 \times 10^{-7} 1.69×10131.69 \times 10^{-13}
[0.50, 0.900][0.50,\ 0.900] 128 1.58×10211.58 \times 10^{-21} 0\mathbf{0} 1.58×10211.58 \times 10^{-21}
[0.50, 0.900][0.50,\ 0.900] 512 1.40×10451.40 \times 10^{-45} 0\mathbf{0} 1.81×10841.81 \times 10^{-84}

fp16 连乘在 BC=128BC = 128α[0.5, 0.9]\alpha \in [0.5,\ 0.9] 时已完全下溢为 0;fp32 在 BC=512BC = 512 时也进入非正规数区间(1.40×10451.40 \times 10^{-45} 已是 fp32 最小非正规数量级)。log 域做加法则不受影响–logΛj=ijlogαi\log \Lambda_j = \sum_{i \le j} \log \alpha_i 是线性增长的负数,表示范围绰绰有余。

因此实现上的标准做法是:门控值全程以 logαt\log \alpha_t 的形式存储,累积积用 cumsum 而非 cumprod,需要衰减因子时用 exp2 还原。这解释了 §5.1 为什么采用 exp2(log2(γ)·(i-j)) 的写法–本级 γ\gamma 是常数,这么写只是为了用上硬件指令;下一级 gt=logαtg_t = \log \alpha_t 本身就是网络输出,log 域是它的原生形式,exp2 从优化手段变成结构必需。

一句话总结:累积积 Λj=ijαi\Lambda_j = \prod_{i \le j} \alpha_i 把区间连乘化归为前缀量之比,这是矩阵并行形式成立的代数基础;但比值一旦被因式分解成 Λi(1/Λj)\Lambda_i \cdot (1/\Lambda_j),就必须物化 1/Λj1/\Lambda_j 这个不小于 1 且指数增长的量。


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

§2.4 的两种写法在实数域完全等价,但落到有限精度上行为分开。本节把问题收窄到常数 γ\gamma 的情形做定量分析。

块内下三角 γij\gamma^{i'-j'} 存在一个看似有利的代数变形–指数可以拆开:

γij=γiγj\gamma^{\,i'-j'} = \gamma^{\,i'} \cdot \gamma^{-j'}

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

Aintra=tril((Qcγi)(Kcγj))A^{\text{intra}} = \operatorname{tril}\big( (Q_c \odot \gamma^{\,i'}) (K_c \odot \gamma^{-j'})^\top \big)

两种实现方式的差别:

方案 做法 块内额外开销 权重数值范围
比值形式(论文) 先 GEMM 得分数,再逐元素乘 Γ\Gamma 一次 BC×BCBC \times BC 逐元素乘 + 一张 BC2BC^2 权重表 (0,1](0, 1]
因式分解 先分别缩放 QQKK 的行,再单次 GEMM 两次 BC×DBC \times D 逐行缩放,无 BC2BC^2 开销 γj\gamma^{-j'} 最大 γ(BC1)\gamma^{-(BC-1)}

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

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

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

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

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

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

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

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

γ\gamma BCBC 最大值 最小非零值 下溢为 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\gamma = 0.5BC=128BC = 128,最小权重 5.88×10395.88 \times 10^{-39} 在 fp32 下仍是正规数,无下溢。下溢比溢出安全得多:权重下溢为 0 意味着「这个远距离贡献可以忽略」,语义上正确;而溢出为 inf 会污染整行输出。

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

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

γ\gamma BC=64BC=64 fp16 状态 BC=128BC=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 采用的做法。

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


4. 四层参考实现

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

参考 实现方式 验证目标
A 逐 token 递归 递推式定义 SγS+kvS \leftarrow \gamma S + k v^\top
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)
# 位置一:块内衰减下三角 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_state = g ** (BC - 1 - i) # 位置二:写入状态时衰减到块尾
w_query = g ** (i + 1) # 位置三:读取状态时补上距上块末尾的步数

# 每块对状态的贡献(已按块尾对齐)
contrib = torch.einsum("bhcmd,bhcmv->bhcdv", Kc * w_state[:, 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 是无权重的前缀和,而本级递推带系数 γBC\gamma^{BC}。这里保留显式循环;若要向量化,需改用加权前缀和的写法:把 Δc\Delta_c 先除以 γcBC\gamma^{cBC}cumsum、最后乘回,代价是又引入 γcBC\gamma^{-cBC} 这个溢出源,cc 大时比 §3 的块内因式分解更危险。这是同一个取舍在跨块层面的重演–它的本质仍是 §2.4 那个 1/Λ1/\Lambda 物化问题,只是尺度从 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)
# ① 流式累加:每轮先整块衰减 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_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,BC=4,γ=0.9B=2, H=2, N=12, D=4, BC=4, \gamma=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\gamma = 1.0)vs 上一级参考 A 3.55×10153.55 \times 10^{-15} 1.62×10161.62 \times 10^{-16}

前三行确认四份实现数学等价,误差均在 fp64 机器精度量级。第四行是退化检验:令 γ=1\gamma = 1 应当精确回到上一级的无衰减实现,这条验证能同时捕获三处权重中任何一处的指数写错–例如把 γBC1m\gamma^{BC-1-m'} 误写成 γBCm\gamma^{BC-m'}γ=1\gamma = 1 时两者都是 1,退化检验通过但 γ=0.9\gamma = 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)
# 新增:衰减矩阵 Γ,f32 保存以避免 §3.2 的 fp16 下溢
Dtri = T.alloc_fragment([BC, BC], accum_dtype)

5.1 预计算衰减下三角

1
2
3
4
5
6
# 位置一:Dtri[i,j] = gamma^(i-j) for j<=i, else 0
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)。原因是 exp2log2 都有单指令硬件实现,而 pow 通常展开成多条指令;log2γ\log_2 \gamma 是循环不变量,编译器会提到循环外。这个改写在下一级会变成必需–如 §2.5 所述,逐 token 门控的 gt=logαtg_t = \log \alpha_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)
# 位置三:整块衰减 gamma^BC(对上一轮累积的状态)
for i, j in T.Parallel(dim, dim):
S_f[i, j] *= gamma_pow_bc # 编译期常量 gamma**BC
# 位置二:K 的第 m 行乘 gamma^(BC-1-m) 后再做 gemm
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 之前。递推式是 Sc=γBCSc1+Δc1S_c = \gamma^{BC} S_{c-1} + \Delta_{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
# ② 跨块:O = (Q_c ⊙ gamma^(i+1)) @ S_prev
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] # 位置三的 query 侧
T.clear(acc_o)
T.gemm(Q_s, S_s, acc_o)

# ③ 块内:注意此处要用未缩放的 Q,故需重新载入
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},而块内那一项用的是未缩放的 QcQ_c。这是本级新增的一处陷阱:上一级 Q_s 在两步之间可以直接复用。

原地缩放省了一个 buffer,代价是必须重载;若寄存器与 shared memory 有余量,另开一个 Q_scaled 更安全。这个取舍在 BC=128BC = 128 时倾向于原地缩放。

5.4 与上一级 kernel 的改动汇总

位置 上一级 本级 新增开销
权重表 Dtri[BC, BC] f32 fragment BC2BC^2 个 f32 寄存器
状态累加 T.gemm 直接累加 S_f *= γ^BC 再 gemm D2D^2 次乘法 / 轮
K 写入状态 原样 逐行乘 γBC1m\gamma^{BC-1-m'} BC×DBC \times D 次乘法 / 轮
Q 读取状态 原样 逐行乘 γi+1\gamma^{i'+1} BC×DBC \times D 次乘法
块内掩码 if_then_else 置 0 逐元素乘 Dtri BC2BC^2 次乘法
Q 复用 步骤②③共用 步骤③必须重载 一次 shared 写入

DtriBC2BC^2 个 f32 寄存器是本级最主要的资源代价。BC=64BC = 64 时是 4096 个 f32、128 线程下每线程 32 个;BC=128BC = 128 时升到 16384 个、每线程 128 个–与 A 加起来已经接近寄存器上限。这是 BC=128BC = 128 在本级比上一级更难用的原因,也是因式分解唯一真正有吸引力的地方(它不需要 BC2BC^2 权重表)。但 §3.1 的 NaN 结论表明这个吸引力不成立。


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

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

L=ln103lnγL = \frac{\ln 10^{-3}}{\ln \gamma}

γ\gamma γ64\gamma^{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\gamma = 0.9 的有效记忆只有 66 token,连一个 chunk(BC=64BC = 64)都刚刚覆盖,长程依赖完全丢失。γ0.99\gamma \ge 0.99 恰好落在 §2.4 与 §3 表格中因式分解仍然安全的区间–这解释了为什么部分实现敢做这个分解:它们隐含假设了门控接近 1。

这个假设在标量门控下大致成立,但在第三级逐通道门控下会失效。逐通道意味着每个 key 通道有独立的 γd\gamma_d,训练中总会有部分通道学到较小的值以实现快速遗忘。此时「所有通道都接近 1」不再成立,因式分解的溢出从例外变成常态。Kimi K3 的处理方式是给门控加下界,这部分留到第三级展开。

一句话总结γ\gamma 的安全区间(接近 1)与有用区间(提供实际遗忘能力)方向相反,标量门控下两者尚可兼顾,逐通道门控下这个矛盾会暴露。


7. 数值验证

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

  1. 四层参考互验(fp64):A/B/C/D 两两对照,误差应在 101510^{-15} 量级;
  2. 退化检验γ=1\gamma = 1 时参考 B 应精确回到上一级实现–这条能捕获三处权重的指数偏移错误;
  3. fp16 失效点复现γ0.8\gamma \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 上完成(见 §2.3、§3.1 与 §4.4 表格)。第 4、5 层需要 CUDA 设备,待真卡跑通后单独补充实测数据,此处不做性能推测。

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


8. 总结

  1. 标量衰减不改变分块恒等式的结构,只在三个位置插入指数权重:块内下三角 γij\gamma^{i'-j'}、写入状态时的块尾对齐 γBC1m\gamma^{BC-1-m'}、跨块递推 γBC\gamma^{BC} 与 query 侧 γi+1\gamma^{i'+1}。四个权重全部 1\le 1,这是本级数值安全的根本原因。
  2. 累积衰减积 Λj=ijαi\Lambda_j = \prod_{i \le j} \alpha_i 是统一两种形式的代数工具。它把区间连乘 i=j+1tαi\prod_{i=j+1}^{t} \alpha_i 化归为前缀量之比 Λt/Λj\Lambda_t / \Lambda_j,于是递归的向量形式与并行的矩阵形式 O=((QK)Γ)VO = ((QK^\top) \odot \Gamma)V 可以相互转换(即 Mamba2 所谓的状态空间对偶性,实测相对 L2 2.03×10162.03 \times 10^{-16})。常数 γ\gammaΛj=γj\Lambda_j = \gamma^j 的特例,此时 Γij=γij\Gamma_{ij} = \gamma^{i-j}
  3. 论文的衰减矩阵 Γij=Λi/Λj\Gamma_{ij} = \Lambda_i/\Lambda_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{q}k\overrightarrow{k}S\overrightarrow{S} 把「衰减到首 / 末位置」直接编码在箭头方向上,基准点选得当就能保证所有指数非正。
  4. 真正的陷阱是把比值因式分解成 Λi(1/Λj)\Lambda_i \cdot (1/\Lambda_j)。这个分解能把两步合成单次 GEMM 并免去 BC2BC^2 权重表,但要求物化 1/Λj1/\Lambda_j,把中间量从 (0,1](0,1] 推到 [1, 1/ΛBC][1,\ 1/\Lambda_{BC}]。实测(BC=64BC=64,fp16 存储加 fp32 累加):γ0.9\gamma \ge 0.9 时两者精度相当(4.14.6×1044.1\text{--}4.6 \times 10^{-4}),γ0.8\gamma \le 0.8 时分解写法产生 NaNγ=0.8\gamma = 0.8γ63=1.27×106\gamma^{-63} = 1.27 \times 10^{6},已远超 fp16 上限 65504。data-dependent αt[0.8, 0.95]\alpha_t \in [0.8,\ 0.95]BC=128BC = 1281/Λ1/\Lambda 最大达 3.40×1083.40 \times 10^{8},同样溢出。结论是保留论文的比值形式,不做分解。
  5. 累积积本身也必须在 log 域计算。fp16 直接 cumprodBC=128BC = 128α[0.5, 0.9]\alpha \in [0.5,\ 0.9] 时已完全下溢为 0,fp32 在 BC=512BC = 512 时进入非正规数区间;logΛj=ijlogαi\log \Lambda_j = \sum_{i \le j} \log \alpha_i 是线性增长的负数,表示范围安全。这就是门控全程存 log 值、用 exp2 还原的原因。
  6. 退化检验是本级新增的关键验证手段:令 γ=1\gamma = 1 应精确回到上一级实现(实测相对 L2 1.62×10161.62 \times 10^{-16})。但它无法单独定案–三处权重中任何一处的指数偏移在 γ=1\gamma = 1 时都不可见,必须同时做 γ=0.9\gamma = 0.9 的四层参考互验(实测 1.822.03×10161.82\text{--}2.03 \times 10^{-16})。
  7. 两处新增实现陷阱S_f *= γ^BC 必须在 T.gemm 之前,否则本块贡献被多衰减一次;步骤②把 Q_s 原地乘了 γi+1\gamma^{i'+1},步骤③用的是未缩放 QcQ_c,必须重新载入。
  8. 有效记忆长度暴露了一个内在矛盾γ\gamma 的数值安全区间(接近 1)与实际有用区间(提供遗忘能力)方向相反。γ=0.9\gamma = 0.9 的有效记忆仅 66 token,连一个 chunk 都刚覆盖;而 γ0.99\gamma \ge 0.99 才落在因式分解仍然安全的区间。标量门控下两者尚可兼顾。

下一篇把 γ\gamma 升级为逐 token 门控 gtg_t:衰减不再是编译期常量,需要在 log 域做 cumsum 得到 §2.1 那个累积积 Λj\Lambda_j,权重表从预计算变成运行时构造,exp2 从优化手段变成必需。此时 §6 的矛盾会真正暴露–每个通道独立的门控使「所有 γ\gamma 接近 1」的假设失效,Λ\Lambda 从标量变成向量,1/Λ1/\Lambda 的溢出从例外变成常态。

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


参考

  • Gated Delta Networks(GDN):arXiv:2412.06464,ICLR 2025。§2.1 给出累积衰减积 γ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} 箭头记号
  • 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 实测数据待补。