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

上一篇实现了不带任何遗忘机制的 chunked 线性注意力,状态单调累加。本篇引入第一个衰减因子:把递推式改为 St=αSt1+vtkt\mathbf{S}_t = \alpha \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^\intercalα\alpha 是一个标量常数。

这个衰减项来自哪里?Vanilla 线性注意力(St=St1+vtkt\mathbf{S}_t = \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^\intercal)在语言建模上远不如 Transformer,为此 Mamba2(Dao & Gu, 2024a)引入了一个数据依赖的逐步遗忘门 αt(0,1)\alpha_t \in (0, 1) 来有选择地丢弃历史信息。本文取它的常数特例 αtα\alpha_t \equiv \alpha–这一档对应的正是 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α\alpha_t \equiv \alpha
γj=i=1jαi\gamma_j = \prod_{i=1}^{j}\alpha_i 累积衰减积 常数情形下 γj=αj\gamma_j = \alpha^{\,j}
Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 衰减感知因果掩码 iji \ge j 时有值,否则 0
CC chunk 长度 代码里对应 BC / blk
[t][t] chunk 序号 Q[t]\mathbf{Q}_{[t]} 即第 tt
r[1,C]r \in [1, C] chunk 位置(1-based) q[t]r:=qtC+r\bm{q}_{[t]}^r := \bm{q}_{tC+r}

两个必须说清的点:

一、γ\gamma 是累积积,不是单步衰减。 这是阅读 GDN 时最容易混淆的一处–很多二次资料把 γ\gamma 当成逐步遗忘系数,但论文里逐步系数是 αt\alpha_tγj\gamma_j 是它们的前缀积。

二、累积积按 chunk 重置,不是从序列起点算。 论文式 (1) 的脚注明确写了 γ[t]j=j=tC+1tC+jαj\gamma_{[t]}^{j} = \prod_{j=tC+1}^{tC+j}\alpha_j,并自认“略微滥用了 γ\gamma 的记号”–每个 chunk 从自己的第一个位置重新开始累乘。这不是细节:若从序列起点算,γj\gamma_j 会随 jj 单调下溢到零(§2.5 有实测),而按 chunk 重置后指数永远不超过 CC,这才是分块算法数值可控的前提。


1. 递推式与分块重写

1.0 Mamba2 的遗忘门:衰减项的来处

上一篇实现的是 vanilla 线性注意力(Katharopoulos et al., 2020),状态只增不减:

St=St1+vtktRdv×dk,ot=StqtRdv\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}

它在语言建模上明显弱于 Transformer。GDN 论文 §2.1 给出的诊断很直接:缺少遗忘历史信息的手段。以 Mamba2(Dao & Gu, 2024a)为例,补救方式是在状态上乘一个逐步衰减项(论文原话是 “up to specific parameterization”,即忽略具体参数化细节):

St=αtSt1+vtkt,ot=Stqt\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

其中 αt(0,1)\alpha_t \in (0, 1)tt 变化的、数据依赖的标量衰减项(a data-dependent scalar-valued decay term that varies with tt)。这一个乘法带来两件事:

  • 有界性。无衰减时 St\|\mathbf{S}_t\|tt 无界增长(每步加一个外积);有了 αt<1\alpha_t < 1,状态范数被压在一个稳态附近。
  • 选择性αt0\alpha_t \to 0 可以快速清空状态(适合上下文切换),αt1\alpha_t \to 1 则保持记忆。Mamba 把这类机制称为 selective mechanism,本质是 gated RNN 里的遗忘门在矩阵值状态上的推广。

这个递推结构不是 Mamba2 独有的,同样出现在 Gated RFA、xLSTM、Gated RetNet 中。而当 αt\alpha_t 与数据无关、退化为常数时,该形式就是 RetNet 与 Lightning-Attention。

本文取的正是这个常数情形:

形态 衰减项 对应架构
无衰减 vanilla 线性注意力(上一篇)
常数标量 αtα\alpha_t \equiv \alpha RetNet / Lightning-Attention(本文)
数据依赖标量 αt=f(xt)\alpha_t = f(x_t) Mamba2 / GLA / Gated RetNet

本文取常数是为了简化:α\alpha 作为编译期常量,权重表可以预计算,kernel 改动最小。

1.1 本文的递推式与显式求和

本文在上一篇的基础上给状态加入一个遗忘系数 α(0,1]\alpha \in (0, 1]

St=αSt1+vtktRdv×dk,ot=StqtRdv\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}

展开成显式求和,注意每个 viki\bm{v}_i\bm{k}_i^{\intercal} 被后续每一步各乘一次 α\alpha,从 iitt 共乘 tit - i 次:

St=itαtiviki,ot=itαtivi(kiqt)\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)

用累积积写就是论文 §2.1 的形式(γj=αj\gamma_j = \alpha^j,所以 γt/γi=αti\gamma_t/\gamma_i = \alpha^{t-i}):

ot=itγtγivi(kiqt)\bm{o}_t = \sum_{i \le t} \frac{\gamma_t}{\gamma_i}\, \bm{v}_i\,(\bm{k}_i^{\intercal}\bm{q}_t)

