TileLang 实战:KDA 从零到一–标量衰减
上一篇实现了不带任何遗忘机制的 chunked 线性注意力,状态单调累加。本篇引入第一个衰减因子:把递推式改为 S t = α S t − 1 + v t k t ⊺ \mathbf{S}_t = \alpha \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^\intercal S t = α S t − 1 + v t k t ⊺ ,α \alpha α 是一个标量常数。
这个衰减项来自哪里?Vanilla 线性注意力(S t = S t − 1 + v t k t ⊺ \mathbf{S}_t = \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^\intercal S t = S t − 1 + v t k t ⊺ )在语言建模上远不如 Transformer,为此 Mamba2(Dao & Gu, 2024a)引入了一个数据依赖的逐步遗忘门 α t ∈ ( 0 , 1 ) \alpha_t \in (0, 1) α t ∈ ( 0 , 1 ) 来有选择地丢弃历史信息。本文取它的常数特例 α t ≡ α \alpha_t \equiv \alpha α t ≡ α –这一档对应的正是 RetNet 与 Lightning-Attention,详见 §1.0。
上一篇见《TileLang 实战:KDA 从零到一–Chunked 线性注意力 》,本文沿用其参考实现的验证框架。
0. 符号约定:与 GDN 论文对齐
本文全程采用 GDN(arXiv:2412.06464v3)§2.1 与式 (1)(2) 的记号。这里先把几个容易混淆的符号定下来。
符号
含义
说明
α t ∈ ( 0 , 1 ) \alpha_t \in (0,1) α t ∈ ( 0 , 1 )
单步 衰减系数
本文用常数,记作 α t ≡ α \alpha_t \equiv \alpha α t ≡ α
γ j = ∏ i = 1 j α i \gamma_j = \prod_{i=1}^{j}\alpha_i γ j = ∏ i = 1 j α i
累积 衰减积
常数情形下 γ j = α j \gamma_j = \alpha^{\,j} γ j = α j
Γ i j = γ i / γ j \Gamma_{ij} = \gamma_i/\gamma_j Γ ij = γ i / γ j
衰减感知因果掩码
i ≥ j i \ge j i ≥ j 时有值,否则 0
C C C
chunk 长度
代码里对应 BC / blk
[ t ] [t] [ t ]
chunk 序号
Q [ t ] \mathbf{Q}_{[t]} Q [ t ] 即第 t t t 块
r ∈ [ 1 , C ] r \in [1, C] r ∈ [ 1 , C ]
chunk 内 位置(1-based)
q [ t ] r : = q t C + r \bm{q}_{[t]}^r := \bm{q}_{tC+r} q [ t ] r := q tC + r
两个必须说清的点:
一、γ \gamma γ 是累积积,不是单步衰减。 这是阅读 GDN 时最容易混淆的一处–很多二次资料把 γ \gamma γ 当成逐步遗忘系数,但论文里逐步系数是 α t \alpha_t α t ,γ j \gamma_j γ j 是它们的前缀积。
二、累积积按 chunk 重置,不是从序列起点算。 论文式 (1) 的脚注明确写了 γ [ t ] j = ∏ j = t C + 1 t C + j α j \gamma_{[t]}^{j} = \prod_{j=tC+1}^{tC+j}\alpha_j γ [ t ] j = ∏ j = tC + 1 tC + j α j ,并自认“略微滥用了 γ \gamma γ 的记号”–每个 chunk 从自己的第一个位置重新开始累乘 。这不是细节:若从序列起点算,γ j \gamma_j γ j 会随 j j j 单调下溢到零(§2.5 有实测),而按 chunk 重置后指数永远不超过 C C C ,这才是分块算法数值可控的前提。
1. 递推式与分块重写
1.0 Mamba2 的遗忘门:衰减项的来处
上一篇实现的是 vanilla 线性注意力(Katharopoulos et al., 2020),状态只增不减:
S t = S t − 1 + v t k t ⊺ ∈ R d v × d k , o t = S t q t ∈ R d v \mathbf{S}_t = \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} \in \mathbb{R}^{d_v \times d_k},
\qquad \bm{o}_t = \mathbf{S}_t\bm{q}_t \in \mathbb{R}^{d_v}
S t = S t − 1 + v t k t ⊺ ∈ R d v × d k , o t = S t q t ∈ R d v
它在语言建模上明显弱于 Transformer。GDN 论文 §2.1 给出的诊断很直接:缺少遗忘历史信息的手段 。以 Mamba2(Dao & Gu, 2024a)为例,补救方式是在状态上乘一个逐步衰减项(论文原话是 “up to specific parameterization”,即忽略具体参数化细节):
S t = α t S t − 1 + v t k t ⊺ , o t = S t q t \mathbf{S}_t = \alpha_t \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal},
\qquad \bm{o}_t = \mathbf{S}_t\bm{q}_t
S t = α t S t − 1 + v t k t ⊺ , o t = S t q t
其中 α t ∈ ( 0 , 1 ) \alpha_t \in (0, 1) α t ∈ ( 0 , 1 ) 是随 t t t 变化的、数据依赖的标量衰减项 (a data-dependent scalar-valued decay term that varies with t t t )。这一个乘法带来两件事:
有界性 。无衰减时 ∥ S t ∥ \|\mathbf{S}_t\| ∥ S t ∥ 随 t t t 无界增长(每步加一个外积);有了 α t < 1 \alpha_t < 1 α t < 1 ,状态范数被压在一个稳态附近。
选择性 。α t → 0 \alpha_t \to 0 α t → 0 可以快速清空状态(适合上下文切换),α t → 1 \alpha_t \to 1 α t → 1 则保持记忆。Mamba 把这类机制称为 selective mechanism ,本质是 gated RNN 里的遗忘门在矩阵值状态上的推广。
这个递推结构不是 Mamba2 独有的,同样出现在 Gated RFA、xLSTM、Gated RetNet 中。而当 α t \alpha_t α t 与数据无关、退化为常数时,该形式就是 RetNet 与 Lightning-Attention。
本文取的正是这个常数情形:
形态
衰减项
对应架构
无衰减
无
vanilla 线性注意力(上一篇)
常数标量
α t ≡ α \alpha_t \equiv \alpha α t ≡ α
RetNet / Lightning-Attention(本文)
数据依赖标量
α t = f ( x t ) \alpha_t = f(x_t) α t = f ( x t )
Mamba2 / GLA / Gated RetNet
本文取常数是为了简化:α \alpha α 作为编译期常量,权重表可以预计算,kernel 改动最小。
1.1 本文的递推式与显式求和
本文在上一篇的基础上给状态加入一个遗忘系数 α ∈ ( 0 , 1 ] \alpha \in (0, 1] α ∈ ( 0 , 1 ] :
S t = α S t − 1 + v t k t ⊺ ∈ R d v × d k , o t = S t q t ∈ R d v \mathbf{S}_t = \alpha \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} \in \mathbb{R}^{d_v \times d_k},
\qquad \bm{o}_t = \mathbf{S}_t\bm{q}_t \in \mathbb{R}^{d_v}
S t = α S t − 1 + v t k t ⊺ ∈ R d v × d k , o t = S t q t ∈ R d v
展开成显式求和,注意每个 v i k i ⊺ \bm{v}_i\bm{k}_i^{\intercal} v i k i ⊺ 被后续每一步各乘一次 α \alpha α ,从 i i i 到 t t t 共乘 t − i t - i t − i 次:
S t = ∑ i ≤ t α t − i v i k i ⊺ , o t = ∑ i ≤ t α t − i v i ( k i ⊺ q t ) \mathbf{S}_t = \sum_{i \le t} \alpha^{\,t-i}\, \bm{v}_i\bm{k}_i^{\intercal}, \qquad
\bm{o}_t = \sum_{i \le t} \alpha^{\,t-i}\, \bm{v}_i\,(\bm{k}_i^{\intercal}\bm{q}_t)
S t = i ≤ t ∑ α t − i v i k i ⊺ , o t = i ≤ t ∑ α t − i v i ( k i ⊺ q t )
用累积积写就是论文 §2.1 的形式(γ j = α j \gamma_j = \alpha^j γ j = α j ,所以 γ t / γ i = α t − i \gamma_t/\gamma_i = \alpha^{t-i} γ t / γ i = α t − i ):
o t = ∑ i ≤ t γ t γ i v i ( k i ⊺ q t ) \bm{o}_t = \sum_{i \le t} \frac{\gamma_t}{\gamma_i}\, \bm{v}_i\,(\bm{k}_i^{\intercal}\bm{q}_t)
o t = i ≤ t ∑ γ i γ t v i ( k i ⊺ q t )
对比无衰减时的 o t = ∑ i ≤ t v i ( k i ⊺ q t ) \bm{o}_t = \sum_{i \le t} \bm{v}_i(\bm{k}_i^{\intercal}\bm{q}_t) o t = ∑ i ≤ t v i ( k i ⊺ q t ) ,唯一变化是每一项多了权重 α t − i \alpha^{t-i} α t − i –距离越远权重越小,这就是「衰减」的含义。α = 1 \alpha = 1 α = 1 时权重恒为 1,退化为不遗忘。
1.2 三处需要插入衰减权重的位置
沿用上一篇的分块框架,序列按 C C C 切分为 N C NC N C 个 chunk。把求和拆成跨块与块内两部分,衰减权重会分别落到三个位置。沿用论文的 chunk 内局部下标 r ∈ [ 1 , C ] r \in [1, C] r ∈ [ 1 , C ] :
位置一–块内下三角(论文的 Γ [ t ] \Gamma_{[t]} Γ [ t ] ) 。同块内 j ≤ i j \le i j ≤ i ,相对距离就是局部下标之差:
( Γ [ t ] ) i j = γ [ t ] i γ [ t ] j = { α i − j j ≤ i 0 j > i (\Gamma_{[t]})_{ij} = \frac{\gamma_{[t]}^{\,i}}{\gamma_{[t]}^{\,j}}
= \begin{cases} \alpha^{\,i-j} & j \le i \\ 0 & j > i \end{cases}
( Γ [ t ] ) ij = γ [ t ] j γ [ t ] i = { α i − j 0 j ≤ i j > i
上一篇这里是 0/1 因果掩码 M \mathbf{M} M ,本文变成衰减感知掩码 Γ \Gamma Γ 。注意权重全部落在 ( 0 , 1 ] (0, 1] ( 0 , 1 ] 区间 –对角线是 α 0 = 1 \alpha^0 = 1 α 0 = 1 ,左下角最小值是 α C − 1 \alpha^{C-1} α C − 1 。
位置二–每块写入状态时的块尾对齐(论文的 k → \overrightarrow{\bm{k}} k ) 。跨块状态需要统一的时间基准,取所属 chunk 的末尾 。第 r r r 个 token 的贡献衰减到块末要乘 γ [ t ] C / γ [ t ] r = α C − r \gamma_{[t]}^{C}/\gamma_{[t]}^{r} = \alpha^{\,C-r} γ [ t ] C / γ [ t ] r = α C − r :
k [ t ] r → = γ [ t ] C γ [ t ] r k [ t ] r , Δ [ t ] = V [ t ] ⊺ K [ t ] → ∈ R d v × d k \overrightarrow{\bm{k}_{[t]}^{r}} = \frac{\gamma_{[t]}^{C}}{\gamma_{[t]}^{r}}\,\bm{k}_{[t]}^{r},
\qquad
\Delta_{[t]} = \mathbf{V}_{[t]}^{\intercal}\, \overrightarrow{\mathbf{K}_{[t]}} \in \mathbb{R}^{d_v \times d_k}
k [ t ] r = γ [ t ] r γ [ t ] C k [ t ] r , Δ [ t ] = V [ t ] ⊺ K [ t ] ∈ R d v × d k
这个权重从哪来?用 C = 4 C = 4 C = 4 手推一遍最清楚。块内每走一格执行一次递推,S [ t ] \mathbf{S}_{[t]} S [ t ] 为进块状态、S [ t + 1 ] \mathbf{S}_{[t+1]} S [ t + 1 ] 为出块状态:
S r = α r S r − 1 + v r k r ⊺ , S 0 = S [ t ] , S C = S [ t + 1 ] \mathbf{S}_r = \alpha_r \mathbf{S}_{r-1} + \bm{v}_r\bm{k}_r^{\intercal},
\qquad \mathbf{S}_0 = \mathbf{S}_{[t]},\quad \mathbf{S}_C = \mathbf{S}_{[t+1]}
S r = α r S r − 1 + v r k r ⊺ , S 0 = S [ t ] , S C = S [ t + 1 ]
正向走一遍没什么可看的。有意思的是站在 S 4 \mathbf{S}_4 S 4 这端逐层倒代换 ,把中间状态一个个拆掉:
S 4 = α 4 S 3 + v 4 k 4 ⊺ = α 4 α 3 S 2 + α 4 v 3 k 3 ⊺ + v 4 k 4 ⊺ = α 4 α 3 α 2 S 1 + α 4 α 3 v 2 k 2 ⊺ + α 4 v 3 k 3 ⊺ + v 4 k 4 ⊺ = α 1 α 2 α 3 α 4 S [ t ] + α 2 α 3 α 4 v 1 k 1 ⊺ + α 3 α 4 v 2 k 2 ⊺ + α 4 v 3 k 3 ⊺ + v 4 k 4 ⊺ \begin{aligned}
\mathbf{S}_4 &= \alpha_4\mathbf{S}_3 + \bm{v}_4\bm{k}_4^{\intercal} \\
&= \alpha_4\alpha_3\mathbf{S}_2 + \alpha_4\bm{v}_3\bm{k}_3^{\intercal} + \bm{v}_4\bm{k}_4^{\intercal} \\
&= \alpha_4\alpha_3\alpha_2\mathbf{S}_1 + \alpha_4\alpha_3\bm{v}_2\bm{k}_2^{\intercal} + \alpha_4\bm{v}_3\bm{k}_3^{\intercal} + \bm{v}_4\bm{k}_4^{\intercal} \\
&= \alpha_1\alpha_2\alpha_3\alpha_4\,\mathbf{S}_{[t]}
+ \alpha_2\alpha_3\alpha_4\,\bm{v}_1\bm{k}_1^{\intercal}
+ \alpha_3\alpha_4\,\bm{v}_2\bm{k}_2^{\intercal}
+ \alpha_4\,\bm{v}_3\bm{k}_3^{\intercal}
+ \bm{v}_4\bm{k}_4^{\intercal}
\end{aligned}
S 4 = α 4 S 3 + v 4 k 4 ⊺ = α 4 α 3 S 2 + α 4 v 3 k 3 ⊺ + v 4 k 4 ⊺ = α 4 α 3 α 2 S 1 + α 4 α 3 v 2 k 2 ⊺ + α 4 v 3 k 3 ⊺ + v 4 k 4 ⊺ = α 1 α 2 α 3 α 4 S [ t ] + α 2 α 3 α 4 v 1 k 1 ⊺ + α 3 α 4 v 2 k 2 ⊺ + α 4 v 3 k 3 ⊺ + v 4 k 4 ⊺
盯着最后一行的系数:v 1 \bm{v}_1 v 1 带的是 α 2 α 3 α 4 \alpha_2\alpha_3\alpha_4 α 2 α 3 α 4 –注意没有 α 1 \alpha_1 α 1 ,因为它自己就是第 1 格写入的,写完后一路经受的是"身后"三道门;历史状态 S [ t ] \mathbf{S}_{[t]} S [ t ] 则带满 α 1 α 2 α 3 α 4 \alpha_1\alpha_2\alpha_3\alpha_4 α 1 α 2 α 3 α 4 。没有任何一项的因子个数取决于它的绝对位置,全部取决于「离第 C C C 格还差几步」。
用累积积改写,公共前缀立刻现形:α 2 α 3 α 4 = γ 4 / γ 1 \alpha_2\alpha_3\alpha_4 = \gamma_4/\gamma_1 α 2 α 3 α 4 = γ 4 / γ 1 ,α 3 α 4 = γ 4 / γ 2 \alpha_3\alpha_4 = \gamma_4/\gamma_2 α 3 α 4 = γ 4 / γ 2 ,α 4 = γ 4 / γ 3 \alpha_4 = \gamma_4/\gamma_3 α 4 = γ 4 / γ 3 ,而 v 4 \bm{v}_4 v 4 那项是 γ 4 / γ 4 = 1 \gamma_4/\gamma_4 = 1 γ 4 / γ 4 = 1 。于是打包收口:
S [ t + 1 ] = γ [ t ] C S [ t ] + ∑ r = 1 C γ [ t ] C γ [ t ] r v r k r ⊺ \mathbf{S}_{[t+1]} = \gamma_{[t]}^{C}\,\mathbf{S}_{[t]} + \sum_{r=1}^{C} \frac{\gamma_{[t]}^{C}}{\gamma_{[t]}^{r}}\, \bm{v}_r\bm{k}_r^{\intercal}
S [ t + 1 ] = γ [ t ] C S [ t ] + r = 1 ∑ C γ [ t ] r γ [ t ] C v r k r ⊺
这就是 k [ t ] r → \overrightarrow{\bm{k}_{[t]}^{r}} k [ t ] r 的全部内容,也就是论文说的「decaying each vector to the last position」。
位置三–跨块状态递推与 query 侧乘累积衰减积(论文的 S → \overrightarrow{\mathbf{S}} S 与 q ← \overleftarrow{\bm{q}} q ) 。相邻块之间隔了整块 C C C 步,因此状态递推是:
S [ t ] → = γ [ t ] C S [ t ] = α C S [ t ] , S [ t + 1 ] = S [ t ] → + V [ t ] ⊺ K [ t ] → \overrightarrow{\mathbf{S}_{[t]}} = \gamma_{[t]}^{C}\,\mathbf{S}_{[t]} = \alpha^{C}\mathbf{S}_{[t]},
\qquad
\mathbf{S}_{[t+1]} = \overrightarrow{\mathbf{S}_{[t]}} + \mathbf{V}_{[t]}^{\intercal}\,\overrightarrow{\mathbf{K}_{[t]}}
S [ t ] = γ [ t ] C S [ t ] = α C S [ t ] , S [ t + 1 ] = S [ t ] + V [ t ] ⊺ K [ t ]
而第 r r r 个 token 读取这个状态时,要衰减到本块首位置 的基准上,乘 γ [ t ] r = α r \gamma_{[t]}^{r} = \alpha^{r} γ [ t ] r = α r :
q [ t ] r ← = γ [ t ] r q [ t ] r , O [ t ] = Q [ t ] ← S [ t ] ⊺ + ( Q [ t ] K [ t ] ⊺ ⊙ Γ [ t ] ) V [ t ] ∈ R C × d v \overleftarrow{\bm{q}_{[t]}^{r}} = \gamma_{[t]}^{r}\,\bm{q}_{[t]}^{r},
\qquad
\mathbf{O}_{[t]} = \overleftarrow{\mathbf{Q}_{[t]}}\,\mathbf{S}_{[t]}^{\intercal}
+ \big(\mathbf{Q}_{[t]}\mathbf{K}_{[t]}^{\intercal} \odot \Gamma_{[t]}\big)\mathbf{V}_{[t]}
\in \mathbb{R}^{C \times d_v}
q [ t ] r = γ [ t ] r q [ t ] r , O [ t ] = Q [ t ] S [ t ] ⊺ + ( Q [ t ] K [ t ] ⊺ ⊙ Γ [ t ] ) V [ t ] ∈ R C × d v
这正是论文式 (1)。三个位置的权重汇总–注意每个权重本质上都是门控的累积连乘 ,常数情形下才收成 α \alpha α 的幂:
位置
论文记号
一般形式(累积积)
常数特例
取值范围
作用对象
块内下三角
( Γ [ t ] ) i j (\Gamma_{[t]})_{ij} ( Γ [ t ] ) ij
γ [ t ] i γ [ t ] j = ∏ r = j + 1 i α r \dfrac{\gamma_{[t]}^{\,i}}{\gamma_{[t]}^{\,j}} = \prod_{r=j+1}^{i}\alpha_r γ [ t ] j γ [ t ] i = ∏ r = j + 1 i α r
α i − j \alpha^{\,i-j} α i − j
[ α C − 1 , 1 ] [\alpha^{C-1},\ 1] [ α C − 1 , 1 ]
C × C C \times C C × C 分数矩阵,逐元素
写入状态
k [ t ] r → \overrightarrow{\bm{k}_{[t]}^{r}} k [ t ] r
γ [ t ] C γ [ t ] r = ∏ u = r + 1 C α u \dfrac{\gamma_{[t]}^{C}}{\gamma_{[t]}^{r}} = \prod_{u=r+1}^{C}\alpha_u γ [ t ] r γ [ t ] C = ∏ u = r + 1 C α u
α C − r \alpha^{\,C-r} α C − r
[ α C − 1 , 1 ] [\alpha^{C-1},\ 1] [ α C − 1 , 1 ]
K [ t ] \mathbf{K}_{[t]} K [ t ] 的行,逐行加权
跨块递推
S [ t ] → \overrightarrow{\mathbf{S}_{[t]}} S [ t ]
γ [ t ] C = ∏ u = 1 C α u \gamma_{[t]}^{C} = \prod_{u=1}^{C}\alpha_u γ [ t ] C = ∏ u = 1 C α u
α C \alpha^{C} α C
标量
整个状态矩阵
读取状态
q [ t ] r ← \overleftarrow{\bm{q}_{[t]}^{r}} q [ t ] r
γ [ t ] r = ∏ u = 1 r α u \gamma_{[t]}^{r} = \prod_{u=1}^{r}\alpha_u γ [ t ] r = ∏ u = 1 r α u
α r \alpha^{r} α r
[ α C , α ] [\alpha^{C},\ \alpha] [ α C , α ]
Q [ t ] \mathbf{Q}_{[t]} Q [ t ] 的行,逐行加权
四个权重全部 ≤ 1 \le 1 ≤ 1 ,指数均为非正。因果约束 i ≥ j i \ge j i ≥ j 使 Γ i j = ∏ r = j + 1 i α r ≤ 1 \Gamma_{ij} = \prod_{r=j+1}^{i}\alpha_r \le 1 Γ ij = ∏ r = j + 1 i α r ≤ 1 ,四个权重同理–论文的写法全程不需要物化任何大于 1 的量 ,这是 §3 讨论的前提。
2. 累积衰减积:从连乘到矩阵并行形式
上一节的推导是直接展开求和得到的,但衰减机制有一个更本质的表述方式,GDN 论文用它统一了递归形式与并行形式。
2.1 累积衰减积的定义
把衰减退回 §1.0 那个一般形式–依赖数据的 α t ∈ ( 0 , 1 ) \alpha_t \in (0, 1) α t ∈ ( 0 , 1 ) ,每个时刻的遗忘强度由输入决定(本文的常数 α \alpha α 是 α t ≡ α \alpha_t \equiv \alpha α t ≡ α 的特例)。定义累积衰减积 :
γ j = ∏ i = 1 j α i \gamma_j = \prod_{i=1}^{j} \alpha_i
γ j = i = 1 ∏ j α i
γ j \gamma_j γ j 的含义是从起点衰减到第 j j j 步的总折扣。有了它,递推式的展开可以写得非常紧凑。展开 S t = α t S t − 1 + v t k t ⊺ \mathbf{S}_t = \alpha_t \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} S t = α t S t − 1 + v t k t ⊺ :
S t = ∑ i ≤ t ( ∏ r = i + 1 t α r ) v i k i ⊺ = ∑ i ≤ t γ t γ i v i k i ⊺ \mathbf{S}_t = \sum_{i \le t} \Big( \prod_{r=i+1}^{t} \alpha_r \Big) \bm{v}_i\bm{k}_i^{\intercal}
= \sum_{i \le t} \frac{\gamma_t}{\gamma_i}\, \bm{v}_i\bm{k}_i^{\intercal}
S t = i ≤ t ∑ ( r = i + 1 ∏ t α r ) v i k i ⊺ = i ≤ t ∑ γ i γ t v i k i ⊺
中间那个连乘 ∏ r = i + 1 t α r \prod_{r=i+1}^{t} \alpha_r ∏ r = i + 1 t α r 正好是两个累积积的比值 γ t / γ i \gamma_t / \gamma_i γ t / γ i –这是累积积定义的全部价值:把「从 i i i 到 t t t 的区间连乘」化归为「两个前缀量之比」 ,于是任意区间的衰减都可以由一个前缀数组 O ( 1 ) O(1) O ( 1 ) 查得,不必对每个 ( i , t ) (i, t) ( i , t ) 对重新连乘。
上面写的是全序列版本。到了分块算法里,累积积要按 chunk 重置 ,即论文式 (1) 脚注的 γ [ t ] j = ∏ j = t C + 1 t C + j α j \gamma_{[t]}^{\,j} = \prod_{j=tC+1}^{tC+j}\alpha_j γ [ t ] j = ∏ j = tC + 1 tC + j α j 。这一步不是为了好看–全序列累乘的 γ j \gamma_j γ j 会随 j j j 单调下溢到零(§2.5 有实测:N = 512 N = 512 N = 512 、α ∈ [ 0.5 , 0.9 ] \alpha \in [0.5, 0.9] α ∈ [ 0.5 , 0.9 ] 时 fp32 已进入非正规数),而按 chunk 重置后指数永远不超过 C C C 。后文出现 γ \gamma γ 时默认指 chunk 内的版本。
2.2 两种等价形式
代入 o t = S t q t \bm{o}_t = \mathbf{S}_t\bm{q}_t o t = S t q t ,同一个结果可以写成两种形式(论文 §2.1):
向量形式(vector form) –逐时刻递归,对应推理阶段:
S t = α t S t − 1 + v t k t ⊺ , o t = S t q t = ∑ i ≤ t v i ( γ t γ i k i ⊺ q t ) \mathbf{S}_t = \alpha_t \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal},
\qquad \bm{o}_t = \mathbf{S}_t\bm{q}_t
= \sum_{i \le t} \bm{v}_i\Big(\frac{\gamma_t}{\gamma_i}\,\bm{k}_i^{\intercal}\bm{q}_t\Big)
S t = α t S t − 1 + v t k t ⊺ , o t = S t q t = i ≤ t ∑ v i ( γ i γ t k i ⊺ q t )
矩阵并行形式(matrix parallel form) –整块一次算出,对应训练与 prefill:
O = ( ( Q K ⊺ ) ⊙ Γ ) V , Γ i j = { γ i γ j i ≥ j 0 i < j \mathbf{O} = \big( (\mathbf{Q} \mathbf{K}^{\intercal}) \odot \Gamma \big) \mathbf{V}, \qquad
\Gamma_{ij} = \begin{cases} \dfrac{\gamma_i}{\gamma_j} & i \ge j \\[2mm] 0 & i < j \end{cases}
O = ( ( Q K ⊺ ) ⊙ Γ ) V , Γ ij = ⎩ ⎨ ⎧ γ j γ i 0 i ≥ j i < j
Γ \Gamma Γ 是一个衰减感知因果掩码(decay-aware causal mask) –把上一篇的 0/1 因果掩码 M \mathbf{M} M 换成了衰减比值。验证第 ( i , j ) (i,j) ( i , j ) 元素:
[ ( Q K ⊺ ) ⊙ Γ ] i j = γ i γ j ( k j ⊺ q i ) ( j ≤ i ) \big[(\mathbf{Q} \mathbf{K}^{\intercal}) \odot \Gamma\big]_{ij} = \frac{\gamma_i}{\gamma_j} (\bm{k}_j^{\intercal}\bm{q}_i) \quad (j \le i)
[ ( Q K ⊺ ) ⊙ Γ ] ij = γ j γ i ( k j ⊺ q i ) ( j ≤ i )
与 §2.1 展开式逐项一致。这个形式的价值是把逐 token 的递归变成一次稠密 GEMM 加一次逐元素乘 ,完全并行,这正是「parallel within each chunk」的含义。递归形式与并行形式的这种等价在 Mamba2 中被称为状态空间对偶性(state space duality, SSD) 。
Γ \Gamma Γ 的双重身份:一个哈达玛积承载两重语义
Γ \Gamma Γ 常被笼统理解为“一个带衰减的掩码”,但它实际上是两个正交语义的乘积 ,只是恰好能合并成一张表:
Γ = M ⏟ 因果性:能不能看 ⊙ D ⏟ 遗忘门:看得多清 , M i j = { 1 i ≥ j 0 i < j , D i j = γ i γ j \Gamma = \underbrace{\mathbf{M}}_{\text{因果性:能不能看}} \odot \underbrace{\mathbf{D}}_{\text{遗忘门:看得多清}},
\qquad
\mathbf{M}_{ij} = \begin{cases}1 & i \ge j\\ 0 & i<j\end{cases},
\quad
\mathbf{D}_{ij} = \frac{\gamma_i}{\gamma_j}
Γ = 因果性:能不能看 M ⊙ 遗忘门:看得多清 D , M ij = { 1 0 i ≥ j i < j , D ij = γ j γ i
M \mathbf{M} M 是离散的、与数据无关的 结构约束–token i i i 不能看到未来的 j > i j > i j > i 。这是 causal 语言模型的硬性要求,α \alpha α 取什么值都不影响它。
D \mathbf{D} D 是连续的、由门控决定的 权重–i i i 与 j j j 相距越远,γ i / γ j \gamma_i/\gamma_j γ i / γ j 越小。这才是遗忘门起作用的地方。
将该矩阵可视化后,结构十分清晰:下三角,且每个值只由 token 序号差 i − j i-j i − j 决定 。
图中三点值得对着公式确认:
对角线恒为 α 0 = 1 \alpha^0 = 1 α 0 = 1 –自己看自己,零衰减。本文的 tril \operatorname{tril} tril 含对角(§1.2 位置一的 j ≤ i j \le i j ≤ i ),与图一致。
同一行自右向左指数缩小 –同行内 i i i 固定,j j j 越小则间隔 i − j i-j i − j 越大、权重越小。所以「衰减」在矩阵上表现为沿对角线方向的等值带 :i − j i-j i − j 相同的格子取值相同。
上三角 j > i j > i j > i 整块置 0 –那是还没发生的 token。这就是「下三角」的全部含义。
面板 ② 用 α = 0.5 \alpha = 0.5 α = 0.5 、C = 4 C = 4 C = 4 给了可验算的数字:γ = ( 0.5 , 0.25 , 0.125 , 0.0625 ) \gamma = (0.5,\ 0.25,\ 0.125,\ 0.0625) γ = ( 0.5 , 0.25 , 0.125 , 0.0625 ) ,格 ( 4 , 1 ) (4,1) ( 4 , 1 ) 即 γ 4 / γ 1 = 0.0625 / 0.5 = 0.125 = α 3 \gamma^4/\gamma^1 = 0.0625/0.5 = 0.125 = \alpha^3 γ 4 / γ 1 = 0.0625/0.5 = 0.125 = α 3 。整张表我用 fp64 复核过,与 α i − j \alpha^{i-j} α i − j 逐格一致、对角恒为 1。
图中「整张表没有一处幂运算,每格只是一次除法」一句还揭示了比值写法的另一重好处,同时解释了为什么论文写 γ i / γ j \gamma_i/\gamma_j γ i / γ j 而不直接写 α i − j \alpha^{i-j} α i − j :后者只在 α \alpha α 为常数时才成立,一旦门控逐 token 变化,「唯一底数」就不存在了,∏ r = j + 1 i α r \prod_{r=j+1}^{i}\alpha_r ∏ r = j + 1 i α r 无法写成任何数的幂;而比值形式原样成立。本文因 α \alpha α 为常数而两种写法皆可,论文则必须采用比值形式。
回到实现。 上面的分解还隐含一个容易忽略的陷阱:若把 M \mathbf{M} M 和 D \mathbf{D} D 真的分开算再相乘,D \mathbf{D} D 的上三角是大于 1 的 (i < j i<j i < j 时指数为正)。取 α = 0.9 \alpha=0.9 α = 0.9 、C = 4 C=4 C = 4 :
D = [ 1 1.111 1.235 1.372 0.9 1 1.111 1.235 0.81 0.9 1 1.111 0.729 0.81 0.9 1 ] → ⊙ M Γ = [ 1 0 0 0 0.9 1 0 0 0.81 0.9 1 0 0.729 0.81 0.9 1 ] \mathbf{D} = \begin{bmatrix}1&\mathbf{1.111}&\mathbf{1.235}&\mathbf{1.372}\\0.9&1&\mathbf{1.111}&\mathbf{1.235}\\0.81&0.9&1&\mathbf{1.111}\\0.729&0.81&0.9&1\end{bmatrix}
\ \xrightarrow{\ \odot\ \mathbf{M}\ }\
\Gamma = \begin{bmatrix}1&0&0&0\\0.9&1&0&0\\0.81&0.9&1&0\\0.729&0.81&0.9&1\end{bmatrix}
D = 1 0.9 0.81 0.729 1.111 1 0.9 0.81 1.235 1.111 1 0.9 1.372 1.235 1.111 1 ⊙ M Γ = 1 0.9 0.81 0.729 0 1 0.9 0.81 0 0 1 0.9 0 0 0 1
C = 4 C = 4 C = 4 时上三角最大才 1.372,但 C = 64 C = 64 C = 64 、α = 0.8 \alpha = 0.8 α = 0.8 时右上角是 0.8 − 63 = 1.27 × 10 6 0.8^{-63} = 1.27 \times 10^{6} 0. 8 − 63 = 1.27 × 1 0 6 ,早已溢出 fp16。于是:
写法
中间量范围
结果
先构造完整 D \mathbf{D} D ,再乘 M \mathbf{M} M
[ α C − 1 , α − ( C − 1 ) ] [\alpha^{C-1},\ \alpha^{-(C-1)}] [ α C − 1 , α − ( C − 1 ) ]
fp16 下上三角可能已 inf,inf × 0 = NaN
只在 i ≥ j i \ge j i ≥ j 处求值,否则直接置 0
( 0 , 1 ] (0,\ 1] ( 0 , 1 ]
安全
因此 kernel 里必须写成条件求值而非「算完再掩」(§5.1 的 T.if_then_else(j <= i, ...) 就是这个原因)–哪怕数学上 M ⊙ D \mathbf{M} \odot \mathbf{D} M ⊙ D 与「只算下三角」完全等价。M \mathbf{M} M 的存在恰好保证了 Γ \Gamma Γ 中每个有效元素的指数 i − j ≥ 0 i-j \ge 0 i − j ≥ 0 ,这才是「Γ i j ≤ 1 \Gamma_{ij} \le 1 Γ ij ≤ 1 」这个结论成立的前提。
换个角度理解:softmax 注意力里掩码是加 − ∞ -\infty − ∞ (因为后面要过 exp),线性注意力里掩码是乘 0(因为没有 exp,直接置零即可)–但两者都不该「先算全量再掩」 ,原因分别是数值溢出与计算浪费。
2.3 对照:Mamba2 官方 kernel 怎么写这个衰减
上面几节的结论不是纸上推演–TileLang 官方 examples/linear_attention/example_mamba_chunk_state.py 就是这么做的。它算的是 Mamba2 chunkwise 的状态更新 那一步,对应本文 §1.2 的位置二(K → \overrightarrow{\mathbf{K}} K ,把整块贡献衰减到块尾)。
参考实现一行 einsum 就说完了:
1 2 3 decay_states = torch.exp((dA_cumsum[:, :, :, -1 :] - dA_cumsum)) return torch.einsum("bclhn,bhcl,bhcl,bclhp->bchpn" , B, decay_states, dt, x)
对着本文的记号读,dA_cumsum 就是 log γ [ t ] r \log \gamma_{[t]}^{r} log γ [ t ] r (Δ A \Delta A Δ A 的累积和,A A A 已经是负的),于是:
decay_states [ r ] = exp ( log γ [ t ] C − log γ [ t ] r ) = γ [ t ] C γ [ t ] r \texttt{decay\_states}[r] = \exp\big(\log\gamma_{[t]}^{C} - \log\gamma_{[t]}^{r}\big)
= \frac{\gamma_{[t]}^{C}}{\gamma_{[t]}^{r}}
decay_states [ r ] = exp ( log γ [ t ] C − log γ [ t ] r ) = γ [ t ] r γ [ t ] C
这正是 §1.2 位置二的 k [ t ] r → \overrightarrow{\bm{k}_{[t]}^{r}} k [ t ] r 权重 ,一字不差。Mamba2 里 k \bm{k} k 的角色由 B 承担、v \bm{v} v 由 x 承担,dt 是离散化步长(本文的常数情形没有这一项)。
kernel 里对应的三行是这样的:
1 2 3 4 5 6 7 8 9 p = 1.44269504 dA_cs_last[0 ] = dA_cumsum[batch_idx, bz, chunk_idx, chunk_size - 1 ] ... for i in T.Parallel(block_K): scale[i] = T.exp2(dA_cs_last[0 ] * p - dA_cumsum_local[i] * p) * dt_local[i] for i, j in T.Parallel(block_M, block_K): xt_local[i, j] = x_local[j, i] * scale[j] T.gemm(xt_local, B_shared, acc_o)
四处细节和本文的结论逐条对上:
官方写法
为什么这么写
输入是 dA_cumsum(log 域累积和),不是 γ \gamma γ 本身
连乘会下溢:fp16 直接 cumprod 在 C = 128 C = 128 C = 128 、α ∈ [ 0.5 , 0.9 ] \alpha \in [0.5,\ 0.9] α ∈ [ 0.5 , 0.9 ] 时归零。log 域改成加法就没这个问题
exp2(log γ_C - log γ_r),先在 log 域相减再取指数
相减的结果因 r ≤ C r \le C r ≤ C 恒 ≤ 0 \le 0 ≤ 0 ,exp2 输出恒 ≤ 1 \le 1 ≤ 1 。若先各自取指数再相除,就要物化 1 / γ r 1/\gamma_r 1/ γ r 这个大于 1 的量
乘 p = 1.44269504 把 e x e^x e x 转成硬件 exp2
1.44269504 = log 2 e 1.44269504 = \log_2 e 1.44269504 = log 2 e ,于是 e x = 2 x log 2 e e^x = 2^{x\log_2 e} e x = 2 x l o g 2 e 。exp2 有单指令实现,pow 通常展开成多条
dA_cs_last 在 T.Pipelined 循环外 读一次
log γ C \log\gamma_C log γ C 是整块共用的常量,每轮重取只是白读一次 shared memory
最值得注意的是第二行。它完全可以写成 exp2(log γ_C * p) / exp2(log γ_r * p)–数学上等价,还能把 γ C \gamma_C γ C 提到循环外省一次减法。但那就等于物化了 1 / γ r 1/\gamma_r 1/ γ r ,正是 §3 实测会 NaN 的写法。官方选择在 log 域相减,指数结果因 r ≤ C r \le C r ≤ C 、log γ \log\gamma log γ 单调递减而恒 ≤ 0 \le 0 ≤ 0 ,exp2 输出恒不大于 1,不存在溢出可能。
顺带一个和本文互补的观察:官方 kernel 把 scale 直接乘到了 xt_local(即 v \bm{v} v 侧)而非 B_shared(k \bm{k} k 侧)。两者数学等价–衰减是标量,挂在哪一侧都行;选 x 是因为那一步本来就要做转置 xt_local[i,j] = x_local[j,i],将衰减乘法合并进转置的同一个 T.Parallel 中,省去一趟对 shared memory 的读写 。这类「把逐元素操作合并进已有的数据搬运」是 tile 编程中常见的优化手法。
3. 定量分析:因式分解在 fp16 下的失效点
论文的写法(先 GEMM 再逐元素乘 Γ \Gamma Γ )中间量恒在 ( 0 , 1 ] (0,1] ( 0 , 1 ] ,是数值安全的。但块内下三角 ( Γ [ t ] ) i j = α i − j (\Gamma_{[t]})_{ij} = \alpha^{\,i-j} ( Γ [ t ] ) ij = α i − j 存在一个看似有利的代数变形–指数可以拆开:
α i − j = α i ⋅ α − j \alpha^{\,i-j} = \alpha^{\,i} \cdot \alpha^{-j}
α i − j = α i ⋅ α − j
于是块内那一项可以写成:
( Q [ t ] K [ t ] ⊺ ⊙ Γ [ t ] ) = tril ( ( Q [ t ] ⊙ α i ) ( K [ t ] ⊙ α − j ) ⊺ ) \big(\mathbf{Q}_{[t]}\mathbf{K}_{[t]}^{\intercal} \odot \Gamma_{[t]}\big)
= \operatorname{tril}\big( (\mathbf{Q}_{[t]} \odot \alpha^{\,i}) (\mathbf{K}_{[t]} \odot \alpha^{-j})^{\intercal} \big)
( Q [ t ] K [ t ] ⊺ ⊙ Γ [ t ] ) = tril ( ( Q [ t ] ⊙ α i ) ( K [ t ] ⊙ α − j ) ⊺ )
两种实现方式的差别:
方案
做法
块内额外开销
权重数值范围
比值形式(论文)
先 GEMM 得分数,再逐元素乘 Γ \Gamma Γ
一次 C × C C \times C C × C 逐元素乘 + 一张 C 2 C^2 C 2 权重表
( 0 , 1 ] (0, 1] ( 0 , 1 ]
因式分解
先给 Q Q Q 、K K K 的行分别乘上衰减因子,再单次 GEMM
两次 C × d k C \times d_k C × d k 逐行加权,无 C 2 C^2 C 2 开销
α − j \alpha^{-j} α − j 最大 α − ( C − 1 ) \alpha^{-(C-1)} α − ( C − 1 )
因式分解在 FLOPs 与寄存器占用上都更优:省掉一张 C × C C \times C C × C 的权重表,逐元素乘的规模从 C 2 C^2 C 2 降到 2 C d k 2Cd_k 2 C d k 。C = 64 C = 64 C = 64 、d k = 64 d_k = 64 d k = 64 时前者是 4096 次乘法,后者 8192 次–乘法次数反而增加,但权重表不必常驻寄存器,这在 C = 128 C = 128 C = 128 时是实质性的压力缓解。
问题在数值范围。α − j \alpha^{-j} α − j 是大于 1 的量 ,且随 j j j 指数增长:
α \alpha α
α − 31 \alpha^{-31} α − 31 (C C C =32)
α − 63 \alpha^{-63} α − 63 (C C C =64)
α − 127 \alpha^{-127} α − 127 (C C C =128)
fp16 安全上限
fp32 安全上限
0.99
1.37 1.37 1.37
1.88 1.88 1.88
3.58 3.58 3.58
C < 1103 C < 1103 C < 1103
C < 8827 C < 8827 C < 8827
0.95
4.90 4.90 4.90
25.3 25.3 25.3
675 675 675
C < 216 C < 216 C < 216
C < 1729 C < 1729 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
C < 105 C < 105 C < 105
C < 842 C < 842 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
C < 49 C < 49 C < 49
C < 397 C < 397 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
C < 15 C < 15 C < 15
C < 127 C < 127 C < 127
fp16 上限是 65504。α = 0.8 \alpha = 0.8 α = 0.8 、C = 64 C = 64 C = 64 时 α − 63 = 1.27 × 10 6 \alpha^{-63} = 1.27 \times 10^6 α − 63 = 1.27 × 1 0 6 已经溢出;α = 0.5 \alpha = 0.5 α = 0.5 时 9.22 × 10 18 9.22 \times 10^{18} 9.22 × 1 0 18 连 fp32 都接近极限。
3.1 实测:单块块内计算的精度对比
单块块内计算的精度对比(C = 64 C = 64 C = 64 ,d k = d v = 64 d_k = d_v = 64 d k = d v = 64 ,fp16 存储 + fp32 累加,20 组随机输入取中位数,参考值为 fp64 精确计算):
α \alpha α
比值形式 相对 L2
因式分解 相对 L2
α − ( C − 1 ) \alpha^{-(C-1)} α − ( 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 \alpha \ge 0.9 α ≥ 0.9 时两者精度相当,因式分解甚至略优(少一次逐元素乘引入的舍入);
α ≤ 0.8 \alpha \le 0.8 α ≤ 0.8 时因式分解在 fp16 下彻底失效 ,产生 NaN 而非精度下降–α − j \alpha^{-j} α − j 溢出成 inf,随后 inf 乘 0 得 NaN;
比值形式的误差在全部 α \alpha α 取值下稳定在 4.1 – 4.6 × 10 − 4 4.1\text{--}4.6 \times 10^{-4} 4.1 – 4.6 × 1 0 − 4 ,与 α \alpha α 无关。
比值形式的误差之所以稳定,是因为 Γ i j \Gamma_{ij} Γ ij 恒在 ( 0 , 1 ] (0,1] ( 0 , 1 ] –这是一个与 α \alpha α 无关的数值保证 。实测各 α \alpha α 下衰减矩阵的取值:
α \alpha α
C C 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 \alpha = 0.5 α = 0.5 、C = 128 C = 128 C = 128 ,最小权重 5.88 × 10 − 39 5.88 \times 10^{-39} 5.88 × 1 0 − 39 在 fp32 下仍是正规数,无下溢。下溢比溢出安全得多 :权重下溢为 0 意味着「这个远距离贡献可以忽略」,语义上正确;而溢出为 inf 会污染整行输出。
3.2 比值形式自身的下溢边界
比值形式也不是完全没有约束。Γ \Gamma Γ 本身若用 fp16 存储,α C − 1 \alpha^{C-1} α C − 1 可能低于 fp16 最小正规数 6.10 × 10 − 5 6.10 \times 10^{-5} 6.10 × 1 0 − 5 :
α \alpha α
C = 64 C=64 C = 64
fp16 状态
C = 128 C=128 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 采用的做法。
一句话总结 :因式分解能省一次 C 2 C^2 C 2 逐元素乘,但把中间量范围从 ( 0 , 1 ] (0,1] ( 0 , 1 ] 推到 [ 1 , α − ( C − 1 ) ] [1, \alpha^{-(C-1)}] [ 1 , α − ( C − 1 ) ] ,α ≤ 0.8 \alpha \le 0.8 α ≤ 0.8 时 fp16 直接 NaN–这个 FLOPs 优化不值得,论文的比值形式应原样保留。
4. 四层参考实现
沿用上一篇的三层框架,本文增加一层专门验证因式分解写法:
参考
实现方式
验证目标
A
逐 token 递归
递推式定义 S t = α S t − 1 + v t k t ⊺ \mathbf{S}_t = \alpha\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} S t = α S t − 1 + v t k t ⊺
B
分块向量化,三处衰减权重
§1.2 三个权重位置的推导
C
逐块独立重算,模拟 grid
kernel 控制流
D
因式分解版块内计算
§3 两种写法的代数等价性
4.0 一个必要的转置:论文的 R d v × d k \mathbb{R}^{d_v \times d_k} R d v × d k vs kernel 的 R d k × d v \mathbb{R}^{d_k \times d_v} R d k × d v
这里要先交代一个容易造成混乱的差异。论文(以及本文 §1–§3 的全部推导)的状态是 S ∈ R d v × d k \mathbf{S} \in \mathbb{R}^{d_v \times d_k} S ∈ R d v × d k :
S t = α t S t − 1 + v t k t ⊺ ∈ R d v × d k , o t = S t q t ∈ R d v \mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} \in \mathbb{R}^{d_v \times d_k},
\qquad \bm{o}_t = \mathbf{S}_t\bm{q}_t \in \mathbb{R}^{d_v}
S t = α t S t − 1 + v t k t ⊺ ∈ R d v × d k , o t = S t q t ∈ R d v
注意是 v k ⊺ \bm{v}\bm{k}^{\intercal} v k ⊺ (value 在外、key 转置在内)、读出是 S q \mathbf{S}\bm{q} S q (状态左乘 query)。而下面所有代码里的 S 存的都是它的转置 :
S = ^ S ⊺ ∈ R d k × d v \texttt{S} \ \widehat{=}\ \mathbf{S}^{\intercal} \in \mathbb{R}^{d_k \times d_v}
S = S ⊺ ∈ R d k × d v
于是两边的写法逐行对应:
论文(S ∈ R d v × d k \mathbf{S} \in \mathbb{R}^{d_v \times d_k} S ∈ R d v × d k )
代码(S ∈ R d k × d v \in \mathbb{R}^{d_k \times d_v} ∈ R d k × d v )
S t = α S t − 1 + v t k t ⊺ \mathbf{S}_t = \alpha\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} S t = α S t − 1 + v t k t ⊺
S = g*S + outer(k, v)
o t = S t q t \bm{o}_t = \mathbf{S}_t\bm{q}_t o t = S t q t
o = q @ S
S [ t + 1 ] = S [ t ] → + V [ t ] ⊺ K [ t ] → \mathbf{S}_{[t+1]} = \overrightarrow{\mathbf{S}_{[t]}} + \mathbf{V}_{[t]}^{\intercal}\overrightarrow{\mathbf{K}_{[t]}} S [ t + 1 ] = S [ t ] + V [ t ] ⊺ K [ t ]
S = g**C * S + Kd.T @ V
O [ t ] = Q [ t ] ← S [ t ] ⊺ + … \mathbf{O}_{[t]} = \overleftarrow{\mathbf{Q}_{[t]}}\mathbf{S}_{[t]}^{\intercal} + \dots O [ t ] = Q [ t ] S [ t ] ⊺ + …
acc = (Qb * w_query) @ S + ...
为何不直接按论文的方向存?因为 d k × d v d_k \times d_v d k × d v 布局下,跨块读出就是 T.gemm(Q_s, S_s, acc_o)–[ C , d k ] × [ d k , d v ] [C, d_k] \times [d_k, d_v] [ C , d k ] × [ d k , d v ] ,不需任何 transpose 标志 ;若按论文方向存,每次读出都要 transpose_B=True。同理状态更新 V ⊺ K → \mathbf{V}^{\intercal}\overrightarrow{\mathbf{K}} V ⊺ K 在转置布局下变成 K → ⊺ V \overrightarrow{\mathbf{K}}^{\intercal}\mathbf{V} K ⊺ V ,正好是 transpose_A=True 一个标志就能表达的形式。这是纯工程选择,不影响任何数学结论 –但读代码时必须心里有这个转置,否则会觉得和论文对不上。
4.1 参考 A:逐 token 递归
1 2 3 4 5 6 7 8 9 10 11 def ref_recurrent (Q, K, V, g ): """严格照抄递推式:S = g*S + k v^T""" B, N, H, D = Q.shape O = torch.empty(B, N, H, D, dtype=torch.float64, device=Q.device) for b in range (B): for h in range (H): S = torch.zeros(D, D, dtype=torch.float64, device=Q.device) for t in range (N): S = g * S + torch.outer(K[b, t, h].double(), V[b, t, h].double()) O[b, t, h] = Q[b, t, h].double() @ S return O
与上一篇的唯一差别是 S = g * S + ... 而非 S += ...。
4.2 参考 B:分块向量化
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 def ref_chunked (Q, K, V, g, BC ): """三处衰减权重分别对应 §1.2 的位置一、二、三""" B, N, H, D = Q.shape NC = N // BC Qc = Q.permute(0 , 2 , 1 , 3 ).double().reshape(B, H, NC, BC, D) Kc = K.permute(0 , 2 , 1 , 3 ).double().reshape(B, H, NC, BC, D) Vc = V.permute(0 , 2 , 1 , 3 ).double().reshape(B, H, NC, BC, D) i = torch.arange(BC, dtype=torch.float64, device=Q.device) decay_tri = torch.where(i[:, None ] >= i[None , :], g ** (i[:, None ] - i[None , :]), torch.zeros((), dtype=torch.float64, device=Q.device)) w_decay = g ** (BC - 1 - i) w_query = g ** (i + 1 ) contrib = torch.einsum("bhcmd,bhcmv->bhcdv" , Kc * w_decay[:, None ], Vc) states_prev = torch.zeros(B, H, NC, D, D, dtype=torch.float64, device=Q.device) for c in range (1 , NC): states_prev[:, :, c] = g ** BC * states_prev[:, :, c - 1 ] + contrib[:, :, c - 1 ] O_inter = torch.einsum("bhcnd,bhcdv->bhcnv" , Qc * w_query[:, None ], states_prev) A = torch.einsum("bhcnd,bhcmd->bhcnm" , Qc, Kc) * decay_tri O_intra = torch.einsum("bhcnm,bhcmv->bhcnv" , A, Vc) return (O_inter + O_intra).reshape(B, H, N, D).permute(0 , 2 , 1 , 3 ).contiguous()
跨块状态在上一篇可以用 torch.cumsum 一次算完,本文不行–cumsum 是无权重的前缀和,而本文的递推带系数 α C \alpha^{C} α C 。这里保留显式循环;若要向量化,需改用加权前缀和 的写法:把 Δ [ t ] \Delta_{[t]} Δ [ t ] 先除以 α t C \alpha^{tC} α tC 再 cumsum、最后乘回,代价是又引入 α − t C \alpha^{-tC} α − tC 这个溢出源,t t t 大时比 §3 的块内因式分解更危险。这是同一个取舍在跨块层面的重演 –它的本质仍是 §3 那个 1 / γ 1/\gamma 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_decay = g ** (BC - 1 - i) w_query = g ** (i + 1 ) for bz in range (B): for by in range (H): for bx in range (NC): sl = slice (bx * BC, (bx + 1 ) * BC) S = torch.zeros(D, D, dtype=torch.float64, device=Q.device) for c in range (bx): cs = slice (c * BC, (c + 1 ) * BC) Kd = K[bz, cs, by, :].double() * w_decay[:, None ] S = g ** BC * S + Kd.T @ V[bz, cs, by, :].double() Qb = Q[bz, sl, by, :].double() Kb = K[bz, sl, by, :].double() Vb = V[bz, sl, by, :].double() acc = (Qb * w_query[:, None ]) @ S acc += ((Qb @ Kb.T) * decay_tri) @ Vb O[bz, sl, by, :] = acc return O
与上一篇参考 C 的差别只有两处:状态累加从 S += 变成 S = g**BC * S + ...,以及三处权重乘法。循环结构完全没变 ,这正是 §1 结论「分块恒等式结构不变」的代码体现。
4.4 四层参考的一致性验证
B = 2 , H = 2 , N = 12 , D = 4 , B C = 4 , α = 0.9 B=2, H=2, N=12, D=4, BC=4, \alpha=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 \alpha = 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
另外用逐 token 变化的 α t ∼ U ( 0.85 , 0.999 ) \alpha_t \sim \mathcal{U}(0.85,\ 0.999) α t ∼ U ( 0.85 , 0.999 ) 复验过一遍(同规模,fp64),确认 §2.2 的矩阵并行形式在门控非常数时同样成立:矩阵形式 vs 向量形式相对 L2 2.03 × 10 − 16 2.03 \times 10^{-16} 2.03 × 1 0 − 16 ,Γ i j \Gamma_{ij} Γ ij 取值范围 [ 0.628 , 1.000 ] [0.628,\ 1.000] [ 0.628 , 1.000 ] 全部 ≤ 1 \le 1 ≤ 1 。
前三行确认四份实现数学等价,误差均在 fp64 机器精度量级。第四行是退化检验 :令 α = 1 \alpha = 1 α = 1 应当精确回到上一篇的无衰减实现,这条验证能同时捕获三处权重中任何一处的指数写错–例如把 α C − r \alpha^{\,C-r} α C − r 误写成 α C − r + 1 \alpha^{\,C-r+1} α C − r + 1 ,α = 1 \alpha = 1 α = 1 时两者都是 1,退化检验通过但 α = 0.9 \alpha = 0.9 α = 0.9 时参考 B 与 A 不符。两条验证必须都做。
5. TileLang kernel 的改动
沿用上一篇 §6.4.2 的 grid 划分:切 ( b v , b ⋅ h ) (bv,\ b\cdot h) ( b v , b ⋅ h ) ,序列轴不进 grid、退回 kernel 内的 T.Pipelined 顺序循环 ,状态作为 loop-carried fragment 常驻寄存器,每 chunk 只做一次 K ⊤ V K^\top V K ⊤ V 。本文在此基础上新增一张 f32 权重表与两个权重向量。
先说清为什么必须切 DV。本文比上一篇多出一个 C × C C \times C C × C 的 Dtri,而它在 DV 方向是共享的 (衰减权重只跟 chunk 内位置有关,与 value 通道无关),因此不随 DV 切分而缩小。d k = d v = 128 d_k = d_v = 128 d k = d v = 128 、C = 64 C = 64 C = 64 、128 线程下点一遍账:
fragment
不切 DV
block D V = 32 \text{block}_{DV} = 32 block D V = 32
S_f
[ 128 , 128 ] [128,128] [ 128 , 128 ] → 128 reg
[ 128 , 32 ] [128,32] [ 128 , 32 ] → 32 reg
acc_o
[ 64 , 128 ] [64,128] [ 64 , 128 ] → 64 reg
[ 64 , 32 ] [64,32] [ 64 , 32 ] → 16 reg
Dtri + A
2 × [ 64 , 64 ] 2 \times [64,64] 2 × [ 64 , 64 ] → 64 reg
64 reg(不变 )
合计
≈ 256 reg/thread ,已超 255 上限
≈ 112 reg/thread
不切 DV 在这个配置下直接编译不过。 代价是 Dtri 被每个 b v bv b v block 各算一遍–D V / block D V = 4 DV/\text{block}_{DV} = 4 D V / block D V = 4 时同一张 4096 元素的表算了 4 遍,合计 16384 次 exp2。这是纯冗余计算,但它不占额外寄存器,且 exp2 是单指令,相比编译不过是划算的交换。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 @tilelang.jit(out_idx=[3 ] ) def linattn_decay_chunk (B, H, S, DK, DV, alpha, blk=64 , block_DV=32 , num_stages=2 , threads=128 , dtype=T.float16, accum_dtype=T.float32 ): C = blk NS = T.ceildiv(S, C) alpha_pow_C = alpha ** C @T.prim_func def main ( Q: T.Tensor([B, S, H, DK], dtype ), K: T.Tensor([B, S, H, DK], dtype ), V: T.Tensor([B, S, H, DV], dtype ), O: T.Tensor([B, S, H, DV], dtype ), ): with T.Kernel(T.ceildiv(DV, block_DV), B * H, threads=threads) as (bv, bbh): bb, bh = bbh // H, bbh % H dv0 = bv * block_DV Q_s = T.alloc_shared([C, DK], dtype) K_s = T.alloc_shared([C, DK], dtype) V_s = T.alloc_shared([C, block_DV], dtype) S_s = T.alloc_shared([DK, block_DV], dtype) O_s = T.alloc_shared([C, block_DV], dtype) S_f = T.alloc_fragment([DK, block_DV], accum_dtype) acc_o = T.alloc_fragment([C, block_DV], accum_dtype) A = T.alloc_fragment([C, C], accum_dtype) A_cast = T.alloc_fragment([C, C], dtype) Dtri = T.alloc_fragment([C, C], accum_dtype) w_decay = T.alloc_fragment([C], accum_dtype) w_query = T.alloc_fragment([C], accum_dtype)
5.1 预计算衰减权重
三张权重表都在 T.Pipelined 之前算一次,整条序列共用:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 lg = T.log2(T.Cast(accum_dtype, alpha)) for i, j in T.Parallel(C, C): Dtri[i, j] = T.if_then_else( j <= i, T.exp2(lg * T.Cast(accum_dtype, i - j)), 0.0 ) for m in T.Parallel(C): w_decay[m] = T.exp2(lg * T.Cast(accum_dtype, C - 1 - m)) for i in T.Parallel(C): w_query[i] = T.exp2(lg * T.Cast(accum_dtype, i + 1 ))
三者的指数都取自 §1.2 的权重表:
变量
形状
元素
指数含义
取值范围
Dtri[i,j]
C × C C \times C C × C
α i − j \alpha^{\,i-j} α i − j (j ≤ i j \le i j ≤ i ,否则 0)
同块内两 token 的间隔
[ α C − 1 , 1 ] [\alpha^{C-1},\ 1] [ α C − 1 , 1 ]
w_decay[m]
C C C
α C − 1 − m \alpha^{\,C-1-m} α C − 1 − m
第 m m m 行到块尾的步数
[ α C − 1 , 1 ] [\alpha^{C-1},\ 1] [ α C − 1 , 1 ]
w_query[i]
C C C
α i + 1 \alpha^{\,i+1} α i + 1
第 i i i 行到上块末尾的步数
[ α C , α ] [\alpha^{C},\ \alpha] [ α C , α ]
注意 w_decay 与 w_query 的方向相反 :w_decay[C-1] = α^0 = 1(块尾那一行本身就是基准,不需衰减),而 w_query[0] = α^1(块首那一行距上块末尾也有 1 步)。这两处若把指数写反或差一,α = 1 \alpha = 1 α = 1 时完全看不出来–§4.4 的退化检验正是为此设计的。
指数用 exp2(log2(α) · k) 而非 pow(α, k):exp2 与 log2 都有单指令硬件实现,而 pow 通常展开成多条指令。
Dtri 的 T.if_then_else 不能省成「先全算再掩」–如 §2.2 所述,j > i j > i j > i 处的 α i − j \alpha^{i-j} α i − j 指数为正,C = 64 C = 64 C = 64 、α = 0.8 \alpha = 0.8 α = 0.8 时右上角已达 1.27 × 10 6 1.27 \times 10^{6} 1.27 × 1 0 6 ,fp16 下溢出成 inf 后再乘 0 会得到 NaN。这里的条件求值同时是正确性和数值安全两重保障。
5.2 序列内循环:三处权重的落点
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 T.clear(S_f) for i_s in T.Pipelined(NS, num_stages=num_stages): s0 = i_s * C T.copy(Q[bb, s0:s0+C, bh, :], Q_s) T.copy(K[bb, s0:s0+C, bh, :], K_s) T.copy(V[bb, s0:s0+C, bh, dv0:dv0+block_DV], V_s) T.copy(S_f, S_s) for i, d in T.Parallel(C, DK): Q_s[i, d] *= w_query[i] T.gemm(Q_s, S_s, acc_o, clear_accum=True ) T.copy(Q[bb, s0:s0+C, bh, :], Q_s) T.gemm(Q_s, K_s, A, transpose_B=True , clear_accum=True ) for i, j in T.Parallel(C, C): A[i, j] *= Dtri[i, j] T.copy(A, A_cast) T.gemm(A_cast, V_s, acc_o) T.copy(acc_o, O_s) T.copy(O_s, O[bb, s0:s0+C, bh, dv0:dv0+block_DV]) for i, j in T.Parallel(DK, block_DV): S_f[i, j] *= alpha_pow_C for m, d in T.Parallel(C, DK): K_s[m, d] *= w_decay[m] T.gemm(K_s, V_s, S_f, transpose_A=True ) return main
四处次序约束,写错任何一处都不会报错、只会算错:
状态更新必须排在输出写回之后 。S f S_f S f 在步骤②被读取时代表 S [ t ] \mathbf{S}_{[t]} S [ t ] (进块状态),本块自己的贡献由 tril \operatorname{tril} tril 那一项负责。提前更新就变成了 inclusive 前缀,块内项被重复计入。
S_f *= alpha_pow_C 必须在 T.gemm(K_s, V_s, S_f) 之前 。递推式是 S [ t + 1 ] = α C S [ t ] + Δ [ t ] \mathbf{S}_{[t+1]} = \alpha^{C}\mathbf{S}_{[t]} + \Delta_{[t]} S [ t + 1 ] = α C S [ t ] + Δ [ t ] ,先衰减旧状态再累加新贡献;若顺序颠倒,本块贡献会被多衰减一次。
步骤③必须重载 Q 。步骤②已把 Q_s 原地乘上 α r \alpha^{r} α r ,而块内项用的是未加权的原始 Q [ t ] \mathbf{Q}_{[t]} Q [ t ] 。原地加权省了一个 buffer,代价是重载一次;若寄存器有余量,另开 Q_scaled 更安全。
T.clear(S_f) 在循环外,clear_accum=True 在循环内 。状态要跨迭代累加,只能循环外清一次;acc_o 与 A 每 chunk 都是全新的,进了流水线循环就不能再用 T.clear(会被排到错误阶段),必须靠 clear_accum=True 在 MMA 那一刻覆盖累加器。
K_s 在循环末尾被原地乘过 w_decay,而步骤③用的是同一个 K_s–这里没有冲突,因为③在①之前执行。但若调整次序或复用,必须重新载入:本文的 K_s 残留内容已被乘过权重 ,比上一篇「残留内容不确定」更危险。
5.3 与上一篇 kernel 的改动汇总
位置
上一篇
本文
新增开销
grid
( N S , H , B ) (NS, H, B) ( N S , H , B )
( D V / block D V , B ⋅ H ) (DV/\text{block}_{DV},\ B \cdot H) ( D V / block D V , B ⋅ H )
序列不进 grid,状态常驻寄存器
权重表
无
Dtri[BC, BC] f32 fragment
C 2 C^2 C 2 个 f32 寄存器
状态累加
T.gemm 直接累加
先 S_f *= alpha_pow_C 再 gemm
d k d v d_k d_v d k d v 次乘法 / 轮
K 写入状态
原样
逐行乘 α C − r \alpha^{\,C-r} α C − r
C × d k C \times d_k C × d k 次乘法 / 轮
Q 读取状态
原样
逐行乘 α r \alpha^{r} α r
C × d k C \times d_k C × d k 次乘法
块内掩码
if_then_else 置 0
逐元素乘 Dtri
C 2 C^2 C 2 次乘法
Q 复用
步骤②③共用
步骤③必须重载
一次 shared 写入
切 DV 后 S_f 从 [ 128 , 128 ] [128,128] [ 128 , 128 ] 降到 [ 128 , 32 ] [128,32] [ 128 , 32 ] ,寄存器压力最大的反而不是 Dtri。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 时逐项数:
fragment
上一篇
本文
变化
S_f
[ 128 , 128 ] [128,128] [ 128 , 128 ] → 128 reg
[ 128 , 32 ] [128,32] [ 128 , 32 ] → 32 reg
下降 96 reg
acc_o
[ 64 , 128 ] [64,128] [ 64 , 128 ] → 64 reg
[ 64 , 32 ] [64,32] [ 64 , 32 ] → 16 reg
下降 48 reg
Dtri + A
无
2 × [ 64 , 64 ] 2 \times [64,64] 2 × [ 64 , 64 ] → 64 reg
新增 64 reg
合计
约 224 reg
约 112 reg
↓ 一半
切 DV 省下的寄存器刚好覆盖 Dtri 的开销,还有余量。 若 C = 128 C = 128 C = 128 ,Dtri + A 升到 256 reg,即使切 DV 到 32 也超过 255 上限–那才是因式分解唯一真正有吸引力的地方(它不需要 C 2 C^2 C 2 权重表),但 §3.1 的 NaN 结论表明这个吸引力不成立。
6. 衰减带来的新维度:有效记忆长度
α \alpha α 引入了一个上一篇不存在的语义参数–状态的遗忘速度。定义有效记忆长度 为权重衰减到 10 − 3 10^{-3} 1 0 − 3 所需的 token 数,即 α L = 10 − 3 \alpha^{L} = 10^{-3} α L = 1 0 − 3 :
L = ln 10 − 3 ln α L = \frac{\ln 10^{-3}}{\ln \alpha}
L = ln α ln 1 0 − 3
α \alpha α
α 64 \alpha^{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 \alpha = 0.9 α = 0.9 的有效记忆只有 66 token,连一个 chunk(C = 64 C = 64 C = 64 )都刚刚覆盖,长程依赖完全丢失。而 α ≥ 0.99 \alpha \ge 0.99 α ≥ 0.99 恰好落在 §3 表格中因式分解仍然安全的区间 –这解释了为什么部分实现敢做这个分解:它们隐含假设了门控接近 1。
这个假设在标量、常数门控下可以接受–反正全局就一个 α \alpha α ,调到 0.99 以上就行。但它是一个隐含假设,不是数值保证 :一旦门控变成数据依赖的(尤其是每个通道各自学一个),训练中总会有部分值降到 0.9 以下以实现快速遗忘,「全部接近 1」就不再成立,因式分解的溢出从例外变成常态。这也是 §3 实测中 α ≤ 0.8 \alpha \le 0.8 α ≤ 0.8 即溢出的原因。结论不变:不要做那个分解,不要依赖「门控总是接近 1」这个前提。
一句话总结 :α \alpha α 的安全区间(接近 1)与有用区间(提供实际遗忘能力)方向相反,常数门控下两者尚可兼顾,但这个兼顾靠的是假设而非保证。
7. 数值验证
验证层次沿用上一篇结构,新增两项:
四层参考互验 (fp64):A/B/C/D 两两对照,误差应在 10 − 15 10^{-15} 1 0 − 15 量级;
退化检验 :α = 1 \alpha = 1 α = 1 时参考 B 应精确回到上一篇实现–这条能捕获三处权重的指数偏移错误;
fp16 失效点复现 :α ≤ 0.8 \alpha \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 上完成(见 §3.1 与 §4.4 表格)。第 4、5 层需要 CUDA 设备,待真卡跑通后单独补充实测数据 ,此处不做性能推测。
本文 kernel 的预期主导误差源与上一篇相同–T.copy(S_f, S_s) 的 f32→f16 降精度。但衰减带来一个有利变化:α C \alpha^{C} α C 每轮把旧状态压缩,早期块的累积误差也随之衰减,因此误差不再像上一篇那样随 b x bx b x 单调增长,而是趋于一个稳态。α = 0.9 \alpha = 0.9 α = 0.9 、C = 64 C = 64 C = 64 时 α C = 1.18 × 10 − 3 \alpha^{C} = 1.18 \times 10^{-3} α C = 1.18 × 1 0 − 3 ,约三轮之后早期误差已不可见。衰减机制顺带改善了数值稳定性 ,这是一个反直觉但合理的副作用。
8. 总结
衰减项来自 Mamba2 的遗忘门 。Vanilla 线性注意力的状态只增不减,语言建模上明显弱于 Transformer;Mamba2 的补救是乘一个数据依赖的标量 α t ∈ ( 0 , 1 ) \alpha_t \in (0,1) α t ∈ ( 0 , 1 ) ,即 S t = α t S t − 1 + v t k t ⊺ \mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} S t = α t S t − 1 + v t k t ⊺ ,同时带来有界性 (状态范数不再无界增长)与选择性 (α t → 0 \alpha_t \to 0 α t → 0 快速清空、→ 1 \to 1 → 1 保持记忆)。本文取其常数特例 α t ≡ α \alpha_t \equiv \alpha α t ≡ α –这正是 RetNet / Lightning-Attention 的形态。
标量衰减不改变分块恒等式的结构 ,只在三个位置插入指数权重(即论文式 (2) 的三个箭头量):块内下三角 ( Γ [ t ] ) i j = α i − j (\Gamma_{[t]})_{ij} = \alpha^{\,i-j} ( Γ [ t ] ) ij = α i − j 、写入状态时的块尾对齐 k → \overrightarrow{\bm{k}} k 即 α C − r \alpha^{\,C-r} α C − r 、跨块递推 S → \overrightarrow{\mathbf{S}} S 即 α C \alpha^{C} α C 与 query 侧 q ← \overleftarrow{\bm{q}} q 即 α r \alpha^{r} α r 。四个权重全部 ≤ 1 \le 1 ≤ 1 ,这是本文数值安全的根本原因。
累积衰减积 γ j = ∏ i = 1 j α i \gamma_j = \prod_{i=1}^{j} \alpha_i γ j = ∏ i = 1 j α i 是统一两种形式的代数工具 。它把区间连乘 ∏ r = i + 1 t α r \prod_{r=i+1}^{t} \alpha_r ∏ r = i + 1 t α r 化归为前缀量之比 γ t / γ i \gamma_t / \gamma_i γ t / γ i ,于是递归的向量形式与并行的矩阵形式 O = ( ( Q K ⊺ ) ⊙ Γ ) V \mathbf{O} = ((\mathbf{Q}\mathbf{K}^{\intercal}) \odot \Gamma)\mathbf{V} O = (( Q K ⊺ ) ⊙ Γ ) V 可以相互转换(即 Mamba2 所谓的状态空间对偶性,实测相对 L2 2.03 × 10 − 16 2.03 \times 10^{-16} 2.03 × 1 0 − 16 )。需注意论文中 α t \alpha_t α t 为单步衰减、γ j \gamma_j γ j 为累积积,二者不可混淆 ;常数情形下 γ j = α j \gamma_j = \alpha^{\,j} γ j = α j ,此时 Γ i j = α i − j \Gamma_{ij} = \alpha^{\,i-j} Γ ij = α i − j 。另外累积积按 chunk 重置(论文式 (1) 脚注),不是从序列起点算。
论文的衰减矩阵 Γ i j = γ i / γ j \Gamma_{ij} = \gamma_i/\gamma_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{\bm{q}} q 、k → \overrightarrow{\bm{k}} k 、S → \overrightarrow{\mathbf{S}} S 把「衰减到首 / 末位置」直接编码在箭头方向上,基准点选得当就能保证所有指数非正。
真正的陷阱是把比值因式分解成 γ i ⋅ ( 1 / γ j ) \gamma_i \cdot (1/\gamma_j) γ i ⋅ ( 1/ γ j ) 。这个分解能把两步合成单次 GEMM 并免去 C 2 C^2 C 2 权重表,但要求物化 1 / γ j 1/\gamma_j 1/ γ j ,把中间量从 ( 0 , 1 ] (0,1] ( 0 , 1 ] 推到 [ 1 , 1 / γ C ] [1,\ 1/\gamma_C] [ 1 , 1/ γ C ] 。实测(C = 64 C=64 C = 64 ,fp16 存储加 fp32 累加):α ≥ 0.9 \alpha \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 \alpha \le 0.8 α ≤ 0.8 时分解写法产生 NaN –α = 0.8 \alpha = 0.8 α = 0.8 时 α − 63 = 1.27 × 10 6 \alpha^{-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 ] 、C = 128 C = 128 C = 128 时 1 / γ j 1/\gamma_j 1/ γ j 最大达 3.40 × 10 8 3.40 \times 10^{8} 3.40 × 1 0 8 ,同样溢出。结论是保留论文的比值形式,不做分解。
累积积本身也必须在 log 域计算 。fp16 直接 cumprod 在 C = 128 C = 128 C = 128 、α ∈ [ 0.5 , 0.9 ] \alpha \in [0.5,\ 0.9] α ∈ [ 0.5 , 0.9 ] 时已完全下溢为 0,fp32 在 C = 512 C = 512 C = 512 时进入非正规数区间;log γ j = ∑ i ≤ j log α i \log \gamma_j = \sum_{i \le j} \log \alpha_i log γ j = ∑ i ≤ j log α i 是线性增长的负数,表示范围安全。这就是门控全程存 log 值、用 exp2 还原的原因。
退化检验是本文新增的关键验证手段 :令 α = 1 \alpha = 1 α = 1 应精确回到上一篇实现(实测相对 L2 1.62 × 10 − 16 1.62 \times 10^{-16} 1.62 × 1 0 − 16 )。但它无法单独定案–三处权重中任何一处的指数偏移在 α = 1 \alpha = 1 α = 1 时都不可见,必须同时做 α = 0.9 \alpha = 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 )。
四处次序约束 (写错都不报错、只算错):状态更新必须排在输出写回之后(否则 inclusive 前缀导致块内项重复计入);S_f *= alpha_pow_C 必须在 T.gemm 之前(否则本块贡献被多衰减一次);步骤③必须重载 Q(步骤②已把 Q_s 原地乘了 α r \alpha^{r} α r );T.clear(S_f) 在流水线循环外、clear_accum=True 在循环内。
必须切 DV :本文新增的 C × C C \times C C × C 权重表在 DV 方向共享、不随切分缩小。d k = d v = 128 d_k = d_v = 128 d k = d v = 128 、C = 64 C = 64 C = 64 时不切 DV 约需 256 reg/thread,已超 255 上限直接编译不过;切到 block D V = 32 \text{block}_{DV} = 32 block D V = 32 降至约 112 reg。代价是 Dtri 被每个 b v bv b v block 各算一遍,属纯冗余但不占额外寄存器。
有效记忆长度暴露了一个内在矛盾 :α \alpha α 的数值安全区间(接近 1)与实际有用区间(提供遗忘能力)方向相反。α = 0.9 \alpha = 0.9 α = 0.9 的有效记忆仅 66 token,连一个 chunk 都刚覆盖;而 α ≥ 0.99 \alpha \ge 0.99 α ≥ 0.99 才落在因式分解仍然安全的区间。常数门控下两者尚可兼顾,但这是假设不是保证–部分实现敢做因式分解,正是隐含依赖了「门控总是接近 1」。
可迁移的启示 :代数上等价的两种写法,数值行为可以完全不同。γ i / γ j \gamma_i / \gamma_j γ i / γ j 与 γ i ⋅ ( 1 / γ j ) \gamma_i \cdot (1/\gamma_j) γ i ⋅ ( 1/ γ j ) 在实数域相等,但前者(j ≤ i j \le i j ≤ i 时)恒不大于 1、后者上界随块长指数增长。论文把衰减写成比值而非乘积并不是记号习惯,而是保证数值安全的必要安排 –那套 ⋅ ← \overleftarrow{\cdot} ⋅ / ⋅ → \overrightarrow{\cdot} ⋅ 箭头记号就是在提醒读者每个衰减都有明确的基准点(chunk 首或 chunk 末)。实现时不要对论文公式做看似等价的代数改写,先问改写后的中间量落在什么范围。
参考 :
Gated Delta Networks(GDN):arXiv:2412.06464,ICLR 2025。§2.1 以 Mamba2 为例给出带衰减的线性递推 S t = α t S t − 1 + v t k t ⊺ \mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^\intercal S t = α t S t − 1 + v t k t ⊺ 、累积衰减积 γ 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} ⋅ 箭头记号
Mamba2:Transformers are SSMs(Dao & Gu, 2024),arXiv:2405.21060。衰减项 α t \alpha_t α t 与状态空间对偶性(SSD)的出处
RetNet:arXiv:2307.08621。常数(数据无关)衰减的代表,即本文实现对应的形态
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 实测数据待补。