对比无衰减时的 ot=itvi(kiqt)\bm{o}_t = \sum_{i \le t} \bm{v}_i(\bm{k}_i^{\intercal}\bm{q}_t),唯一变化是每一项多了权重 αti\alpha^{t-i}–距离越远权重越小,这就是「衰减」的含义。α=1\alpha = 1 时权重恒为 1,退化为不遗忘。

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

沿用上一篇的分块框架,序列按 CC 切分为 NCNC 个 chunk。把求和拆成跨块与块内两部分,衰减权重会分别落到三个位置。沿用论文的 chunk 内局部下标 r[1,C]r \in [1, C]

位置一–块内下三角(论文的 Γ[t]\Gamma_{[t]}。同块内 jij \le i,相对距离就是局部下标之差:

(Γ[t])ij=γ[t]iγ[t]j={αijji0j>i(\Gamma_{[t]})_{ij} = \frac{\gamma_{[t]}^{\,i}}{\gamma_{[t]}^{\,j}} = \begin{cases} \alpha^{\,i-j} & j \le i \\ 0 & j > i \end{cases}

上一篇这里是 0/1 因果掩码 M\mathbf{M},本文变成衰减感知掩码 Γ\Gamma注意权重全部落在 (0,1](0, 1] 区间–对角线是 α0=1\alpha^0 = 1,左下角最小值是 αC1\alpha^{C-1}

位置二–每块写入状态时的块尾对齐(论文的 k\overrightarrow{\bm{k}}。跨块状态需要统一的时间基准,取所属 chunk 的末尾。第 rr 个 token 的贡献衰减到块末要乘 γ[t]C/γ[t]r=αCr\gamma_{[t]}^{C}/\gamma_{[t]}^{r} = \alpha^{\,C-r}

k[t]r=γ[t]Cγ[t]rk[t]r,Δ[t]=V[t]K[t]Rdv×dk\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}

这个权重从哪来?用 C=4C = 4 手推一遍最清楚。块内每走一格执行一次递推,S[t]\mathbf{S}_{[t]} 为进块状态、S[t+1]\mathbf{S}_{[t+1]} 为出块状态:

Sr=αrSr1+vrkr,S0=S[t],SC=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]}

正向走一遍没什么可看的。有意思的是站在 S4\mathbf{S}_4 这端逐层倒代换,把中间状态一个个拆掉:

S4=α4S3+v4k4=α4α3S2+α4v3k3+v4k4=α4α3α2S1+α4α3v2k2+α4v3k3+v4k4=α1α2α3α4S[t]+α2α3α4v1k1+α3α4v2k2+α4v3k3+v4k4\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}

盯着最后一行的系数:v1\bm{v}_1 带的是 α2α3α4\alpha_2\alpha_3\alpha_4注意没有 α1\alpha_1,因为它自己就是第 1 格写入的,写完后一路经受的是"身后"三道门;历史状态 S[t]\mathbf{S}_{[t]} 则带满 α1α2α3α4\alpha_1\alpha_2\alpha_3\alpha_4没有任何一项的因子个数取决于它的绝对位置,全部取决于「离第 CC 格还差几步」。

用累积积改写,公共前缀立刻现形:α2α3α4=γ4/γ1\alpha_2\alpha_3\alpha_4 = \gamma_4/\gamma_1α3α4=γ4/γ2\alpha_3\alpha_4 = \gamma_4/\gamma_2α4=γ4/γ3\alpha_4 = \gamma_4/\gamma_3,而 v4\bm{v}_4 那项是 γ4/γ4=1\gamma_4/\gamma_4 = 1。于是打包收口:

S[t+1]=γ[t]CS[t]+r=1Cγ[t]Cγ[t]rvrkr\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}

这就是 k[t]r\overrightarrow{\bm{k}_{[t]}^{r}} 的全部内容,也就是论文说的「decaying each vector to the last position」。

位置三–跨块状态递推与 query 侧乘累积衰减积(论文的 S\overrightarrow{\mathbf{S}}q\overleftarrow{\bm{q}}。相邻块之间隔了整块 CC 步,因此状态递推是:

S[t]=γ[t]CS[t]=αCS[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]}}

而第 rr 个 token 读取这个状态时,要衰减到本块首位置的基准上,乘 γ[t]r=αr\gamma_{[t]}^{r} = \alpha^{r}

q[t]r=γ[t]rq[t]r,O[t]=Q[t]S[t]+(Q[t]K[t]Γ[t])V[t]RC×dv\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}

这正是论文式 (1)。三个位置的权重汇总–注意每个权重本质上都是门控的累积连乘,常数情形下才收成 α\alpha 的幂:

位置 论文记号 一般形式(累积积) 常数特例 取值范围 作用对象
块内下三角 (Γ[t])ij(\Gamma_{[t]})_{ij} γ[t]iγ[t]j=r=j+1iαr\dfrac{\gamma_{[t]}^{\,i}}{\gamma_{[t]}^{\,j}} = \prod_{r=j+1}^{i}\alpha_r αij\alpha^{\,i-j} [αC1, 1][\alpha^{C-1},\ 1] C×CC \times C 分数矩阵,逐元素
写入状态 k[t]r\overrightarrow{\bm{k}_{[t]}^{r}} γ[t]Cγ[t]r=u=r+1Cαu\dfrac{\gamma_{[t]}^{C}}{\gamma_{[t]}^{r}} = \prod_{u=r+1}^{C}\alpha_u αCr\alpha^{\,C-r} [αC1, 1][\alpha^{C-1},\ 1] K[t]\mathbf{K}_{[t]} 的行,逐行加权
跨块递推 S[t]\overrightarrow{\mathbf{S}_{[t]}} γ[t]C=u=1Cαu\gamma_{[t]}^{C} = \prod_{u=1}^{C}\alpha_u αC\alpha^{C} 标量 整个状态矩阵
读取状态 q[t]r\overleftarrow{\bm{q}_{[t]}^{r}} γ[t]r=u=1rαu\gamma_{[t]}^{r} = \prod_{u=1}^{r}\alpha_u αr\alpha^{r} [αC, α][\alpha^{C},\ \alpha] Q[t]\mathbf{Q}_{[t]} 的行,逐行加权

四个权重全部 1\le 1,指数均为非正。因果约束 iji \ge j 使 Γij=r=j+1iαr1\Gamma_{ij} = \prod_{r=j+1}^{i}\alpha_r \le 1,四个权重同理–论文的写法全程不需要物化任何大于 1 的量,这是 §3 讨论的前提。


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

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

2.1 累积衰减积的定义

把衰减退回 §1.0 那个一般形式–依赖数据的 αt(0,1)\alpha_t \in (0, 1),每个时刻的遗忘强度由输入决定(本文的常数 α\alphaαtα\alpha_t \equiv \alpha 的特例)。定义累积衰减积

γj=i=1jαi\gamma_j = \prod_{i=1}^{j} \alpha_i

γj\gamma_j 的含义是从起点衰减到第 jj 步的总折扣。有了它,递推式的展开可以写得非常紧凑。展开 St=αtSt1+vtkt\mathbf{S}_t = \alpha_t \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal}

St=it(r=i+1tαr)viki=itγtγiviki\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}

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

上面写的是全序列版本。到了分块算法里,累积积要按 chunk 重置,即论文式 (1) 脚注的 γ[t]j=j=tC+1tC+jαj\gamma_{[t]}^{\,j} = \prod_{j=tC+1}^{tC+j}\alpha_j。这一步不是为了好看–全序列累乘的 γj\gamma_j 会随 jj 单调下溢到零(§2.5 有实测:N=512N = 512α[0.5,0.9]\alpha \in [0.5, 0.9] 时 fp32 已进入非正规数),而按 chunk 重置后指数永远不超过 CC。后文出现 γ\gamma 时默认指 chunk 内的版本。

2.2 两种等价形式

代入 ot=Stqt\bm{o}_t = \mathbf{S}_t\bm{q}_t,同一个结果可以写成两种形式(论文 §2.1):

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

St=αtSt1+vtkt,ot=Stqt=itvi(γtγikiqt)\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)

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

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

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

[(QK)Γ]ij=γiγj(kjqi)(ji)\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)

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

Γ\Gamma 的双重身份:一个哈达玛积承载两重语义

Γ\Gamma 常被笼统理解为“一个带衰减的掩码”,但它实际上是两个正交语义的乘积,只是恰好能合并成一张表:

Γ=M因果性:能不能看D遗忘门:看得多清,Mij={1ij0i<j,Dij=γ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\mathbf{M}离散的、与数据无关的结构约束–token ii 不能看到未来的 j>ij > i。这是 causal 语言模型的硬性要求,α\alpha 取什么值都不影响它。
  • D\mathbf{D}连续的、由门控决定的权重–iijj 相距越远,γi/γj\gamma_i/\gamma_j 越小。这才是遗忘门起作用的地方。

将该矩阵可视化后,结构十分清晰:下三角,且每个值只由 token 序号差 iji-j 决定

图中三点值得对着公式确认:

  1. 对角线恒为 α0=1\alpha^0 = 1–自己看自己,零衰减。本文的 tril\operatorname{tril} 含对角(§1.2 位置一的 jij \le i),与图一致。
  2. 同一行自右向左指数缩小–同行内 ii 固定,jj 越小则间隔 iji-j 越大、权重越小。所以「衰减」在矩阵上表现为沿对角线方向的等值带iji-j 相同的格子取值相同。
  3. 上三角 j>ij > i 整块置 0–那是还没发生的 token。这就是「下三角」的全部含义。

面板 ② 用 α=0.5\alpha = 0.5C=4C = 4 给了可验算的数字:γ=(0.5, 0.25, 0.125, 0.0625)\gamma = (0.5,\ 0.25,\ 0.125,\ 0.0625),格 (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。整张表我用 fp64 复核过,与 αij\alpha^{i-j} 逐格一致、对角恒为 1。

图中「整张表没有一处幂运算,每格只是一次除法」一句还揭示了比值写法的另一重好处,同时解释了为什么论文写 γi/γj\gamma_i/\gamma_j 而不直接写 αij\alpha^{i-j}:后者只在 α\alpha 为常数时才成立,一旦门控逐 token 变化,「唯一底数」就不存在了,r=j+1iαr\prod_{r=j+1}^{i}\alpha_r 无法写成任何数的幂;而比值形式原样成立。本文因 α\alpha 为常数而两种写法皆可,论文则必须采用比值形式。

回到实现。 上面的分解还隐含一个容易忽略的陷阱:若把 M\mathbf{M}D\mathbf{D} 真的分开算再相乘,D\mathbf{D} 的上三角是大于 1 的i<ji<j 时指数为正)。取 α=0.9\alpha=0.9C=4C=4

D=[11.1111.2351.3720.911.1111.2350.810.911.1110.7290.810.91]   M  Γ=[10000.91000.810.9100.7290.810.91]\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}

C=4C = 4 时上三角最大才 1.372,但 C=64C = 64α=0.8\alpha = 0.8 时右上角是 0.863=1.27×1060.8^{-63} = 1.27 \times 10^{6},早已溢出 fp16。于是:

写法 中间量范围 结果
先构造完整 D\mathbf{D},再乘 M\mathbf{M} [αC1, α(C1)][\alpha^{C-1},\ \alpha^{-(C-1)}] fp16 下上三角可能已 inf,inf × 0 = NaN
只在 iji \ge j 处求值,否则直接置 0 (0, 1](0,\ 1] 安全

因此 kernel 里必须写成条件求值而非「算完再掩」(§5.1 的 T.if_then_else(j <= i, ...) 就是这个原因)–哪怕数学上 MD\mathbf{M} \odot \mathbf{D} 与「只算下三角」完全等价。M\mathbf{M} 的存在恰好保证了 Γ\Gamma 中每个有效元素的指数 ij0i-j \ge 0,这才是「Γij1\Gamma_{ij} \le 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}},把整块贡献衰减到块尾)。

参考实现一行 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}ΔA\Delta A 的累积和,AA 已经是负的),于是:

decay_states[r]=exp(logγ[t]Clogγ[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}}

这正是 §1.2 位置二的 k[t]r\overrightarrow{\bm{k}_{[t]}^{r}} 权重,一字不差。Mamba2 里 k\bm{k} 的角色由 B 承担、v\bm{v}x 承担,dt 是离散化步长(本文的常数情形没有这一项)。

kernel 里对应的三行是这样的:

1
2
3
4
5
6
7
8
9
p = 1.44269504                     # log2(e)

dA_cs_last[0] = dA_cumsum[batch_idx, bz, chunk_idx, chunk_size - 1] # log γ_C,循环外取一次
...
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) # 加权后才进 Tensor Core

四处细节和本文的结论逐条对上:

官方写法 为什么这么写
输入是 dA_cumsum(log 域累积和),不是 γ\gamma 本身 连乘会下溢:fp16 直接 cumprodC=128C = 128α[0.5, 0.9]\alpha \in [0.5,\ 0.9] 时归零。log 域改成加法就没这个问题
exp2(log γ_C - log γ_r)先在 log 域相减再取指数 相减的结果因 rCr \le C0\le 0exp2 输出恒 1\le 1。若先各自取指数再相除,就要物化 1/γr1/\gamma_r 这个大于 1 的量
p = 1.44269504exe^x 转成硬件 exp2 1.44269504=log2e1.44269504 = \log_2 e,于是 ex=2xlog2ee^x = 2^{x\log_2 e}exp2 有单指令实现,pow 通常展开成多条
dA_cs_lastT.Pipelined 循环外读一次 logγC\log\gamma_C 是整块共用的常量,每轮重取只是白读一次 shared memory

最值得注意的是第二行。它完全可以写成 exp2(log γ_C * p) / exp2(log γ_r * p)–数学上等价,还能把 γC\gamma_C 提到循环外省一次减法。但那就等于物化了 1/γr1/\gamma_r,正是 §3 实测会 NaN 的写法。官方选择在 log 域相减,指数结果因 rCr \le Clogγ\log\gamma 单调递减而恒 0\le 0exp2 输出恒不大于 1,不存在溢出可能。

顺带一个和本文互补的观察:官方 kernel 把 scale 直接乘到了 xt_local(即 v\bm{v} 侧)而非 B_sharedk\bm{k} 侧)。两者数学等价–衰减是标量,挂在哪一侧都行;选 x 是因为那一步本来就要做转置 xt_local[i,j] = x_local[j,i]将衰减乘法合并进转置的同一个 T.Parallel 中,省去一趟对 shared memory 的读写。这类「把逐元素操作合并进已有的数据搬运」是 tile 编程中常见的优化手法。


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

论文的写法(先 GEMM 再逐元素乘 Γ\Gamma)中间量恒在 (0,1](0,1],是数值安全的。但块内下三角 (Γ[t])ij=αij(\Gamma_{[t]})_{ij} = \alpha^{\,i-j} 存在一个看似有利的代数变形–指数可以拆开:

αij=αiαj\alpha^{\,i-j} = \alpha^{\,i} \cdot \alpha^{-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)

两种实现方式的差别:

方案 做法 块内额外开销 权重数值范围
比值形式(论文) 先 GEMM 得分数,再逐元素乘 Γ\Gamma 一次 C×CC \times C 逐元素乘 + 一张 C2C^2 权重表 (0,1](0, 1]
因式分解 先给 QQKK 的行分别乘上衰减因子,再单次 GEMM 两次 C×dkC \times d_k 逐行加权,无 C2C^2 开销 αj\alpha^{-j} 最大 α(C1)\alpha^{-(C-1)}

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

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

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

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

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

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

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

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

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

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

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

α\alpha C=64C=64 fp16 状态 C=128C=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 采用的做法。

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


4. 四层参考实现

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

参考 实现方式 验证目标
A 逐 token 递归 递推式定义 St=αSt1+vtkt\mathbf{S}_t = \alpha\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal}
B 分块向量化,三处衰减权重 §1.2 三个权重位置的推导
C 逐块独立重算,模拟 grid kernel 控制流
D 因式分解版块内计算 §3 两种写法的代数等价性

4.0 一个必要的转置:论文的 Rdv×dk\mathbb{R}^{d_v \times d_k} vs kernel 的 Rdk×dv\mathbb{R}^{d_k \times d_v}

这里要先交代一个容易造成混乱的差异。论文(以及本文 §1–§3 的全部推导)的状态是 SRdv×dk\mathbf{S} \in \mathbb{R}^{d_v \times d_k}

St=αtSt1+vtktRdv×dk,ot=StqtRdv\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}

注意是 vk\bm{v}\bm{k}^{\intercal}(value 在外、key 转置在内)、读出是 Sq\mathbf{S}\bm{q}(状态左乘 query)。而下面所有代码里的 S 存的都是它的转置

S =^ SRdk×dv\texttt{S} \ \widehat{=}\ \mathbf{S}^{\intercal} \in \mathbb{R}^{d_k \times d_v}

于是两边的写法逐行对应:

论文(SRdv×dk\mathbf{S} \in \mathbb{R}^{d_v \times d_k} 代码(S Rdk×dv\in \mathbb{R}^{d_k \times d_v}
St=αSt1+vtkt\mathbf{S}_t = \alpha\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} S = g*S + outer(k, v)
ot=Stqt\bm{o}_t = \mathbf{S}_t\bm{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 = g**C * S + Kd.T @ V
O[t]=Q[t]S[t]+\mathbf{O}_{[t]} = \overleftarrow{\mathbf{Q}_{[t]}}\mathbf{S}_{[t]}^{\intercal} + \dots acc = (Qb * w_query) @ S + ...

为何不直接按论文的方向存?因为 dk×dvd_k \times d_v 布局下,跨块读出就是 T.gemm(Q_s, S_s, acc_o)[C,dk]×[dk,dv][C, d_k] \times [d_k, d_v]不需任何 transpose 标志;若按论文方向存,每次读出都要 transpose_B=True。同理状态更新 VK\mathbf{V}^{\intercal}\overrightarrow{\mathbf{K}} 在转置布局下变成 KV\overrightarrow{\mathbf{K}}^{\intercal}\mathbf{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)
# 位置一:块内衰减下三角 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_decay = g ** (BC - 1 - i) # 位置二:写入状态时衰减到块尾
w_query = g ** (i + 1) # 位置三:读取状态时补上距上块末尾的步数

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

另外用逐 token 变化的 αtU(0.85, 0.999)\alpha_t \sim \mathcal{U}(0.85,\ 0.999) 复验过一遍(同规模,fp64),确认 §2.2 的矩阵并行形式在门控非常数时同样成立:矩阵形式 vs 向量形式相对 L2 2.03×10162.03 \times 10^{-16}Γij\Gamma_{ij} 取值范围 [0.628, 1.000][0.628,\ 1.000] 全部 1\le 1

前三行确认四份实现数学等价,误差均在 fp64 机器精度量级。第四行是退化检验:令 α=1\alpha = 1 应当精确回到上一篇的无衰减实现,这条验证能同时捕获三处权重中任何一处的指数写错–例如把 αCr\alpha^{\,C-r} 误写成 αCr+1\alpha^{\,C-r+1}α=1\alpha = 1 时两者都是 1,退化检验通过但 α=0.9\alpha = 0.9 时参考 B 与 A 不符。两条验证必须都做。


5. TileLang kernel 的改动

沿用上一篇 §6.4.2 的 grid 划分:(bv, bh)(bv,\ b\cdot h),序列轴不进 grid、退回 kernel 内的 T.Pipelined 顺序循环,状态作为 loop-carried fragment 常驻寄存器,每 chunk 只做一次 KVK^\top V。本文在此基础上新增一张 f32 权重表与两个权重向量。

先说清为什么必须切 DV。本文比上一篇多出一个 C×CC \times CDtri,而它在 DV 方向是共享的(衰减权重只跟 chunk 内位置有关,与 value 通道无关),因此不随 DV 切分而缩小。dk=dv=128d_k = d_v = 128C=64C = 64、128 线程下点一遍账:

fragment 不切 DV blockDV=32\text{block}_{DV} = 32
S_f [128,128][128,128] → 128 reg [128,32][128,32] → 32 reg
acc_o [64,128][64,128] → 64 reg [64,32][64,32] → 16 reg
Dtri + A 2×[64,64]2 \times [64,64] → 64 reg 64 reg(不变
合计 ≈ 256 reg/thread,已超 255 上限 ≈ 112 reg/thread

不切 DV 在这个配置下直接编译不过。 代价是 Dtri 被每个 bvbv block 各算一遍–DV/blockDV=4DV/\text{block}_{DV} = 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),
):
# grid:DV 切块 × (batch·head);序列轴不在这里
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) # loop-carried
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)
# 新增:衰减矩阵 Γ,f32 保存以避免 §3.2 的 fp16 下溢
Dtri = T.alloc_fragment([C, C], accum_dtype)
# 新增:两个长度为 C 的权重向量(位置二、位置三)
w_decay = T.alloc_fragment([C], accum_dtype) # α^(C-1-m),写入状态用
w_query = T.alloc_fragment([C], accum_dtype) # α^(i+1),读取状态用

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))       # log2(α),三处复用

# 位置一:Dtri[i,j] = α^(i-j) for j<=i, else 0
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)

# 位置二:写入状态时,第 m 行要衰减到块尾 → α^(C-1-m)
for m in T.Parallel(C):
w_decay[m] = T.exp2(lg * T.Cast(accum_dtype, C - 1 - m))

# 位置三:读取状态时,第 i 行距上块末尾 i+1 步 → α^(i+1)
for i in T.Parallel(C):
w_query[i] = T.exp2(lg * T.Cast(accum_dtype, i + 1))

三者的指数都取自 §1.2 的权重表:

变量 形状 元素 指数含义 取值范围
Dtri[i,j] C×CC \times C αij\alpha^{\,i-j}jij \le i,否则 0) 同块内两 token 的间隔 [αC1, 1][\alpha^{C-1},\ 1]
w_decay[m] CC αC1m\alpha^{\,C-1-m} mm 行到块尾的步数 [αC1, 1][\alpha^{C-1},\ 1]
w_query[i] CC αi+1\alpha^{\,i+1} ii 行到上块末尾的步数 [αC, α][\alpha^{C},\ \alpha]

注意 w_decayw_query方向相反w_decay[C-1] = α^0 = 1(块尾那一行本身就是基准,不需衰减),而 w_query[0] = α^1(块首那一行距上块末尾也有 1 步)。这两处若把指数写反或差一,α=1\alpha = 1 时完全看不出来–§4.4 的退化检验正是为此设计的。

指数用 exp2(log2(α) · k) 而非 pow(α, k)exp2log2 都有单指令硬件实现,而 pow 通常展开成多条指令。

DtriT.if_then_else 不能省成「先全算再掩」–如 §2.2 所述,j>ij > i 处的 αij\alpha^{i-j} 指数为正,C=64C = 64α=0.8\alpha = 0.8 时右上角已达 1.27×1061.27 \times 10^{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) # 全 DK
T.copy(K[bb, s0:s0+C, bh, :], K_s) # 全 DK
T.copy(V[bb, s0:s0+C, bh, dv0:dv0+block_DV], V_s) # 仅 dv 竖条

# ② 跨块项:此刻 S_f 仍是 S_prev(更新在循环末尾)
T.copy(S_f, S_s)
for i, d in T.Parallel(C, DK):
Q_s[i, d] *= w_query[i] # 位置三:query 侧
T.gemm(Q_s, S_s, acc_o, clear_accum=True)

# ③ 块内项:要用未加权的原始 Q,故重新载入
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 # 位置三:整块衰减 α^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

四处次序约束,写错任何一处都不会报错、只会算错:

  1. 状态更新必须排在输出写回之后SfS_f 在步骤②被读取时代表 S[t]\mathbf{S}_{[t]}(进块状态),本块自己的贡献由 tril\operatorname{tril} 那一项负责。提前更新就变成了 inclusive 前缀,块内项被重复计入。
  2. S_f *= alpha_pow_C 必须在 T.gemm(K_s, V_s, S_f) 之前。递推式是 S[t+1]=αCS[t]+Δ[t]\mathbf{S}_{[t+1]} = \alpha^{C}\mathbf{S}_{[t]} + \Delta_{[t]},先衰减旧状态再累加新贡献;若顺序颠倒,本块贡献会被多衰减一次。
  3. 步骤③必须重载 Q。步骤②已把 Q_s 原地乘上 αr\alpha^{r},而块内项用的是未加权的原始 Q[t]\mathbf{Q}_{[t]}。原地加权省了一个 buffer,代价是重载一次;若寄存器有余量,另开 Q_scaled 更安全。
  4. T.clear(S_f) 在循环外,clear_accum=True 在循环内。状态要跨迭代累加,只能循环外清一次;acc_oA 每 chunk 都是全新的,进了流水线循环就不能再用 T.clear(会被排到错误阶段),必须靠 clear_accum=True 在 MMA 那一刻覆盖累加器。

K_s 在循环末尾被原地乘过 w_decay,而步骤③用的是同一个 K_s–这里没有冲突,因为③在①之前执行。但若调整次序或复用,必须重新载入:本文的 K_s 残留内容已被乘过权重,比上一篇「残留内容不确定」更危险。

5.3 与上一篇 kernel 的改动汇总

位置 上一篇 本文 新增开销
grid (NS,H,B)(NS, H, B) (DV/blockDV, BH)(DV/\text{block}_{DV},\ B \cdot H) 序列不进 grid,状态常驻寄存器
权重表 Dtri[BC, BC] f32 fragment C2C^2 个 f32 寄存器
状态累加 T.gemm 直接累加 S_f *= alpha_pow_C 再 gemm dkdvd_k d_v 次乘法 / 轮
K 写入状态 原样 逐行乘 αCr\alpha^{\,C-r} C×dkC \times d_k 次乘法 / 轮
Q 读取状态 原样 逐行乘 αr\alpha^{r} C×dkC \times d_k 次乘法
块内掩码 if_then_else 置 0 逐元素乘 Dtri C2C^2 次乘法
Q 复用 步骤②③共用 步骤③必须重载 一次 shared 写入

切 DV 后 S_f[128,128][128,128] 降到 [128,32][128,32],寄存器压力最大的反而不是 Dtridk=dv=128d_k = d_v = 128C=64C = 64blockDV=32\text{block}_{DV} = 32 时逐项数:

fragment 上一篇 本文 变化
S_f [128,128][128,128] → 128 reg [128,32][128,32] → 32 reg 下降 96 reg
acc_o [64,128][64,128] → 64 reg [64,32][64,32] → 16 reg 下降 48 reg
Dtri + A 2×[64,64]2 \times [64,64] → 64 reg 新增 64 reg
合计 约 224 reg 约 112 reg ↓ 一半

切 DV 省下的寄存器刚好覆盖 Dtri 的开销,还有余量。C=128C = 128Dtri + A 升到 256 reg,即使切 DV 到 32 也超过 255 上限–那才是因式分解唯一真正有吸引力的地方(它不需要 C2C^2 权重表),但 §3.1 的 NaN 结论表明这个吸引力不成立。


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

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

L=ln103lnαL = \frac{\ln 10^{-3}}{\ln \alpha}

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

这个假设在标量、常数门控下可以接受–反正全局就一个 α\alpha,调到 0.99 以上就行。但它是一个隐含假设,不是数值保证:一旦门控变成数据依赖的(尤其是每个通道各自学一个),训练中总会有部分值降到 0.9 以下以实现快速遗忘,「全部接近 1」就不再成立,因式分解的溢出从例外变成常态。这也是 §3 实测中 α0.8\alpha \le 0.8 即溢出的原因。结论不变:不要做那个分解,不要依赖「门控总是接近 1」这个前提。

一句话总结α\alpha 的安全区间(接近 1)与有用区间(提供实际遗忘能力)方向相反,常数门控下两者尚可兼顾,但这个兼顾靠的是假设而非保证。


7. 数值验证

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

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

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


8. 总结

  1. 衰减项来自 Mamba2 的遗忘门。Vanilla 线性注意力的状态只增不减,语言建模上明显弱于 Transformer;Mamba2 的补救是乘一个数据依赖的标量 αt(0,1)\alpha_t \in (0,1),即 St=αtSt1+vtkt\mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal},同时带来有界性(状态范数不再无界增长)与选择性αt0\alpha_t \to 0 快速清空、1\to 1 保持记忆)。本文取其常数特例 αtα\alpha_t \equiv \alpha–这正是 RetNet / Lightning-Attention 的形态。
  2. 标量衰减不改变分块恒等式的结构,只在三个位置插入指数权重(即论文式 (2) 的三个箭头量):块内下三角 (Γ[t])ij=αij(\Gamma_{[t]})_{ij} = \alpha^{\,i-j}、写入状态时的块尾对齐 k\overrightarrow{\bm{k}}αCr\alpha^{\,C-r}、跨块递推 S\overrightarrow{\mathbf{S}}αC\alpha^{C} 与 query 侧 q\overleftarrow{\bm{q}}αr\alpha^{r}。四个权重全部 1\le 1,这是本文数值安全的根本原因。
  3. 累积衰减积 γj=i=1jαi\gamma_j = \prod_{i=1}^{j} \alpha_i 是统一两种形式的代数工具。它把区间连乘 r=i+1tαr\prod_{r=i+1}^{t} \alpha_r 化归为前缀量之比 γt/γi\gamma_t / \gamma_i,于是递归的向量形式与并行的矩阵形式 O=((QK)Γ)V\mathbf{O} = ((\mathbf{Q}\mathbf{K}^{\intercal}) \odot \Gamma)\mathbf{V} 可以相互转换(即 Mamba2 所谓的状态空间对偶性,实测相对 L2 2.03×10162.03 \times 10^{-16})。需注意论文中 αt\alpha_t 为单步衰减、γj\gamma_j 为累积积,二者不可混淆;常数情形下 γj=αj\gamma_j = \alpha^{\,j},此时 Γij=αij\Gamma_{ij} = \alpha^{\,i-j}。另外累积积按 chunk 重置(论文式 (1) 脚注),不是从序列起点算。
  4. 论文的衰减矩阵 Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_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{\bm{q}}k\overrightarrow{\bm{k}}S\overrightarrow{\mathbf{S}} 把「衰减到首 / 末位置」直接编码在箭头方向上,基准点选得当就能保证所有指数非正。
  5. 真正的陷阱是把比值因式分解成 γi(1/γj)\gamma_i \cdot (1/\gamma_j)。这个分解能把两步合成单次 GEMM 并免去 C2C^2 权重表,但要求物化 1/γj1/\gamma_j,把中间量从 (0,1](0,1] 推到 [1, 1/γC][1,\ 1/\gamma_C]。实测(C=64C=64,fp16 存储加 fp32 累加):α0.9\alpha \ge 0.9 时两者精度相当(4.14.6×1044.1\text{--}4.6 \times 10^{-4}),α0.8\alpha \le 0.8 时分解写法产生 NaNα=0.8\alpha = 0.8α63=1.27×106\alpha^{-63} = 1.27 \times 10^{6},已远超 fp16 上限 65504。data-dependent αt[0.8, 0.95]\alpha_t \in [0.8,\ 0.95]C=128C = 1281/γj1/\gamma_j 最大达 3.40×1083.40 \times 10^{8},同样溢出。结论是保留论文的比值形式,不做分解。
  6. 累积积本身也必须在 log 域计算。fp16 直接 cumprodC=128C = 128α[0.5, 0.9]\alpha \in [0.5,\ 0.9] 时已完全下溢为 0,fp32 在 C=512C = 512 时进入非正规数区间;logγj=ijlogαi\log \gamma_j = \sum_{i \le j} \log \alpha_i 是线性增长的负数,表示范围安全。这就是门控全程存 log 值、用 exp2 还原的原因。
  7. 退化检验是本文新增的关键验证手段:令 α=1\alpha = 1 应精确回到上一篇实现(实测相对 L2 1.62×10161.62 \times 10^{-16})。但它无法单独定案–三处权重中任何一处的指数偏移在 α=1\alpha = 1 时都不可见,必须同时做 α=0.9\alpha = 0.9 的四层参考互验(实测 1.822.03×10161.82\text{--}2.03 \times 10^{-16})。
  8. 四处次序约束(写错都不报错、只算错):状态更新必须排在输出写回之后(否则 inclusive 前缀导致块内项重复计入);S_f *= alpha_pow_C 必须在 T.gemm 之前(否则本块贡献被多衰减一次);步骤③必须重载 Q(步骤②已把 Q_s 原地乘了 αr\alpha^{r});T.clear(S_f) 在流水线循环外、clear_accum=True 在循环内。
  9. 必须切 DV:本文新增的 C×CC \times C 权重表在 DV 方向共享、不随切分缩小。dk=dv=128d_k = d_v = 128C=64C = 64 时不切 DV 约需 256 reg/thread,已超 255 上限直接编译不过;切到 blockDV=32\text{block}_{DV} = 32 降至约 112 reg。代价是 Dtri 被每个 bvbv block 各算一遍,属纯冗余但不占额外寄存器。
  10. 有效记忆长度暴露了一个内在矛盾α\alpha 的数值安全区间(接近 1)与实际有用区间(提供遗忘能力)方向相反。α=0.9\alpha = 0.9 的有效记忆仅 66 token,连一个 chunk 都刚覆盖;而 α0.99\alpha \ge 0.99 才落在因式分解仍然安全的区间。常数门控下两者尚可兼顾,但这是假设不是保证–部分实现敢做因式分解,正是隐含依赖了「门控总是接近 1」。

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


参考

  • Gated Delta Networks(GDN):arXiv:2412.06464,ICLR 2025。§2.1 以 Mamba2 为例给出带衰减的线性递推 St=αtSt1+vtkt\mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^\intercal、累积衰减积 γ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} 箭头记号
  • Mamba2:Transformers are SSMs(Dao & Gu, 2024),arXiv:2405.21060。衰减项 αt\alpha_t 与状态空间对偶性(SSD)的出处
  • RetNet:arXiv:2307.08621。常数(数据无关)衰减的代表,即本文实现对应的形态
  • 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 实测数据待补。