ggaaooppeenngg

为什么计算机科学是无限的但生命是有限的

线性注意力的线性代数前置知识

1. Householder 矩阵

1.1 定义

uRn\bm{u} \in \mathbb{R}^nuu=1\bm{u}^{\intercal}\bm{u} = 1(即 u2=1\|\bm{u}\|_2 = 1),定义

H=I2uuRn×n\mathbf{H} = \mathbf{I} - 2\bm{u}\bm{u}^{\intercal} \in \mathbb{R}^{n \times n}

题外一句:Householder 矩阵也是做 QR 分解的重要工具——逐列用反射把对角线以下的元素打成零,就得到 A=QR\mathbf{A} = \mathbf{Q}\mathbf{R},比 Gram-Schmidt 数值上稳得多。

1.1.1 矩阵形式是怎么被「提」出来的

定义式 H=I2uu\mathbf{H} = \mathbf{I} - 2\bm{u}\bm{u}^{\intercal} 常被当成天降公式,其实它只是把「反射两步走」这句话提取公因子的结果。取 n=2n = 2u=e1=[1,0]\bm{u} = \bm{e}_1 = [1, 0]^{\intercal}x=[3,2]\bm{x} = [3, 2]^{\intercal} 走一遍:

几何上:u=e1\bm{u} = \bm{e}_1 时镜面就是 x2x_2 轴。从 x\bm{x} 出发沿 u-\bm{u} 方向走一步 (ux)u-(\bm{u}^{\intercal}\bm{x})\bm{u} 落到镜面上(此时第一坐标归零),再走同样的一步就到 x\bm{x}'。两步等长,这正是系数 2 的来源。

推导上唯一需要的一步是 ux\bm{u}^{\intercal}\bm{x}标量,标量与向量相乘可换序,所以 (ux)u=u(ux)(\bm{u}^{\intercal}\bm{x})\bm{u} = \bm{u}(\bm{u}^{\intercal}\bm{x})。换序之后 x\bm{x} 变成右端公因子,反着用分配律提出来:

x=x2(ux)u=x2u(ux)=(I2uu)x=Hx\bm{x}' = \bm{x} - 2(\bm{u}^{\intercal}\bm{x})\bm{u} = \bm{x} - 2\bm{u}(\bm{u}^{\intercal}\bm{x}) = \big(\mathbf{I} - 2\bm{u}\bm{u}^{\intercal}\big)\bm{x} = \mathbf{H}\bm{x}

换序是全部的技巧所在(ux)u(\bm{u}^{\intercal}\bm{x})\bm{u}x\bm{x} 被夹在中间提不出来,写成 u(ux)\bm{u}(\bm{u}^{\intercal}\bm{x})uu\bm{u}\bm{u}^{\intercal} 自然聚成一个矩阵。

1.2 对称性证明

H=(I2uu)=I2(uu)=I2uu=H\mathbf{H}^{\intercal} = (\mathbf{I} - 2\bm{u}\bm{u}^{\intercal})^{\intercal} = \mathbf{I}^{\intercal} - 2(\bm{u}\bm{u}^{\intercal})^{\intercal} = \mathbf{I} - 2\bm{u}\bm{u}^{\intercal} = \mathbf{H} \qquad \blacksquare

关键一步用的是乘积转置法则 (AB)=BA(\mathbf{A}\mathbf{B})^{\intercal} = \mathbf{B}^{\intercal}\mathbf{A}^{\intercal}(转置后因子次序颠倒),以及 (A)=A(\mathbf{A}^{\intercal})^{\intercal} = \mathbf{A}

(uu)=(u)u=uu(\bm{u}\bm{u}^{\intercal})^{\intercal} = (\bm{u}^{\intercal})^{\intercal}\bm{u}^{\intercal} = \bm{u}\bm{u}^{\intercal}

外积矩阵因此天然对称。

1.3 广义形式:把归一化收进一个系数 τ\tau

H=I2uu\mathbf{H} = \mathbf{I} - 2\bm{u}\bm{u}^{\intercal} 要求 u2=1\|\bm{u}\|_2 = 1。想对任意非零 u\bm{u} 都能写,把归一化带在系数里就行:

H=I2uuuu=I+τuu,τ:=2uu\mathbf{H} = \mathbf{I} - \frac{2\bm{u}\bm{u}^{\intercal}}{\bm{u}^{\intercal}\bm{u}} = \mathbf{I} + \tau\,\bm{u}\bm{u}^{\intercal}, \qquad \tau := -\frac{2}{\bm{u}^{\intercal}\bm{u}}

τ\tau 当成一个普通标量系数写在外积前面,而不是留在分母上,后面的推导会干净很多。u2=1\|\bm{u}\|_2 = 1τ=2\tau = -2,退回 §1.1 的定义。

系数取 2/(uu)-2/(\bm{u}^{\intercal}\bm{u}) 不是美学选择,而是被「H\mathbf{H} 是反射」逆推出来的。反射必须满足对合性 H2=I\mathbf{H}^2 = \mathbf{I}(照两次镜子回到原处),展开:

(I+τuu)2=I+2τuu+τ2u(uu)标量u=I+τ(2+τuu)uu(\mathbf{I} + \tau\bm{u}\bm{u}^{\intercal})^2 = \mathbf{I} + 2\tau\bm{u}\bm{u}^{\intercal} + \tau^2\bm{u}\underbrace{(\bm{u}^{\intercal}\bm{u})}_{\text{标量}}\bm{u}^{\intercal} = \mathbf{I} + \tau\big(2 + \tau\,\bm{u}^{\intercal}\bm{u}\big)\bm{u}\bm{u}^{\intercal}

要它等于 I\mathbf{I},就要 τ(2+τuu)=0\tau(2 + \tau\bm{u}^{\intercal}\bm{u}) = 0,除去平常解 τ=0\tau = 0 只剩

τ=2uu\tau = -\frac{2}{\bm{u}^{\intercal}\bm{u}}

所以 τ\tau 取任何值时 H\mathbf{H} 都对称(外积天然对称,与系数无关),但只有这一个 τ\tau 才让它成为对合的正交反射、才有「沿法向走两步等长」的几何意义。其余 τ\tau 给出的是一般的秩一修正 I+τuu\mathbf{I} + \tau\bm{u}\bm{u}^{\intercal},不再是等距变换;但 §3 的推导对 τ\tau 不做任何限制,对这类一般秩一修正同样成立。

2. WY 表达式:把一串反射压成一个低秩修正

结论先给pp 个 Householder 反射的乘积,永远能写成单位矩阵加一个秩不超过 pp 的修正

H1H2Hp=I+WY,W,YRn×p\mathbf{H}_1\mathbf{H}_2\cdots\mathbf{H}_p = \mathbf{I} + \mathbf{W}\mathbf{Y}^{\intercal}, \qquad \mathbf{W},\,\mathbf{Y} \in \mathbb{R}^{n \times p}

证明用归纳构造,一列一列把反射「吸收」进 W,Y\mathbf{W},\mathbf{Y}

2.1 基例 p=1p = 1

H1=I2u1u1=I+W1Y1\mathbf{H}_1 = \mathbf{I} - 2\bm{u}_1\bm{u}_1^{\intercal} = \mathbf{I} + \mathbf{W}_1\mathbf{Y}_1^{\intercal},取

W1=[2u1],Y1=[u1]\mathbf{W}_1 = \big[-2\bm{u}_1\big], \qquad \mathbf{Y}_1 = \big[\bm{u}_1\big]

均为 n×1n \times 1。命题对 p=1p = 1 成立。

2.2 归纳步:一列一列「吸收」反射

归纳假设:已存在 Wk,YkRn×k\mathbf{W}_k,\,\mathbf{Y}_k \in \mathbb{R}^{n \times k} 使 H1Hk=I+WkYk\mathbf{H}_1\cdots\mathbf{H}_k = \mathbf{I} + \mathbf{W}_k\mathbf{Y}_k^{\intercal}。要证 H1HkHk+1\mathbf{H}_1\cdots\mathbf{H}_k\mathbf{H}_{k+1} 也能写成 I+Wk+1Yk+1\mathbf{I} + \mathbf{W}_{k+1}\mathbf{Y}_{k+1}^{\intercal},其中 Wk+1,Yk+1Rn×(k+1)\mathbf{W}_{k+1},\mathbf{Y}_{k+1} \in \mathbb{R}^{n \times (k+1)}

记新反射方向 u=uk+1\bm{u} = \bm{u}_{k+1},并记

w:=2(I+WkYk)u\bm{w} := -2\,(\mathbf{I} + \mathbf{W}_k\mathbf{Y}_k^{\intercal})\,\bm{u}

则把 (I+WkYk)(I2uu)(\mathbf{I} + \mathbf{W}_k\mathbf{Y}_k^{\intercal})(\mathbf{I} - 2\bm{u}\bm{u}^{\intercal}) 乘开后的四项可以合并(用 (I+WkYk)uu=[(I+WkYk)u]u(\mathbf{I} + \mathbf{W}_k\mathbf{Y}_k^{\intercal})\bm{u}\bm{u}^{\intercal} = \big[(\mathbf{I} + \mathbf{W}_k\mathbf{Y}_k^{\intercal})\bm{u}\big]\bm{u}^{\intercal}u\bm{u}^{\intercal} 原样留在最右)

I+WkYk2uu2WkYkuu=I+WkYk+wu\mathbf{I} + \mathbf{W}_k\mathbf{Y}_k^{\intercal} - 2\bm{u}\bm{u}^{\intercal} - 2\mathbf{W}_k\mathbf{Y}_k^{\intercal}\bm{u}\bm{u}^{\intercal} = \mathbf{I} + \mathbf{W}_k\mathbf{Y}_k^{\intercal} + \bm{w}\bm{u}^{\intercal}

于是

H1Hk+1=I+[Wk    w][Yk    u]\mathbf{H}_1\cdots\mathbf{H}_{k+1} = \mathbf{I} + \big[\,\mathbf{W}_k \;\; \bm{w}\,\big]\big[\,\mathbf{Y}_k \;\; \bm{u}\,\big]^{\intercal}

Wk+1=[Wk, w]\mathbf{W}_{k+1} = [\mathbf{W}_k,\ \bm{w}]Yk+1=[Yk, uk+1]\mathbf{Y}_{k+1} = [\mathbf{Y}_k,\ \bm{u}_{k+1}]。定理得证。\blacksquare

3. compact WY:把纠缠收进一个上三角因子 T\mathbf{T}

WY 有个不便:W\mathbf{W} 的列互相纠缠。由 §2.2 的递推回代可得 wt=2H1Ht1ut\bm{w}_t = -2\mathbf{H}_1\cdots\mathbf{H}_{t-1}\bm{u}_t,每根列都带着一个前缀乘积,W\mathbf{W} 里存的已经不是原始法向了。

1989 年 Schreiber–Van Loan 的 compact WY 把全部纠缠塞进一个 p×pp \times p 小矩阵,让两个 n×pn \times p 因子里只剩下原始法向:

H1Hp=I+YTY,Y=[u1,,up],TRp×p 上三角\mathbf{H}_1\cdots\mathbf{H}_p = \mathbf{I} + \mathbf{Y}\mathbf{T}\mathbf{Y}^{\intercal}, \qquad \mathbf{Y} = [\,\bm{u}_1, \dots, \bm{u}_p\,], \quad \mathbf{T} \in \mathbb{R}^{p \times p}\ \text{上三角}

下面用 §1.3 的广义形式 H=I+τuu\mathbf{H} = \mathbf{I} + \tau\bm{u}\bm{u}^{\intercal} 写,ut\bm{u}_t 不需归一,τt\tau_t 也不限于 2/(utut)-2/(\bm{u}_t^{\intercal}\bm{u}_t),当成一串给定系数即可。

递推构造

T1=[τ1],Tk+1=[Tkτk+1Tk(Ykuk+1)0τk+1]\mathbf{T}_1 = [\tau_1], \qquad \mathbf{T}_{k+1} = \begin{bmatrix} \mathbf{T}_k & \tau_{k+1}\,\mathbf{T}_k\,(\mathbf{Y}_k^{\intercal}\bm{u}_{k+1}) \\ \bm{0} & \tau_{k+1} \end{bmatrix}

证明与 §2.2 是同一个机制,只是新增的一列不往 nn 维因子里塞,而往 T\mathbf{T} 里塞。记 u=uk+1\bm{u} = \bm{u}_{k+1}τ=τk+1\tau = \tau_{k+1},乘开:

(I+YkTkYk)(I+τuu)=I+YkTkYk+τuu+τYkTk(Yku)u(\mathbf{I} + \mathbf{Y}_k\mathbf{T}_k\mathbf{Y}_k^{\intercal})(\mathbf{I} + \tau\bm{u}\bm{u}^{\intercal}) = \mathbf{I} + \mathbf{Y}_k\mathbf{T}_k\mathbf{Y}_k^{\intercal} + \tau\bm{u}\bm{u}^{\intercal} + \tau\,\mathbf{Y}_k\mathbf{T}_k(\mathbf{Y}_k^{\intercal}\bm{u})\,\bm{u}^{\intercal}

三个修正项正好是 [Yk,u]\big[\mathbf{Y}_k,\, \bm{u}\big] 夹一个 (k+1)×(k+1)(k+1) \times (k+1) 矩阵再夹 [Yk,u]\big[\mathbf{Y}_k,\, \bm{u}\big]^{\intercal} 的展开结果:YkTkYk\mathbf{Y}_k\mathbf{T}_k\mathbf{Y}_k^{\intercal} 对应左上块,τuu\tau\bm{u}\bm{u}^{\intercal} 对应右下角,交叉项 τYkTk(Yku)u\tau\mathbf{Y}_k\mathbf{T}_k(\mathbf{Y}_k^{\intercal}\bm{u})\bm{u}^{\intercal} 对应右上列,左下为零:

=I+[Yk,u][TkτTk(Yku)0τ][Yk,u]= \mathbf{I} + \big[\mathbf{Y}_k,\, \bm{u}\big] \begin{bmatrix} \mathbf{T}_k & \tau\mathbf{T}_k(\mathbf{Y}_k^{\intercal}\bm{u}) \\ \bm{0} & \tau \end{bmatrix} \big[\mathbf{Y}_k,\, \bm{u}\big]^{\intercal}

Yk+1=[Yk, uk+1]\mathbf{Y}_{k+1} = [\mathbf{Y}_k,\ \bm{u}_{k+1}]Tk+1\mathbf{T}_{k+1} 如上,定理得证。\blacksquare

三个副产品直接从递推读出来:

  • T\mathbf{T} 上三角。左下块永远是 0\bm{0},上三角性逐层遗传。
  • 对角线就是系数diag(T)=(τ1,,τp)\mathrm{diag}(\mathbf{T}) = (\tau_1, \dots, \tau_p),单位向量情形下全是 2-2
  • 与 §2 的关系是 W=YT\mathbf{W} = \mathbf{Y}\mathbf{T}。两个表达式并不独立:compact WY 把 WY 里那些前缀乘积分解成「原始法向 ×\times 上三角系数」。纠缠没有消失,只是从 n×pn \times pW\mathbf{W} 里搬到了 p×pp \times pT\mathbf{T} 里,而 pnp \ll n

4. 另一条路:用正交条件直接解出 T\mathbf{T}

§3 的递推是串行的:要算 Tk+1\mathbf{T}_{k+1} 必须先有 Tk\mathbf{T}_k。Joffrain 等人(2006)指出 T\mathbf{T}闭式,不靠归纳、不靠前缀,一步得出。

为与文献记号一致,本节写成减号形式:

Q=H1Hp=IUSU,U=[u1,,up]\mathbf{Q} = \mathbf{H}_1\cdots\mathbf{H}_p = \mathbf{I} - \mathbf{U}\mathbf{S}\mathbf{U}^{\intercal}, \qquad \mathbf{U} = [\,\bm{u}_1, \dots, \bm{u}_p\,]

对照 §3:把那里的上三角因子记作 TSVL\mathbf{T}_{\text{SVL}},则 U=Y\mathbf{U} = \mathbf{Y}S=TSVL\mathbf{S} = -\mathbf{T}_{\text{SVL}},同为上三角;由 diag(S)=(2/u1u1,,2/upup)\mathrm{diag}(\mathbf{S}) = (2/\bm{u}_1^{\intercal}\bm{u}_1, \dots, 2/\bm{u}_p^{\intercal}\bm{u}_p) 全非零,且三角阵行列式等于对角元之积,知 S\mathbf{S} 可逆。本节的 T:=S1\mathbf{T} := \mathbf{S}^{-1}TSVL\mathbf{T}_{\text{SVL}} 差一个求逆和符号,不是同一个矩阵。以下设 U\mathbf{U} 列满秩。

4.1 可用的条件是正交,不是对合

先排除一个很容易误用的条件。每个 Hi\mathbf{H}_i 都对称(§1.2),但对称矩阵的乘积一般不对称

(H1H2)=H2H1=H2H1H1H2(\mathbf{H}_1\mathbf{H}_2)^{\intercal} = \mathbf{H}_2^{\intercal}\mathbf{H}_1^{\intercal} = \mathbf{H}_2\mathbf{H}_1 \ne \mathbf{H}_1\mathbf{H}_2

转置把因子次序翻了过来,而反射一般不交换。所以 Q=Q\mathbf{Q}^{\intercal} = \mathbf{Q} 不成立,Q2=I\mathbf{Q}^2 = \mathbf{I} 也不成立——两面镜子依次反射是一个旋转,转两次不回原处。真正能用的是正交性:正交矩阵的乘积仍正交,所以

QQ=I\mathbf{Q}^{\intercal}\mathbf{Q} = \mathbf{I}

好在下面的代数只需要这一条,结论一字不改。

4.2 展开正交条件

Q=IUSU\mathbf{Q} = \mathbf{I} - \mathbf{U}\mathbf{S}\mathbf{U}^{\intercal} 代入 QQ=I\mathbf{Q}^{\intercal}\mathbf{Q} = \mathbf{I}

(IUSU)(IUSU)=IUSUUSU+US(UU)SU=I(\mathbf{I} - \mathbf{U}\mathbf{S}^{\intercal}\mathbf{U}^{\intercal})(\mathbf{I} - \mathbf{U}\mathbf{S}\mathbf{U}^{\intercal}) = \mathbf{I} - \mathbf{U}\mathbf{S}^{\intercal}\mathbf{U}^{\intercal} - \mathbf{U}\mathbf{S}\mathbf{U}^{\intercal} + \mathbf{U}\mathbf{S}^{\intercal}(\mathbf{U}^{\intercal}\mathbf{U})\mathbf{S}\mathbf{U}^{\intercal} = \mathbf{I}

消去 I\mathbf{I},左提 U\mathbf{U}、右提 U\mathbf{U}^{\intercal}

U(SS+S(UU)S)U=0\mathbf{U}\big(-\mathbf{S}^{\intercal} - \mathbf{S} + \mathbf{S}^{\intercal}(\mathbf{U}^{\intercal}\mathbf{U})\mathbf{S}\big)\mathbf{U}^{\intercal} = \bm{0}

U\mathbf{U} 列满秩(Ux=0x=0\mathbf{U}\bm{x} = \bm{0} \Rightarrow \bm{x} = \bm{0}),因此括号内的 p×pp \times p 矩阵必须为零:

S+S=S(UU)S\mathbf{S}^{\intercal} + \mathbf{S} = \mathbf{S}^{\intercal}(\mathbf{U}^{\intercal}\mathbf{U})\mathbf{S}

S\mathbf{S} 可逆,左乘 S\mathbf{S}^{-\intercal}、右乘 S1\mathbf{S}^{-1},右端夹着的 S,S\mathbf{S}^{\intercal},\mathbf{S} 被抵消:

S(S+S)S1=UUS1+S=UU\mathbf{S}^{-\intercal}(\mathbf{S}^{\intercal} + \mathbf{S})\mathbf{S}^{-1} = \mathbf{U}^{\intercal}\mathbf{U} \qquad\Longrightarrow\qquad \mathbf{S}^{-1} + \mathbf{S}^{-\intercal} = \mathbf{U}^{\intercal}\mathbf{U}

T=S1\mathbf{T} = \mathbf{S}^{-1}(上三角阵的逆仍上三角),得到全部机关:

T+T=UU\mathbf{T} + \mathbf{T}^{\intercal} = \mathbf{U}^{\intercal}\mathbf{U}

4.3 一个方程定住整个 T\mathbf{T}

T\mathbf{T} 上三角、T\mathbf{T}^{\intercal} 下三角,两者相加时不同位置互不干扰,所以逐位置比对就能定住每一项(记 G=UU\mathbf{G} = \mathbf{U}^{\intercal}\mathbf{U},它对称):

  • 严格上三角(i<ji < j):T\mathbf{T}^{\intercal} 在此为零,所以 Tij=Gij=uiuj\mathbf{T}_{ij} = \mathbf{G}_{ij} = \bm{u}_i^{\intercal}\bm{u}_j
  • 对角线:两项相等,2Tii=Gii2\mathbf{T}_{ii} = \mathbf{G}_{ii},所以 Tii=12uiui\mathbf{T}_{ii} = \tfrac12\bm{u}_i^{\intercal}\bm{u}_i
  • 严格下三角:Tij=0\mathbf{T}_{ij} = 0(上三角的定义)。

写成一行:

T=striu(UU)+12diag(UU),Q=IUT1U\mathbf{T} = \mathrm{striu}(\mathbf{U}^{\intercal}\mathbf{U}) + \tfrac12\,\mathrm{diag}(\mathbf{U}^{\intercal}\mathbf{U}), \qquad \mathbf{Q} = \mathbf{I} - \mathbf{U}\mathbf{T}^{-1}\mathbf{U}^{\intercal}

其中 striu\mathrm{striu} 取严格上三角部分。换句话说,T\mathbf{T} 就是 Gram 矩阵 UU\mathbf{U}^{\intercal}\mathbf{U} 的「上三角一半」。单位向量时对角线全为 12\tfrac12S=T1\mathbf{S} = \mathbf{T}^{-1} 对角线全为 22,正对应 §1.1 的系数。T\mathbf{T} 由方程唯一确定,因此使 Q=IUSU\mathbf{Q} = \mathbf{I} - \mathbf{U}\mathbf{S}\mathbf{U}^{\intercal} 成立的上三角可逆 S\mathbf{S} 也唯一。

4.4 为什么这个闭式适合并行

两条路得到的是同一个 T\mathbf{T},但计算结构完全不同:

维度 §3 递推(Schreiber–Van Loan) §4 闭式(Joffrain)
依赖结构 Tk+1\mathbf{T}_{k+1} 依赖 Tk\mathbf{T}_k,严格串行 只依赖 U\mathbf{U},无链式依赖
主体运算 pp 次矩阵-向量(TkYku\mathbf{T}_k\mathbf{Y}_k^{\intercal}\bm{u} 一次 UU\mathbf{U}^{\intercal}\mathbf{U},即一个 GEMM
并行度 逐列展开 G\mathbf{G} 的所有元素可同时算
分块合并 需重跑递推 直接拼,见下

关键差异在于 UU\mathbf{U}^{\intercal}\mathbf{U} 是一个普通的 Gram 矩阵乘法:p2p^2 个内积互不依赖,能一次 GEMM 算完,而递推的第 k+1k+1 列必须等第 kk 列。代价是多了一个 p×pp \times p 上三角求逆(或三角求解),但 pnp \ll n,这部分开销可忽略。

更有用的是分块可合并。把 U=[U1,U2]\mathbf{U} = [\mathbf{U}_1, \mathbf{U}_2] 分两块,则

UU=[U1U1U1U2U2U1U2U2]T=[T1U1U20T2]\mathbf{U}^{\intercal}\mathbf{U} = \begin{bmatrix} \mathbf{U}_1^{\intercal}\mathbf{U}_1 & \mathbf{U}_1^{\intercal}\mathbf{U}_2 \\ \mathbf{U}_2^{\intercal}\mathbf{U}_1 & \mathbf{U}_2^{\intercal}\mathbf{U}_2 \end{bmatrix} \qquad\Longrightarrow\qquad \mathbf{T} = \begin{bmatrix} \mathbf{T}_1 & \mathbf{U}_1^{\intercal}\mathbf{U}_2 \\ \bm{0} & \mathbf{T}_2 \end{bmatrix}

两个子块的 T1,T2\mathbf{T}_1, \mathbf{T}_2 可以并行算,合并时只需补上跨块耦合块 U1U2\mathbf{U}_1^{\intercal}\mathbf{U}_2——这正是 LAPACK 分块 QR 里 dlarft 合并相邻 panel 的做法,也是把反射串映到 Tensor Core 上时想要的结构。

最后注意一个方向约定:上三角对应的是 H1H2Hp\mathbf{H}_1\mathbf{H}_2\cdots\mathbf{H}_p 这个顺序。反序乘积 HpH1\mathbf{H}_p\cdots\mathbf{H}_1 换成 T\mathbf{T}^{-\intercal},即下三角那一半;弄反会得到错的矩阵,实现时这是最容易踩的坑。

TileLang 实战:KDA 从零到一–Kimi Delta Attention

前三篇分别实现了无遗忘的 chunked 线性注意力、标量衰减、以及带删除项的 Gated DeltaNet。本篇是系列收尾:把 GDN 的标量αt\alpha_t 换成逐通道向量atRdk\bm{a}_t \in \mathbb{R}^{d_k}

St=St1Diag(at)(Iβtktkt)+βtvtkt\mathbf{S}_t = \mathbf{S}_{t-1}\operatorname{Diag}(\bm{a}_t)\big(\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal}\big) + \beta_t\bm{v}_t\bm{k}_t^{\intercal}

改动只有一处:αt\alpha_t 变成了 Diag(at)\operatorname{Diag}(\bm{a}_t)。但这一处把前三篇积累的所有便利拆掉了–衰减不再能从矩阵里外提,Γ\GammaC×CC\times C 变成 C×C×dkC \times C \times d_k,而累积衰减积的通道离散度会直接把 fp16 打穿。

本篇聚焦实现。KDA 的数学推导(递推式、逐通道 WY 表示、UT 变换、下界衰减与满秩门控的动机)见《KDA 的来龙去脉》§3–§4,这里不重复;本文只做一件事:把那些公式落成能跑的 kernel,并量化每一处数值边界

前三篇见 Chunked 线性注意力标量衰减Gated DeltaNet


0. 符号与三处变化

沿用前三篇:βt\beta_t 写入强度、CC 块长、[t][t] 块序号、r[1,C]r \in [1,C] 块内位置、SRdk×dv\mathbf{S}^\intercal \in \mathbb{R}^{d_k \times d_v} 为 kernel 存储布局。

本篇的核心量改为累积 log 衰减(沿用《来龙去脉》的记号):

γi=s=1igsRdk,gs=logas<0 (逐分量)\bm{\gamma}_i = \sum_{s=1}^{i}\bm{g}_s \in \mathbb{R}^{d_k}, \qquad \bm{g}_s = \log\bm{a}_s < \bm{0}\ \text{(逐分量)}

注意 γi\bm{\gamma}_i向量–这是与前两篇最本质的区别。前三篇的 γr\gamma^r 是标量,Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 是一张 C×CC \times C 的表;本篇 γiγj\bm{\gamma}_i - \bm{\gamma}_jdkd_k 维向量,"衰减矩阵"概念上是 C×C×dkC \times C \times d_k不可能物化

三处结构性变化:

GDN(标量门) KDA(逐通道门)
累积积 γr\gamma^r 标量,cumsum 后 CC 个数 γrRdk\bm{\gamma}_r \in \mathbb{R}^{d_k},cumsum 后 C×dkC \times d_k
衰减掩码 ΓRC×C\Gamma \in \mathbb{R}^{C\times C},可物化 概念上 C×C×dkC\times C\times d_k必须融进 GEMM
KK 矩阵 KKΓ\mathbf{K}\mathbf{K}^\intercal \odot \Gamma,衰减可外提 Mci=(kceγcγi)kiM_{ci} = (\bm{k}_c \odot e^{\bm{\gamma}_c - \bm{\gamma}_i})^\intercal\bm{k}_i,衰减长在内部

第三行是全篇的技术核心。


1. 衰减长在内部:问题与出路

1.1 朴素做法的代价

Mci=(kceγcγi)kiM_{ci} = (\bm{k}_c \odot e^{\bm{\gamma}_c - \bm{\gamma}_i})^\intercal\bm{k}_i 里的指数依赖 (c,i,d)(c, i, d) 三个下标。直接算就是 C2/2C^2/2 次长度 dkd_k 的加权内积,每次都要现算 dkd_kexp

这不是常数因子问题–它把一次 GEMM 变成了 C2/2C^2/2 个独立的向量运算,完全用不上 Tensor Core。C=64C = 64dk=128d_k = 128 时是 2016 次加权内积、约 26 万次 exp2

1.2 出路:指数可分离

关键观察是指数可以按通道拆开

eγcγi=eγceγie^{\bm{\gamma}_c - \bm{\gamma}_i} = e^{\bm{\gamma}_c} \odot e^{-\bm{\gamma}_i}

于是:

Mci=d(kc[d]eγc[d])K~+[c,d](ki[d]eγi[d])K~[i,d]=(K~+K~)ciM_{ci} = \sum_{d}\underbrace{\big(k_c[d]\,e^{\gamma_c[d]}\big)}_{\widetilde{K}^{+}[c,d]}\underbrace{\big(k_i[d]\,e^{-\gamma_i[d]}\big)}_{\widetilde{K}^{-}[i,d]} = \big(\widetilde{\mathbf{K}}^{+}\widetilde{\mathbf{K}}^{-\intercal}\big)_{ci}

一次 GEMM 解决,前置两个 C×dkC \times d_k 的逐元素加权。实测拆分与朴素计算的差异 1.94×10161.94 \times 10^{-16},代数上完全等价。

1.3 代价:eγie^{-\bm{\gamma}_i} 是大于 1 的量

第二篇讲过一个教训:把比值 γi/γj\gamma_i/\gamma_j 因式分解成 γi(1/γj)\gamma_i \cdot (1/\gamma_j) 会物化一个指数增长的量,fp16 下 α0.8\alpha \le 0.8 即 NaN。这里是同一个陷阱的逐通道版本eγi[d]1e^{-\gamma_i[d]} \ge 1,且随 ii 与通道衰减强度指数增长。

但本篇的处境和第二篇不同:那里有替代方案(保留比值形式),这里没有。逐通道衰减无法外提,不拆就用不上 Tensor Core。所以问题从"要不要拆"变成了"怎样让拆分在数值上安全"。

实测 eγe^{-\bm{\gamma}} 的最大值(C=64C = 64dk=128d_k = 128):

门控下界 maxeγ\max e^{-\bm{\gamma}} fp16 fp32
0.999 1.041.04 OK OK
0.99 1.481.48 OK OK
0.95 7.147.14 OK OK
0.9 54.9554.95 OK OK
0.8 4.20×1034.20 \times 10^{3} OK OK
0.5 3.04×10103.04 \times 10^{10} 溢出 OK

门控下界在 0.9 以上时,eγe^{-\bm{\gamma}} 连 fp16 都装得下。 这就把 KDA 论文里"下界衰减(lower-bounded decay)"这个设计从模型层面的技巧,变成了 kernel 能否用 Tensor Core 的前提条件。


2. 下界衰减:不是精度调优,是可行性前提

2.1 无下界时 fp16 全线归零

eγe^{-\bm{\gamma}} 会溢出,另一头 eγe^{\bm{\gamma}} 会下溢。而逐通道门让后者严重得多–总有一些通道学到很小的 aa,它们的 γ\gamma 累积得最快。

实测 mineγC\min e^{\bm{\gamma}_C} 与 fp16 下归零的通道数(dk=128d_k = 128):

a\bm{a} 采样区间 CC mineγC\min e^{\bm{\gamma}_C} fp16 归零通道
[0.9, 0.999][0.9,\ 0.999] 64 1.70×1021.70\times10^{-2} 0 / 128
[0.9, 0.999][0.9,\ 0.999] 128 4.84×1044.84\times10^{-4} 0 / 128
[0.5, 0.999][0.5,\ 0.999] 64 2.20×10112.20\times10^{-11} 118 / 128
[0.5, 0.999][0.5,\ 0.999] 128 1.59×10201.59\times10^{-20} 128 / 128
[0.1, 0.999][0.1,\ 0.999] 64 7.39×10287.39\times10^{-28} 128 / 128
[0.01, 0.999][0.01,\ 0.999] 128 6.44×10656.44\times10^{-65} 128 / 128

C=128C = 128、下界 0.5 时全部 128 个通道在 fp16 下归零–状态被彻底清空,kernel 输出恒为块内项,跨块信息完全丢失。

加下界后:

下界 mineγC\min e^{\bm{\gamma}_C} fp16 归零
1.84×10661.84\times10^{-66} 128 / 128
0.5 9.07×10329.07\times10^{-32} 128 / 128
0.9 1.72×1061.72\times10^{-6} 0 / 128
0.95 1.43×1031.43\times10^{-3} 0 / 128

下界 0.9 是分界线。 这一个数字同时解决了两头:eγe^{\bm{\gamma}} 不下溢、eγe^{-\bm{\gamma}} 不溢出。

2.2 通道离散度:逐通道门特有的问题

标量门下 γC\gamma^C 是一个数;逐通道门下它是 dkd_k 个数,而它们的极差决定了同一个 chunk 内不同通道的数值尺度差异:

a\bm{a} 区间 CC γC\bm{\gamma}_C 极差(log 域) e极差e^{\text{极差}}
[0.9, 0.999][0.9,\ 0.999] 64 1.28 3.583.58
[0.9, 0.999][0.9,\ 0.999] 128 1.74 5.695.69
[0.5, 0.999][0.5,\ 0.999] 64 8.17 3.53×1033.53\times10^{3}
[0.5, 0.999][0.5,\ 0.999] 128 11.25 7.69×1047.69\times10^{4}
[0.01, 0.999][0.01,\ 0.999] 128 47.46 4.10×10204.10\times10^{20}

极差 4747 意味着同一个 fragment 里最强和最弱通道的数值相差 20 个数量级–任何浮点格式都无法同时表示。这是逐通道门独有的病:标量门只需要担心整体的尺度漂移,逐通道门还要担心通道之间的尺度撕裂。

下界 0.9 把极差压到 1.74(e1.74=5.7e^{1.74} = 5.7),完全可控。

2.3 sub-chunk:第二道保险

即使有下界,C=128C = 128maxeγ=1.79×103\max e^{-\bm{\gamma}} = 1.79\times10^3(下界 0.9)–fp16 装得下但余量不多。把块内再切成 sub-chunk,每个 sub-chunk 内部重新起算 γ\bm{\gamma}

划分 maxeγ\max e^{-\bm{\gamma}}
C=128C = 128 整块 1.79×1031.79\times10^{3}
sub-chunk =64= 64 5.34×1015.34\times10^{1}
sub-chunk =32= 32 9.029.02
sub-chunk =16= 16 3.413.41

指数跨度只取决于 sub-chunk 长度,与总块长无关。这与第二篇"累积积按 chunk 重置"是同一个道理,只是又下降了一层:chunk 重置控制 γ\bm{\gamma} 本身,sub-chunk 重置控制 eγe^{-\bm{\gamma}} 的动态范围。

代价是 sub-chunk 之间需要额外的状态传递,块内变成两层循环。C=64C = 64 + 下界 0.9 时 maxeγ=55\max e^{-\bm{\gamma}} = 55不需要 sub-chunkC=128C = 128 时建议切 32 或 64。


3. 四层参考实现

沿用前三篇框架。

3.1 参考 A:逐 token 递归

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
def ref_A(Q, K, V, g, beta):
"""g: (B,N,H,DK) 逐通道 log 门控,g < 0"""
B, N, H, DK = Q.shape
DV = V.shape[-1]
O = np.zeros((B, N, H, DV))
for b in range(B):
for h in range(H):
S = np.zeros((DV, DK)) # 论文方向
for t in range(N):
k, v, q = K[b, t, h], V[b, t, h], Q[b, t, h]
a, bt = np.exp(g[b, t, h]), beta[b, t, h]
# Diag(a) 在 S 与 Householder 之间
S = S @ np.diag(a) @ (np.eye(DK) - bt * np.outer(k, k)) \
+ bt * np.outer(v, k)
O[b, t, h] = S @ q
return O

Diag(at)\operatorname{Diag}(\bm{a}_t) 的位置很关键:它夹在 St1\mathbf{S}_{t-1} 与 Householder 之间。写成 Diag(a)S\operatorname{Diag}(\bm{a})\mathbf{S} 或挪到括号外都会算错。

3.2 参考 B:chunkwise + 逐通道 UT 变换

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
def ref_B(Q, K, V, g, beta, C):
B, N, H, DK = Q.shape
DV = V.shape[-1]
O = np.zeros((B, N, H, DV))
for b in range(B):
for h in range(H):
S = np.zeros((DK, DV)) # 转置布局
for c in range(N // C):
sl = slice(c * C, (c + 1) * C)
Qc, Kc, Vc = Q[b, sl, h], K[b, sl, h], V[b, sl, h]
gc, bt = g[b, sl, h], beta[b, sl, h]

gam = np.cumsum(gc, axis=0) # γ_i ∈ R^{DK}
gC = gam[-1]

# M[c,i] = (k_c ⊙ e^{γ_c-γ_i})·k_i, i < c
M = np.zeros((C, C))
for cc in range(C):
for i in range(cc):
M[cc, i] = np.dot(Kc[cc] * np.exp(gam[cc] - gam[i]), Kc[i])

L = np.tril(bt[None, :] * M, -1) # L[c,i] = β_i M[c,i]
Tm = np.linalg.inv(np.eye(C) + L)
Ah = bt[:, None] * Tm # diag(β) T

W = Ah @ (Kc * np.exp(gam)) # C × DK
U = Ah @ Vc # C × DV
Vt = U - W @ S # 伪值 Ṽ

# A^qk[c,j] = (q_c ⊙ e^{γ_c-γ_j})·k_j, j ≤ c
Aqk = np.zeros((C, C))
for cc in range(C):
for j in range(cc + 1):
Aqk[cc, j] = np.dot(Qc[cc] * np.exp(gam[cc] - gam[j]), Kc[j])

O[b, sl, h] = (Qc * np.exp(gam)) @ S + Aqk @ Vt
S = np.diag(np.exp(gC)) @ S + (Kc * np.exp(gC - gam)).T @ Vt
return O

这份参考刻意用朴素双循环M\mathbf{M}Aqk\mathbf{A}^{qk}–慢,但与公式逐字对应,作为基准可信。§3.3 的 kernel 镜像才用 GEMM 化写法。

Lci=βiMciL_{ci} = \beta_i M_{ci}A^=diag(β)T\hat{\mathbf{A}} = \operatorname{diag}(\beta)\mathbf{T} 这一对下标必须配套。实测另一种等价写法是 Lci=βcMciL_{ci} = \beta_c M_{ci}A^=Tdiag(β)\hat{\mathbf{A}} = \mathbf{T}\operatorname{diag}(\beta),两者都对,混用则错。

3.3 参考 C:kernel 控制流镜像(GEMM 化)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
def ref_C(Q, K, V, g, beta, C, block_DV):
B, N, H, DK = Q.shape
DV = V.shape[-1]
O = np.zeros((B, N, H, DV))
for bv in range(DV // block_DV): # ← grid.x
dv = slice(bv * block_DV, (bv + 1) * block_DV)
for bbh in range(B * H): # ← grid.y
b, h = bbh // H, bbh % H
S = np.zeros((DK, block_DV))
for c in range(N // C): # ← T.Pipelined
sl = slice(c * C, (c + 1) * C)
Qc, Kc = Q[b, sl, h], K[b, sl, h]
Vc = V[b, sl, h, dv]
gc, bt = g[b, sl, h], beta[b, sl, h]

gam = np.cumsum(gc, axis=0)
gC = gam[-1]
Ep, Em = np.exp(gam), np.exp(-gam) # e^{γ}, e^{-γ}

# ── 指数可分离:两次 GEMM 代替 C²/2 次加权内积 ──
Kp, Km = Kc * Ep, Kc * Em
M = np.tril(Kp @ Km.T, -1) # 严格下三角
Aqk = np.tril((Qc * Ep) @ Km.T) # 含对角

L = np.tril(bt[None, :] * M, -1)
# 前向替换解 T @ [Kp | Vc],再逐行乘 β(因为 Â = diag(β)T)
RHS = np.concatenate([Kc * Ep, Vc], axis=1)
X = np.zeros_like(RHS)
for r in range(C):
X[r] = RHS[r]
for j in range(r):
X[r] -= L[r, j] * X[j]
X = bt[:, None] * X # ← β 在解完之后乘
W, U = X[:, :DK], X[:, DK:]

Vt = U - W @ S
O[b, sl, h, dv] = (Qc * Ep) @ S + Aqk @ Vt
S = Ep[-1][:, None] * S + (Kc * np.exp(gC - gam)).T @ Vt
return O

三处与参考 B 不同的实现选择:

  1. M\mathbf{M}Aqk\mathbf{A}^{qk} 走 GEMMK~+K~\widetilde{\mathbf{K}}^{+}\widetilde{\mathbf{K}}^{-\intercal}Aqk\mathbf{A}^{qk} 含对角(jcj \le c),M\mathbf{M} 不含(i<ci < c)。
  2. W\mathbf{W}U\mathbf{U} 拼成一次前向替换:两者共用同一个 (I+L)(\mathbf{I}+\mathbf{L}),拼接后只解一遍,省一半串行开销。
  3. β\beta 必须在解完三角系统之后乘A^=diag(β)T\hat{\mathbf{A}} = \operatorname{diag}(\beta)\mathbf{T} 展开是「先 T\mathbf{T} 作用、再逐行乘 β\beta」;若把 β\beta 提前乘进右端项,算的就是 Tdiag(β)\mathbf{T}\operatorname{diag}(\beta)–那是另一种配对(需搭配 Lci=βcMciL_{ci} = \beta_c M_{ci}),混用则错。这个错误在 β\beta 全部相等时完全看不出来,我第一次写就踩了:β\beta 随机时相对 L2 达 1.3×1011.3\times10^{-1}
  4. 状态更新用逐行乘代替 Diag\operatorname{Diag}Ep[-1][:, None] * S 就是 Diag(eγC)S\operatorname{Diag}(e^{\bm{\gamma}_C})\mathbf{S},不物化对角矩阵。

3.4 一致性验证

B=2,H=2,N=12,dk=dv=4,C=4B=2, H=2, N=12, d_k=d_v=4, C=4aU(0.90,0.999)dk\bm{a} \sim \mathcal{U}(0.90, 0.999)^{d_k}βU(0.1,0.9)\beta \sim \mathcal{U}(0.1, 0.9)q,k\bm{q},\bm{k} 已 L2 归一化,fp64:

比较 max abs 误差 相对 L2
B chunkwise vs A 逐 token 3.89×10163.89 \times 10^{-16} 2.46×10162.46 \times 10^{-16}
C kernel 镜像(GEMM 化)vs A 4.44×10164.44 \times 10^{-16} 2.57×10162.57 \times 10^{-16}
指数分离 GEMM vs 朴素加权内积 1.11×10161.11 \times 10^{-16}
L=βiML = \beta_i M + diag(β)T\operatorname{diag}(\beta)\mathbf{T} 1.11×10161.11 \times 10^{-16}
L=βcML = \beta_c M + Tdiag(β)\mathbf{T}\operatorname{diag}(\beta) 1.11×10161.11 \times 10^{-16}

三个退化检验(本篇比 GDN 多一个):

退化 应回到 检验的部分 实测
glogα1\bm{g} \equiv \log\alpha \cdot \bm{1}(所有通道同值) GDN 逐通道的指数分离 4.44×10164.44\times10^{-16}
β0\beta \to 0 逐通道纯衰减 整个 UT 与三角系统 1.21×10271.21\times10^{-27}
g0\bm{g} \to \bm{0}β\beta 保留 纯 DeltaNet 所有衰减权重 4.44×10164.44\times10^{-16}

三个 blockDV\text{block}_{DV}(1 / 2 / 4)的相对 L2 分别为 2.482.48 / 2.572.57 / 2.57×10162.57 \times 10^{-16}–切 DV 零依赖,与第一篇的结论一致。

第一个是本篇独有且最重要的:把逐通道门退化成标量门,必须精确回到第三篇的结果。这一步能抓住"指数分离时把通道维和位置维搞混"这类错误–而那类错误在通道值本来就相同时会隐身。


4. TileLang kernel

4.0 寄存器账

dk=dv=128d_k = d_v = 128C=64C = 64blockDV=32\text{block}_{DV} = 32、128 线程:

fragment GDN KDA 说明
S_f [dk,bDV][d_k,\text{bDV}] 32 32 不变
acc_o [C,bDV][C,\text{bDV}] 16 16 不变
Gam [C,C][C,C] 32 0 逐通道无法物化,取消
A [C,C][C,C] 32 32 Aqk\mathbf{A}^{qk}
Amat [C,C][C,C] 32 32 三角系统系数 L\mathbf{L}
Ep/Em [C,dk][C,d_k] 2×64 = 128 新增:e±γe^{\pm\bm{\gamma}}
W [C,dk][C,d_k] 64 新增
Delta/Vt [C,bDV][C,\text{bDV}] 16 16
合计 ≈ 160 ≈ 320 上限 255

320 超限了。 逐通道门带来的 [C,dk][C, d_k] 量级 fragment(e±γe^{\pm\bm{\gamma}}W\mathbf{W})比 C2C^2 表更吃寄存器–C×dk=64×128C \times d_k = 64\times128C2=642C^2 = 64^2 的两倍。

三条出路:

  1. eγe^{-\bm{\gamma}} 不常驻:它只在构造 M\mathbf{M}Aqk\mathbf{A}^{qk} 时用,用完即弃,可以放 shared memory。省 64 reg。
  2. W\mathbf{W} 走 shared:它是 UWS\mathbf{U} - \mathbf{W}\mathbf{S} 的中间量,不参与后续 GEMM 的累加器。省 64 reg。
  3. CC 降到 32:所有 CC 相关量减半。

组合 1+2 后约 192 reg,可行。下面的 kernel 采用这个方案。

4.1 声明与累积衰减

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
@tilelang.jit(out_idx=[5])
def kda_chunk(B, H, S, DK, DV, blk=64, block_DV=32,
num_stages=2, threads=128,
dtype=T.float16, accum_dtype=T.float32):
C = blk
NS = T.ceildiv(S, C)

@T.prim_func
def main(
Q: T.Tensor([B, S, H, DK], dtype),
K: T.Tensor([B, S, H, DK], dtype),
V: T.Tensor([B, S, H, DV], dtype),
G: T.Tensor([B, S, H, DK], accum_dtype), # 逐通道 log 门控,< 0
Beta: T.Tensor([B, S, H], accum_dtype),
O: T.Tensor([B, S, H, DV], dtype),
):
with T.Kernel(T.ceildiv(DV, block_DV), B * H, threads=threads) as (bv, bbh):
bb, bh = bbh // H, bbh % H
dv0 = bv * block_DV

Q_s = T.alloc_shared([C, DK], dtype)
K_s = T.alloc_shared([C, DK], dtype)
V_s = T.alloc_shared([C, block_DV], dtype)
S_s = T.alloc_shared([DK, block_DV], dtype)
O_s = T.alloc_shared([C, block_DV], dtype)
# 逐通道特有:e^{±γ} 与 W 放 shared,避免寄存器超限
Ep_s = T.alloc_shared([C, DK], accum_dtype) # e^{γ_r}
Em_s = T.alloc_shared([C, DK], accum_dtype) # e^{-γ_r}
Kp_s = T.alloc_shared([C, DK], dtype) # k ⊙ e^{γ}
Km_s = T.alloc_shared([C, DK], dtype) # k ⊙ e^{-γ}
W_s = T.alloc_shared([C, DK], accum_dtype)
Vt_s = T.alloc_shared([C, block_DV], dtype)

S_f = T.alloc_fragment([DK, block_DV], accum_dtype) # loop-carried
acc_o = T.alloc_fragment([C, block_DV], accum_dtype)
Aqk = T.alloc_fragment([C, C], accum_dtype)
Aq_c = T.alloc_fragment([C, C], dtype)
Lmat = T.alloc_fragment([C, C], accum_dtype) # 三角系数
Vt = T.alloc_fragment([C, block_DV], accum_dtype)
RHS_v = T.alloc_fragment([C, block_DV], accum_dtype)
gam = T.alloc_fragment([C, DK], accum_dtype) # 累积 log
be_f = T.alloc_fragment([C], accum_dtype)

累积衰减是逐通道的前缀和dkd_k 个通道各自独立 cumsum,通道之间完全并行:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
T.clear(S_f)

for i_s in T.Pipelined(NS, num_stages=num_stages):
s0 = i_s * C
T.copy(Q[bb, s0:s0+C, bh, :], Q_s)
T.copy(K[bb, s0:s0+C, bh, :], K_s)
T.copy(V[bb, s0:s0+C, bh, dv0:dv0+block_DV], V_s)

# ── 逐通道累积 log 衰减:通道维并行,位置维串行 ──
for d in T.Parallel(DK):
gam[0, d] = G[bb, s0, bh, d]
for r in T.serial(1, C):
for d in T.Parallel(DK):
gam[r, d] = gam[r - 1, d] + G[bb, s0 + r, bh, d]

for r, d in T.Parallel(C, DK):
Ep_s[r, d] = T.exp2(gam[r, d] * 1.4426950408889634) # e^{γ}
Em_s[r, d] = T.exp2(-gam[r, d] * 1.4426950408889634) # e^{-γ}
Kp_s[r, d] = T.Cast(dtype, K_s[r, d] * Ep_s[r, d])
Km_s[r, d] = T.Cast(dtype, K_s[r, d] * Em_s[r, d])
for r in T.Parallel(C):
be_f[r] = Beta[bb, s0 + r, bh]

对比 GDN:那里前缀和是 CC 步串行、每步 1 个数;这里是 CC 步串行、每步 dkd_k 个通道并行。串行长度不变,但并行度从 1 涨到 128–逐通道门在这一处反而更适合 GPU。

1.4426950408889634log2e\log_2 e,把 exe^x 转成硬件 exp2(与第三篇一致)。

4.2 指数分离:两次 GEMM 取代双重循环

1
2
3
4
5
6
7
8
9
10
11
12
# ── M = strictLower(Kp @ Km^T):衰减 KK 矩阵 ──
T.gemm(Kp_s, Km_s, Lmat, transpose_B=True, clear_accum=True)
for i, j in T.Parallel(C, C):
Lmat[i, j] = T.if_then_else(
j < i, be_f[j] * Lmat[i, j], 0.0) # L[c,i] = β_i M[c,i]

# ── A^qk = tril(Qp @ Km^T),含对角 ──
for r, d in T.Parallel(C, DK):
Q_s[r, d] = T.Cast(dtype, Q_s[r, d] * Ep_s[r, d]) # q ⊙ e^{γ}
T.gemm(Q_s, Km_s, Aqk, transpose_B=True, clear_accum=True)
for i, j in T.Parallel(C, C):
Aqk[i, j] = T.if_then_else(j <= i, Aqk[i, j], 0.0)

这是全篇最关键的两行 GEMM。 朴素写法需要 C2/2C^2/2 次长度 dkd_k 的加权内积(C=64,dk=128C=64, d_k=128 时约 26 万次 exp);分离后 e±γe^{\pm\bm{\gamma}} 各算一次(C×dk=8192C \times d_k = 8192exp2),剩下交给 Tensor Core。

注意两个掩码的边界不同:Lmatj<ij < i(严格下三角,对角会重复计入 βr\beta_r),Aqkjij \le i(含对角,自己看自己零衰减)。

4.3 前向替换:W 与 U 合并求解

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
# ── 右端项:[Kp | V],注意 β 不在这里乘 ──
for r, d in T.Parallel(C, DK):
W_s[r, d] = K_s[r, d] * Ep_s[r, d]
for r, d in T.Parallel(C, block_DV):
RHS_v[r, d] = V_s[r, d]

# ── 前向替换,一次解出 T@[Kp | V] ──
for r in T.serial(C):
for j in T.serial(r):
for d in T.Parallel(DK):
W_s[r, d] -= Lmat[r, j] * W_s[j, d]
for d in T.Parallel(block_DV):
RHS_v[r, d] -= Lmat[r, j] * RHS_v[j, d]

# ── 解完之后才逐行乘 β(Â = diag(β)T)──
for r, d in T.Parallel(C, DK):
W_s[r, d] *= be_f[r]
for r, d in T.Parallel(C, block_DV):
RHS_v[r, d] *= be_f[r]

# ── 伪值 Ṽ = U - W S ──
T.copy(S_f, S_s)
T.gemm(W_s, S_s, Vt, clear_accum=True) # W S
for r, d in T.Parallel(C, block_DV):
Vt[r, d] = RHS_v[r, d] - Vt[r, d]
T.copy(Vt, Vt_s)

W\mathbf{W}U\mathbf{U} 共用同一个 (I+L)(\mathbf{I}+\mathbf{L})拼在一次前向替换里只付一遍串行代价。这一步的并行度是 dk+blockDV=160d_k + \text{block}_{DV} = 160,比 GDN 的 32 好得多–逐通道门在这里又一次因为多了通道维而受益。

4.4 输出与状态更新

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
# ② 块间 + ③ 块内
T.gemm(Q_s, S_s, acc_o, clear_accum=True) # (Q ⊙ e^γ) S
T.copy(Aqk, Aq_c)
T.gemm(Aq_c, Vt_s, acc_o) # A^qk Ṽ

T.copy(acc_o, O_s)
T.copy(O_s, O[bb, s0:s0+C, bh, dv0:dv0+block_DV])

# ① 状态更新:Diag(e^{γ_C}) S + (K ⊙ e^{γ_C-γ_r})^T Ṽ
for d, j in T.Parallel(DK, block_DV):
S_f[d, j] *= Ep_s[C - 1, d] # 逐通道!不是标量
for r, d in T.Parallel(C, DK):
Kp_s[r, d] = T.Cast(dtype,
K_s[r, d] * Ep_s[C - 1, d] * Em_s[r, d]) # e^{γ_C-γ_r}
T.gemm(Kp_s, Vt_s, S_f, transpose_A=True)

最后一段体现了逐通道门与前三篇最直观的差别:

  • GDN:S_f[i,j] *= alpha_pow_C,一个标量乘整个状态
  • KDA:S_f[d,j] *= Ep_s[C-1,d]每个 dkd_k 行乘各自的衰减

状态的不同行按不同速率淡出–这就是"逐通道"在 kernel 层面的全部含义。

e^{\bm{\gamma}_C - \bm{\gamma}_r}Ep_s[C-1,d] * Em_s[r,d] 算,即 eγCeγre^{\gamma_C}\cdot e^{-\gamma_r}。这里复用了已有的两张表,不必重算 exp2;但它也是 eγe^{-\bm{\gamma}} 参与的第三处,进一步说明为什么下界不可省。

4.5 七处次序约束

比 GDN 多两处:

  1. 状态更新排在输出写回之后(右移语义,四篇一致)。
  2. S_f *= Ep_s[C-1,:]T.gemm 之前
  3. Q_s 被原地乘 eγe^{\bm{\gamma}} 后不可再用于块内项–但本篇的块内项恰好也用 qeγ\bm{q}\odot e^{\bm{\gamma}}Aqk\mathbf{A}^{qk} 的定义里就带),所以不需要重载 Q。这是与 GDN 的一处反差,容易照抄出错。
  4. Lmat 只取严格下三角j<ij < i),Aqk 含对角(jij \le i)。
  5. K_s 必须保持原始值Kp\mathbf{Kp}Km\mathbf{Km}、右端项、状态更新四处都从 K_s 派生,任何一处原地修改都会污染后续。本篇的做法是始终写入独立的 Kp_s/Km_sK_s 只读–比 GDN 的"污染两次再重载"更清晰。
  6. β\beta 在前向替换之后乘,不能提前混进右端项A^=diag(β)T\hat{\mathbf{A}} = \operatorname{diag}(\beta)\mathbf{T}Tdiag(β)\mathbf{T}\operatorname{diag}(\beta) 是两种不同配对,各自要搭配 Lci=βiMciL_{ci} = \beta_i M_{ci}βcMci\beta_c M_{ci}β\beta 全相等时这个错误完全隐身。
  7. T.clear(S_f) 在循环外,clear_accum=True 在循环内

第 3、5、6 条都是"照抄上一篇会错"的地方,其中第 6 条我实际踩过。


5. 四篇对照

线性注意力 标量衰减 GDN KDA
衰减 α\alpha 常数 αt\alpha_t 标量 atRdk\bm{a}_t \in \mathbb{R}^{d_k}
删除 Iβkk\mathbf{I}-\beta\bm{k}\bm{k}^\intercal 同 GDN
衰减掩码 M\mathbf{M}(0/1) Γ\GammaC2C^2 表) Γ\GammaC2C^2 表) 无表,融进 GEMM
累积积 编译期常量 CC 步串行,宽度 1 CC 步串行,宽度 dkd_k
前向替换并行度 bDV=32\text{bDV} = 32 dk+bDV=160d_k + \text{bDV} = 160
主要 fragment C2C^2 2C22C^2 3C23C^2 2C2+3Cdk2C^2 + 3Cd_k
可用 CC 128 128 64 64(需 shared 卸载)
数值命门 1/γ1/\gamma 溢出 三角系统条件数 e±γe^{\pm\bm{\gamma}} 双向 + 通道离散
关键前提 累积积按 chunk 重置 k\bm{k} L2 归一化 门控下界 ≥ 0.9

四篇的箭头量语义完全一致q\overleftarrow{\bm{q}}eγre^{\bm{\gamma}_r}k\overrightarrow{\bm{k}}eγCγre^{\bm{\gamma}_C-\bm{\gamma}_r}S\overrightarrow{\mathbf{S}}eγCe^{\bm{\gamma}_C}。从第一篇的"无衰减"到本篇的"逐通道",衰减插入的三个位置从未改变,改变的只是每个位置乘的是标量还是向量。


6. 总结

  1. 逐通道门只改一个符号,却拆掉了前三篇所有便利αtDiag(at)\alpha_t \to \operatorname{Diag}(\bm{a}_t) 让累积积从标量变向量,Γ\GammaC×CC\times C 表变成概念上的 C×C×dkC\times C\times d_k不可能物化,必须融进 GEMM。
  2. 衰减长在 KK 矩阵内部,这是 KDA chunkwise 的核心难点Mci=(kceγcγi)kiM_{ci} = (\bm{k}_c \odot e^{\bm{\gamma}_c-\bm{\gamma}_i})^\intercal\bm{k}_i 的指数依赖三个下标,无法像标量门那样外提成逐元素乘。
  3. 出路是指数可分离eγcγi=eγceγie^{\bm{\gamma}_c-\bm{\gamma}_i} = e^{\bm{\gamma}_c}\odot e^{-\bm{\gamma}_i},于是 M=K~+K~\mathbf{M} = \widetilde{\mathbf{K}}^{+}\widetilde{\mathbf{K}}^{-\intercal} 一次 GEMM 解决(实测等价,误差 1.94×10161.94\times10^{-16})。朴素写法要 C2/2C^2/2 次加权内积、约 26 万次 exp;分离后只需 Cdk=8192C d_k = 8192exp2 加两次 GEMM。
  4. 代价是必须物化 eγe^{-\bm{\gamma}}–第二篇批判过的 1/γ1/\gamma 陷阱的逐通道版本。但这次没有替代方案:不拆就用不上 Tensor Core。问题从"要不要拆"变成"如何让拆分数值安全"。
  5. 门控下界是 kernel 可行性的前提,不是精度调优。实测 C=128C=128、无下界时 fp16 下 eγCe^{\bm{\gamma}_C} 128/128 通道全部归零,状态被彻底清空。下界 0.9 时归零通道数为 0,同时 maxeγ=55\max e^{-\bm{\gamma}} = 55 不溢出–一个数字同时管住了下溢与溢出两头
  6. 通道离散度是逐通道门独有的病a[0.01,0.999]\bm{a}\in[0.01,0.999]C=128C=128γC\bm{\gamma}_C 的通道极差达 47(log 域),即同一 fragment 内最强与最弱通道相差 20 个数量级,任何浮点格式都无法同时表示。下界 0.9 把极差压到 1.74。
  7. sub-chunk 是第二道保险eγe^{-\bm{\gamma}} 的动态范围只取决于 sub-chunk 长度:C=128C=128 整块 1.79×1031.79\times10^3,切 32 后降到 9.02。这与第二篇"累积积按 chunk 重置"同理,只是又降一层。C=64C=64 + 下界 0.9 时不需要。
  8. 逐通道门在两处反而更适合 GPU。累积前缀和:GDN 是 CC 步串行 × 宽度 1,KDA 是 CC 步串行 × 宽度 dkd_k,串行长度不变而并行度从 1 涨到 128。前向替换:W\mathbf{W}U\mathbf{U} 共用 (I+L)(\mathbf{I}+\mathbf{L}) 拼成一次求解,并行度 dk+bDV=160d_k + \text{bDV} = 160,远高于 GDN 的 32。
  9. 寄存器压力换了主角。GDN 的瓶颈是三张 C2C^2 表;KDA 是三张 C×dkC \times d_ke±γe^{\pm\bm{\gamma}}W\mathbf{W}),64×12864\times12864264^2 的两倍。朴素分配约 320 reg 超限,把 eγe^{-\bm{\gamma}}W\mathbf{W} 卸载到 shared memory 后降到约 192。
  10. 三处"照抄上一篇会错":块内项的 Aqk\mathbf{A}^{qk} 定义里本就带 qeγ\bm{q}\odot e^{\bm{\gamma}}不需要像 GDN 那样重载 QK_s 应始终只读、派生量写独立 buffer,而非 GDN 的"污染再重载";β\beta 必须在解完三角系统之后乘diag(β)T\operatorname{diag}(\beta)\mathbf{T}Tdiag(β)\mathbf{T}\operatorname{diag}(\beta) 需搭配不同的 LL 下标,写反时 β\beta 全相等则隐身、β\beta 随机则相对 L2 达 1.3×1011.3\times10^{-1}(我第一次写就踩了这个)。
  11. 三个退化检验,比 GDN 多一个g\bm{g} 各通道同值 → 回到 GDN(检验指数分离是否搞混通道维与位置维)、β0\beta\to0 → 逐通道纯衰减、g0\bm{g}\to\bm{0} → 纯 DeltaNet。第一个最重要且本篇独有。
  12. 衰减的三处落点四篇未变q\overleftarrow{\bm{q}}k\overrightarrow{\bm{k}}S\overrightarrow{\mathbf{S}} 从第一篇到第四篇完全一致,变的只是每处乘标量还是乘向量。这是 GDN 那套箭头记号真正的价值–它把"衰减"隔离成了一个独立于门控形态的层面。

可迁移的启示:把一个标量参数升级成向量,代价从来不在参数量。它会连带改变哪些量能被预计算、哪些表能被物化、哪些运算能进 Tensor Core。KDA 的 αtDiag(at)\alpha_t \to \operatorname{Diag}(\bm{a}_t) 只多了 dkd_k 个数,却让衰减掩码从"一张可复用的表"变成"必须融进 GEMM 的隐式结构",并且逼出了一个模型层面的约束(门控下界)作为 kernel 可行性的前提。当一个数值技巧成为架构设计的必要条件时,它就不再是实现细节。

参考

TileLang 实战:KDA 从零到一–Gated DeltaNet

前两篇分别实现了无遗忘的 chunked 线性注意力,以及带标量衰减的版本。两者的状态更新都只做加法:新的键值对累加进状态,旧信息靠衰减系数被动遗忘。本篇补上最后一块–删除

递推式变成:

St=St1(αt(Iβtktkt))+βtvtkt\mathbf{S}_t = \mathbf{S}_{t-1}\Big(\alpha_t\big(\mathbf{I} - \beta_t \bm{k}_t\bm{k}_t^{\intercal}\big)\Big) + \beta_t\bm{v}_t\bm{k}_t^{\intercal}

比上一篇多出的是 Iβtktkt\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal} 这个广义 Householder 变换。它带来的不是又一个逐元素权重,而是一个矩阵乘在状态右侧–这一改把整个 chunkwise 算法的结构改掉了:块内不再能靠一张下三角权重表解决,需要解一个 C×CC \times C 的三角系统。

上两篇见《Chunked 线性注意力》与《标量衰减》,本文沿用其符号约定与四层参考验证框架。


0. 符号约定

沿用 GDN(arXiv:2412.06464v3)§2.2、§3.1、§3.3 与附录 A 的记号,本篇新增两个:

符号 含义 说明
βt(0,1)\beta_t \in (0,1) 写入强度(writing strength) 也是 delta rule 视角下的学习率
T[t]RC×C\mathbf{T}_{[t]} \in \mathbb{R}^{C \times C} UT 变换矩阵 下三角系统的逆,本篇的核心开销
U~[t]RC×dv\widetilde{\mathbf{U}}_{[t]} \in \mathbb{R}^{C \times d_v} 修正后的 value T\mathbf{T} 作用在 diag(β)V\operatorname{diag}(\beta)\mathbf{V}
W[t]RC×dk\mathbf{W}_{[t]} \in \mathbb{R}^{C \times d_k} 修正后的 key T\mathbf{T} 作用在 diag(β)K\operatorname{diag}(\beta)\mathbf{K}

沿用的:αt\alpha_t 单步衰减、γ[t]r=i=1rαi\gamma^r_{[t]} = \prod_{i=1}^{r}\alpha_i 累积衰减积(按 chunk 重置)、Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 衰减感知因果掩码、CC 块长、[t][t] 块序号、r[1,C]r \in [1,C] 块内位置。

一个前提:论文对 q,k\bm{q}, \bm{k} 做 L2 归一化,即 kt=1\|\bm{k}_t\| = 1。这不只是训练稳定性的考虑–下面 §1.2 会看到它直接决定了 Householder 变换会不会把状态越推越大。


1. 从加法到删除:delta rule 在做什么

1.1 三种更新方式的对照

把三篇的递推式并排放,差异一目了然:

形态 递推式 状态如何变化
线性注意力 St=St1+vtkt\mathbf{S}_t = \mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} 只增不减
标量衰减 St=αtSt1+vtkt\mathbf{S}_t = \alpha_t\mathbf{S}_{t-1} + \bm{v}_t\bm{k}_t^{\intercal} 整体等比遗忘
gated delta rule St=St1(αt(Iβtktkt))+βtvtkt\mathbf{S}_t = \mathbf{S}_{t-1}\big(\alpha_t(\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal})\big) + \beta_t\bm{v}_t\bm{k}_t^{\intercal} 定向替换 + 整体遗忘

标量衰减的问题在于它不分对象αt\alpha_t 一乘,所有历史信息按同一比例衰减。要腾出空间写入新内容,只能把无关的旧信息一起冲淡。GDN 论文的说法是,门控擅长「快速擦除」,delta rule 擅长「定向修改」,两者互补。

delta rule 的定向性来自哪里?把它拆开看:

St=St1(St1kt)vtoldkt+(βtvt+(1βt)St1kt)vtnewkt\mathbf{S}_t = \mathbf{S}_{t-1} - \underbrace{(\mathbf{S}_{t-1}\bm{k}_t)}_{\bm{v}_t^{\text{old}}}\bm{k}_t^{\intercal} + \underbrace{\big(\beta_t\bm{v}_t + (1-\beta_t)\mathbf{S}_{t-1}\bm{k}_t\big)}_{\bm{v}_t^{\text{new}}}\bm{k}_t^{\intercal}

读法是:先把 kt\bm{k}_t 这个键上原有的值 vtold=St1kt\bm{v}^{\text{old}}_t = \mathbf{S}_{t-1}\bm{k}_t 减掉,再写入新值 vtnew\bm{v}^{\text{new}}_t,而新值是旧值与目标值的凸组合,βt\beta_t 控制替换的彻底程度。βt1\beta_t \to 1 是完全覆盖,βt0\beta_t \to 0 是不动。

关键在于这个减法只作用在 kt\bm{k}_t 方向上,与 kt\bm{k}_t 正交的记忆分毫不动。这就是「定向」–相比标量衰减的一刀切,delta rule 只擦掉要覆盖的那一条。

1.2 为什么是 Householder,以及 L2 归一化的作用

Iβtktkt\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal} 是广义 Householder 变换。kt=1\|\bm{k}_t\| = 1 时,它的特征值只有两种取值,结构一目了然:

  • 沿 kt\bm{k}_t 方向:特征值 1βt1 - \beta_t
  • kt\bm{k}_t 正交的 dk1d_k - 1 个方向:特征值 11

于是 βt(0,1)\beta_t \in (0,1) 时全部特征值的绝对值都不超过 1,状态每步只会被压缩、不会被放大,递推因此稳定。这正是论文对 k\bm{k} 做 L2 归一化的深层原因–若 kt1\|\bm{k}_t\| \ne 1,沿 kt\bm{k}_t 的特征值变成 1βtkt21 - \beta_t\|\bm{k}_t\|^2βtkt2>2\beta_t\|\bm{k}_t\|^2 > 2 时就会翻到 1-1 以下,递推放大。

顺带一提,论文脚注提到可以放开到 βt(0,2)\beta_t \in (0,2) 以允许负特征值,那是为了解锁状态跟踪能力(state tracking)。本文按 (0,1)(0,1) 处理。

1.3 test-time SGD 视角

论文给了一个很有启发的解释:把状态 S\mathbf{S} 看成一个快速权重矩阵,delta rule 就是在做在线回归的一步梯度下降。目标是 L(St)=12Stktvt2\mathcal{L}(\mathbf{S}_t) = \frac{1}{2}\|\mathbf{S}_t\bm{k}_t - \bm{v}_t\|^2,那么:

StβtL(St)=Stβt(Stktvt)kt=St(Iβtktkt)+βtvtkt\mathbf{S}_t - \beta_t\nabla\mathcal{L}(\mathbf{S}_t) = \mathbf{S}_t - \beta_t(\mathbf{S}_t\bm{k}_t - \bm{v}_t)\bm{k}_t^{\intercal} = \mathbf{S}_t(\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal}) + \beta_t\bm{v}_t\bm{k}_t^{\intercal}

βt\beta_t 就是学习率,αt\alpha_t 就是 weight decay。 这个视角下 gated delta rule 没有任何神秘之处–它是带权重衰减的 test-time SGD。


2. Chunkwise 形式:为什么需要解三角系统

2.1 展开递推:转移矩阵不再是标量

按块展开 rr 步(论文式 10):

S[t]r=S[t]i=1rα[t]i(Iβ[t]ik[t]ik[t]i)F[t]r+i=1rβ[t]iv[t]ik[t]ij=i+1rα[t]j(Iβ[t]jk[t]jk[t]j)G[t]r\mathbf{S}_{[t]}^{r} = \mathbf{S}_{[t]}\underbrace{\prod_{i=1}^{r}\alpha^i_{[t]}\big(\mathbf{I} - \beta^i_{[t]}\bm{k}^i_{[t]}\bm{k}^{i\intercal}_{[t]}\big)}_{\mathbf{F}^r_{[t]}} + \underbrace{\sum_{i=1}^{r}\beta^i_{[t]}\bm{v}^i_{[t]}\bm{k}^{i\intercal}_{[t]}\prod_{j=i+1}^{r}\alpha^j_{[t]}\big(\mathbf{I} - \beta^j_{[t]}\bm{k}^j_{[t]}\bm{k}^{j\intercal}_{[t]}\big)}_{\mathbf{G}^r_{[t]}}

对比上一篇:那里的转移量是标量 αri\alpha^{r-i},可以直接查表。这里是矩阵连乘F\mathbf{F}G\mathbf{G} 都不能靠逐元素权重表达。

衰减部分可以先提出来–αi=γr\prod\alpha_i = \gamma^r 是标量,与 Householder 部分可交换:

F[t]r=γ[t]rP[t]r,P[t]r=i=1r(Iβikiki)\mathbf{F}^r_{[t]} = \gamma^r_{[t]}\,\mathbf{P}^r_{[t]}, \qquad \mathbf{P}^r_{[t]} = \prod_{i=1}^{r}\big(\mathbf{I} - \beta^i\bm{k}^i\bm{k}^{i\intercal}\big)

剩下的 P\mathbf{P} 是纯 DeltaNet 的部分,靠 WY 表示处理–这正是论文 §2.2 的内容,下一节完整展开。

2.2 论文 §2.2:无门控 DeltaNet 的 WY 表示

在处理带门控的版本之前,先把论文 §2.2 那套无门控 DeltaNet 的 chunkwise 推导完整走一遍。GDN 的做法本质上是在这套框架上打补丁,先看清基线,后面的改动才有参照。

本节 αt1\alpha_t \equiv 1(无遗忘门),递推退化为 St=St1(Iβtktkt)+βtvtkt\mathbf{S}_t = \mathbf{S}_{t-1}(\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^{\intercal}) + \beta_t\bm{v}_t\bm{k}_t^{\intercal}

2.2.1 部分展开:两个连乘(论文式 3)

按块部分展开递推:

S[t]r=S[t](i=1r(Iβ[t]ik[t]ik[t]i)):=P[t]r+i=1rβ[t]iv[t]ik[t]ij=i+1r(Iβ[t]jk[t]jk[t]j):=H[t]r\mathbf{S}^r_{[t]} = \mathbf{S}_{[t]}\underbrace{\Big(\prod_{i=1}^{r}\big(\mathbf{I} - \beta^i_{[t]}\bm{k}^i_{[t]}\bm{k}^{i\intercal}_{[t]}\big)\Big)}_{\textstyle :=\mathbf{P}^r_{[t]}} + \underbrace{\sum_{i=1}^{r}\beta^i_{[t]}\bm{v}^i_{[t]}\bm{k}^{i\intercal}_{[t]}\prod_{j=i+1}^{r}\big(\mathbf{I} - \beta^j_{[t]}\bm{k}^j_{[t]}\bm{k}^{j\intercal}_{[t]}\big)}_{\textstyle :=\mathbf{H}^r_{[t]}}

两个部分的角色不同,值得分清:

形状 含义 结构
P[t]r\mathbf{P}^r_{[t]} dk×dkd_k \times d_k 历史状态的遗忘算子:入口状态经过 rr 步 Householder 后剩下什么 Householder 的纯连乘
H[t]r\mathbf{H}^r_{[t]} dv×dkd_v \times d_k 块内新写入的累积:前 rr 个 token 写进来的内容(互相已扣除重叠) 连乘的加权和

P\mathbf{P} 是纯连乘,H\mathbf{H} 的每一项后面还挂着一截连乘尾巴–两者都不能直接算,C=64C = 64P\mathbf{P} 要 64 个 dk×dkd_k\times d_k 矩阵相乘。WY 表示的作用就是把这两个连乘各自压成一次求和。

2.2.2 经典 WY:把 Householder 连乘压成秩-CC 更新(论文式 4)

这是 Bischof & Van Loan (1985) 的经典结果。核心事实:CC 个 Householder 矩阵的乘积可以写成单位矩阵减去一个秩至多 CC 的修正

P[t]r=Ii=1rw[t]ik[t]iRdk×dk,w[t]r=β[t]r(k[t]ri=1r1w[t]i(k[t]ik[t]r))Rdk\mathbf{P}^r_{[t]} = \mathbf{I} - \sum_{i=1}^{r}\bm{w}^i_{[t]}\bm{k}^{i\intercal}_{[t]} \in \mathbb{R}^{d_k\times d_k}, \qquad \bm{w}^r_{[t]} = \beta^r_{[t]}\Big(\bm{k}^r_{[t]} - \sum_{i=1}^{r-1}\bm{w}^i_{[t]}\big(\bm{k}^{i\intercal}_{[t]}\bm{k}^r_{[t]}\big)\Big) \in \mathbb{R}^{d_k}

为什么成立,看一步归纳就够:

Pr=Pr1(Iβrkrkr)=Pr1βrPr1krwrkr\mathbf{P}^{r} = \mathbf{P}^{r-1}\big(\mathbf{I} - \beta_r\bm{k}_r\bm{k}_r^{\intercal}\big) = \mathbf{P}^{r-1} - \underbrace{\beta_r\mathbf{P}^{r-1}\bm{k}_r}_{\textstyle \bm{w}_r}\bm{k}_r^{\intercal}

于是 wr=βrPr1kr\bm{w}_r = \beta_r\mathbf{P}^{r-1}\bm{k}_r,把 Pr1=Ii<rwiki\mathbf{P}^{r-1} = \mathbf{I} - \sum_{i<r}\bm{w}_i\bm{k}_i^{\intercal} 代入就得到上面那个递推。每多一个 Householder,秩只增加 1–这就是"连乘变求和"的全部内容。

2.2.3 同一个模具:H\mathbf{H} 的 WY 表示(论文式 5)

H\mathbf{H} 的推导结构与 P\mathbf{P} 完全平行

H[t]r=i=1ru[t]ik[t]iRdv×dk,u[t]r=β[t]r(v[t]ri=1r1u[t]i(k[t]ik[t]r))Rdv\mathbf{H}^r_{[t]} = \sum_{i=1}^{r}\bm{u}^i_{[t]}\bm{k}^{i\intercal}_{[t]} \in \mathbb{R}^{d_v\times d_k}, \qquad \bm{u}^r_{[t]} = \beta^r_{[t]}\Big(\bm{v}^r_{[t]} - \sum_{i=1}^{r-1}\bm{u}^i_{[t]}\big(\bm{k}^{i\intercal}_{[t]}\bm{k}^r_{[t]}\big)\Big) \in \mathbb{R}^{d_v}

w\bm{w} 的递推和 u\bm{u} 的递推并排看,会发现它们是同一个式子

wr=βr(kri<rwi(kikr)),ur=βr(vri<rui(kikr))\bm{w}_r = \beta_r\Big(\underline{\bm{k}_r} - \sum_{i<r}\bm{w}_i(\bm{k}_i^{\intercal}\bm{k}_r)\Big), \qquad \bm{u}_r = \beta_r\Big(\underline{\bm{v}_r} - \sum_{i<r}\bm{u}_i(\bm{k}_i^{\intercal}\bm{k}_r)\Big)

只有下划线处不同:一个是 kr\bm{k}_r、一个是 vr\bm{v}_r系数完全一样–这个观察是后面式 (6)(7) 能共用一个 T\mathbf{T} 的全部原因,也是 kernel 里 W\mathbf{W}U\mathbf{U} 能拼进一次前向替换的依据(第四篇 KDA 用的就是这个技巧)。

写成矩阵形式,两个连乘都消失了:

P[t]=IW[t]K[t]Rdk×dk,H[t]=U[t]K[t]Rdv×dk\mathbf{P}_{[t]} = \mathbf{I} - \mathbf{W}_{[t]}^{\intercal}\mathbf{K}_{[t]} \in \mathbb{R}^{d_k\times d_k}, \qquad \mathbf{H}_{[t]} = \mathbf{U}_{[t]}^{\intercal}\mathbf{K}_{[t]} \in \mathbb{R}^{d_v\times d_k}

2.2.4 UT 变换:递推也不必串行(论文式 6、7)

式 (4)(5) 虽然把连乘变成了求和,但 wr\bm{w}_r / ur\bm{u}_r 自身还是串行递推。Joffrain et al. (2006) 的 UT 变换把它变成一次矩阵求逆:

T[t]=[I+strictLower(diag(β[t])K[t]K[t])]1diag(β[t])RC×C\mathbf{T}_{[t]} = \Big[\mathbf{I} + \operatorname{strictLower}\big(\operatorname{diag}(\beta_{[t]})\mathbf{K}_{[t]}\mathbf{K}_{[t]}^{\intercal}\big)\Big]^{-1}\operatorname{diag}(\beta_{[t]}) \in \mathbb{R}^{C\times C}

W[t]=T[t]K[t]RC×dk,U[t]=T[t]V[t]RC×dv\mathbf{W}_{[t]} = \mathbf{T}_{[t]}\mathbf{K}_{[t]} \in \mathbb{R}^{C\times d_k}, \qquad \mathbf{U}_{[t]} = \mathbf{T}_{[t]}\mathbf{V}_{[t]} \in \mathbb{R}^{C\times d_v}

W\mathbf{W}U\mathbf{U} 共用同一个 T\mathbf{T},正是因为 §2.2.3 那两个递推的系数相同。这一步的意义在于:T\mathbf{T} 只跟 K\mathbf{K}β\beta 有关,与 V\mathbf{V} 无关。

2.2.5 代回式 3:可用 Tensor Core 的算法(论文式 8、9)

S[t+1]=S[t]P[t]+H[t]=S[t]+(U[t]W[t]S[t])K[t]Rdv×dkO[t]=Q[t]S[t]+(Q[t]K[t]M)(U[t]W[t]S[t])RC×dv\begin{aligned} \mathbf{S}_{[t+1]} &= \mathbf{S}_{[t]}\mathbf{P}_{[t]} + \mathbf{H}_{[t]} = \mathbf{S}_{[t]} + \big(\mathbf{U}_{[t]} - \mathbf{W}_{[t]}\mathbf{S}_{[t]}^{\intercal}\big)^{\intercal}\mathbf{K}_{[t]} \in \mathbb{R}^{d_v\times d_k} \\[4pt] \mathbf{O}_{[t]} &= \mathbf{Q}_{[t]}\mathbf{S}_{[t]}^{\intercal} + \big(\mathbf{Q}_{[t]}\mathbf{K}_{[t]}^{\intercal}\odot\mathbf{M}\big)\big(\mathbf{U}_{[t]} - \mathbf{W}_{[t]}\mathbf{S}_{[t]}^{\intercal}\big) \in \mathbb{R}^{C\times d_v} \end{aligned}

式 (8) 的化简值得看一眼–SP=S(IWK)=S(WS)K\mathbf{S}\mathbf{P} = \mathbf{S}(\mathbf{I}-\mathbf{W}^\intercal\mathbf{K}) = \mathbf{S} - (\mathbf{W}\mathbf{S}^\intercal)^\intercal\mathbf{K},与 H=UK\mathbf{H} = \mathbf{U}^\intercal\mathbf{K} 合并后 K\mathbf{K} 被提到右侧,括号里剩下 UWS\mathbf{U} - \mathbf{W}\mathbf{S}^\intercal这个量在式 (8) 和式 (9) 里是同一个,算一次用两回。

M\mathbf{M} 是下三角全 1 的因果掩码。注意此处没有任何衰减–这正是 GDN 要改的地方。

2.2.6 数值验证:九个等式逐一核对

dk=6,dv=5,C=7d_k=6, d_v=5, C=7βU(0.1,0.9)\beta \sim \mathcal{U}(0.1,0.9)k\bm{k} 已 L2 归一化,入口状态 S00\mathbf{S}_0 \ne \mathbf{0}(比论文附录 A 的首块假设更严格),fp64:

论文式 内容 max abs 误差
(3) Sr=S0Pr+Hr\mathbf{S}^r = \mathbf{S}_0\mathbf{P}^r + \mathbf{H}^r 4.44×10164.44\times10^{-16}
(4) Pr=Iiwiki\mathbf{P}^r = \mathbf{I} - \sum_i\bm{w}_i\bm{k}_i^\intercal 4.44×10164.44\times10^{-16}
(5) Hr=iuiki\mathbf{H}^r = \sum_i\bm{u}_i\bm{k}_i^\intercal 3.33×10163.33\times10^{-16}
矩阵形式 P=IWK\mathbf{P} = \mathbf{I} - \mathbf{W}^\intercal\mathbf{K} 2.22×10162.22\times10^{-16}
矩阵形式 H=UK\mathbf{H} = \mathbf{U}^\intercal\mathbf{K} 2.78×10162.78\times10^{-16}
(6)(7) W=TK\mathbf{W} = \mathbf{T}\mathbf{K}(UT 变换 vs 式 4 递推) 5.55×10175.55\times10^{-17}
(6)(7) U=TV\mathbf{U} = \mathbf{T}\mathbf{V} 2.22×10162.22\times10^{-16}
(8) S[t+1]\mathbf{S}_{[t+1]} 完整式 4.44×10164.44\times10^{-16}
(9) O[t]\mathbf{O}_{[t]} 完整式 8.88×10168.88\times10^{-16}

另外验证了 P=IWK\mathbf{P} = \mathbf{I}-\mathbf{W}^\intercal\mathbf{K} 的特征值绝对值分别为 0.9918,0.9082,0.5805,0.3131,0.1951,0.06760.9918, 0.9082, 0.5805, 0.3131, 0.1951, 0.0676全部 1\le 1,与 §1.2 说的"每步只压缩不放大"一致。

2.2.7 GDN 改了哪两处

把门控加回来(αt1\alpha_t \ne 1),论文 §3.3 的做法只动了式 (6)(7) 两个地方

无门控(式 6、7) GDN(门控版)
T\mathbf{T} 里的 KK 矩阵 KK\mathbf{K}\mathbf{K}^{\intercal} ΓKK\Gamma \odot \mathbf{K}\mathbf{K}^{\intercal}
W\mathbf{W} 的输入 K\mathbf{K} diag(γ)K=K\operatorname{diag}(\gamma)\mathbf{K} = \overleftarrow{\mathbf{K}}
U\mathbf{U} 的输入 V\mathbf{V} V\mathbf{V}(不变)
式 (8) 的旧状态项 S[t]\mathbf{S}_{[t]} γCS[t]=S\gamma^C\mathbf{S}_{[t]} = \overrightarrow{\mathbf{S}}
式 (8) 的 K\mathbf{K} K\mathbf{K} γCγrK=K\frac{\gamma^C}{\gamma^r}\mathbf{K} = \overrightarrow{\mathbf{K}}
式 (9) 的 Q\mathbf{Q} 与掩码 Q\mathbf{Q}M\mathbf{M} γrQ=Q\gamma^r\mathbf{Q} = \overleftarrow{\mathbf{Q}}Γ\Gamma

实测这个替换的正确性:门控版 S[t+1]\mathbf{S}_{[t+1]} 与逐 token 递归差 3.33×10163.33\times10^{-16};令 α1\alpha\equiv1 时,Γ\Gamma 与下三角全 1 掩码 M\mathbf{M} 完全相等(差 0.0)T\mathbf{T} 与无门控版完全相等(差 0.0)W\mathbf{W} 也完全相等。三个 0.0 说明门控版是无门控版的严格推广,没有引入任何额外近似。

换个角度看Γ\Gamma 就是把式 (9) 那个 0/1 因果掩码 M\mathbf{M} 升级成了"带衰减的因果掩码"。Mij{0,1}\mathbf{M}_{ij} \in \{0,1\} 只回答"jj 能否影响 ii",Γij=γi/γj\Gamma_{ij} = \gamma_i/\gamma_j 还回答"影响衰减了多少"。前两篇反复出现的那个 Γ\Gamma,在这个视角下就是 M\mathbf{M} 的自然推广。


2.3 手推一遍:Householder 连乘怎么变成三角系统

直接推 C=4C = 4 的情形最清楚。定义每步的修正量 dr\bm{d}_r,使得递推写成加法形式:

Sr=αrSr1+drkr,dr=βrvrαrβrSr1kr\mathbf{S}^r = \alpha_r\mathbf{S}^{r-1} + \bm{d}_r\bm{k}_r^{\intercal}, \qquad \bm{d}_r = \beta_r\bm{v}_r - \alpha_r\beta_r\mathbf{S}^{r-1}\bm{k}_r

这一步只是把 Sr1(αr(Iβrkk))+βrvk\mathbf{S}^{r-1}(\alpha_r(\mathbf{I}-\beta_r\bm{k}\bm{k}^\intercal)) + \beta_r\bm{v}\bm{k}^\intercal 重新分组,恒等变形。注意 dr\bm{d}_r 里带 αr\alpha_r–删除项作用在已经衰减过的状态上,这个细节写错不会报错、只会算错。

有了加法形式,就能像上一篇那样倒代换(每项系数是「距块尾的步数」):

S[t+1]=γCS[t]+r=1CγCγrdrkr\mathbf{S}_{[t+1]} = \gamma^C\mathbf{S}_{[t]} + \sum_{r=1}^{C}\frac{\gamma^C}{\gamma^r}\,\bm{d}_r\bm{k}_r^{\intercal}

问题在于 dr\bm{d}_r 依赖 Sr1\mathbf{S}^{r-1},而 Sr1\mathbf{S}^{r-1} 又依赖 d1,,dr1\bm{d}_1, \ldots, \bm{d}_{r-1}–串行依赖,无法并行。把 Sr1\mathbf{S}^{r-1} 也倒代换开(利用 αrγr1=γr\alpha_r\gamma^{r-1} = \gamma^r):

αrSr1=γrS[t]+i<rγrγidiki\alpha_r\mathbf{S}^{r-1} = \gamma^r\mathbf{S}_{[t]} + \sum_{i<r}\frac{\gamma^r}{\gamma^i}\bm{d}_i\bm{k}_i^{\intercal}

代回 dr\bm{d}_r 的定义:

dr=βr(vrγrS[t]kri<rγrγi(kikr)di)\bm{d}_r = \beta_r\Big(\bm{v}_r - \gamma^r\mathbf{S}_{[t]}\bm{k}_r - \sum_{i<r}\frac{\gamma^r}{\gamma^i}\big(\bm{k}_i^{\intercal}\bm{k}_r\big)\bm{d}_i\Big)

这是一个下三角线性系统dr\bm{d}_r 只依赖 di (i<r)\bm{d}_i\ (i < r),系数是 βrγrγi(kikr)\beta_r\frac{\gamma^r}{\gamma^i}(\bm{k}_i^\intercal\bm{k}_r)。写成矩阵形式,令 Ari=βrγrγi(kikr)\mathbf{A}_{ri} = \beta_r\frac{\gamma^r}{\gamma^i}(\bm{k}_i^\intercal\bm{k}_r) 的严格下三角部分:

(I+strictLower(A))Δ=diag(β)(Vdiag(γ)KS[t])(\mathbf{I} + \operatorname{strictLower}(\mathbf{A}))\,\Delta = \operatorname{diag}(\beta)\big(\mathbf{V} - \operatorname{diag}(\gamma)\mathbf{K}\mathbf{S}_{[t]}^{\intercal}\big)

其中 ΔRC×dv\Delta \in \mathbb{R}^{C \times d_v} 的第 rr 行是 dr\bm{d}_r。注意 A=diag(β)(ΓKK)\mathbf{A} = \operatorname{diag}(\beta)\big(\Gamma \odot \mathbf{K}\mathbf{K}^{\intercal}\big)衰减感知掩码 Γ\Gamma 在这里第二次出现,这次是嵌在三角系统的系数矩阵里。

2.4 论文附录 A:扩展 WY 表示的归纳证明

上面 §2.3 是"把递推硬拆开"的推法。论文附录 A 给了一个更漂亮的等价路线:先猜出闭式,再用数学归纳法证明。这一节按论文原文复述(GDN 论文附录 A,为减少符号负担,论文同样只考虑首块,即 S0=0\mathbf{S}_0 = \mathbf{0})。

命题(扩展 WY 表示).St\mathbf{S}_t

St=i=1tγtγiuiki,ut=βt(vti=1t1γtγiuikikt)\mathbf{S}_t = \sum_{i=1}^{t}\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}, \qquad \bm{u}_t = \beta_t\Big(\bm{v}_t - \sum_{i=1}^{t-1}\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}\bm{k}_t\Big)

证明.tt 作归纳。

St+1=St(αt+1(Iβt+1kt+1kt+1))+βt+1vt+1kt+1=αt+1(i=1tγtγiuiki)αt+1βt+1(i=1tγtγiuikikt+1kt+1)+βt+1vt+1kt+1=i=1tγt+1γiuiki+βt+1(vt+1i=1tγt+1γiuikikt+1)ut+1kt+1=i=1tγt+1γiuiki+γt+1γt+11ut+1kt+1  =  i=1t+1γt+1γiuiki\begin{aligned} \mathbf{S}_{t+1} &= \mathbf{S}_t\Big(\alpha_{t+1}\big(\mathbf{I} - \beta_{t+1}\bm{k}_{t+1}\bm{k}_{t+1}^{\intercal}\big)\Big) + \beta_{t+1}\bm{v}_{t+1}\bm{k}_{t+1}^{\intercal} \\[4pt] &= \alpha_{t+1}\Big(\sum_{i=1}^{t}\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}\Big) - \alpha_{t+1}\beta_{t+1}\Big(\sum_{i=1}^{t}\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}\bm{k}_{t+1}\bm{k}_{t+1}^{\intercal}\Big) + \beta_{t+1}\bm{v}_{t+1}\bm{k}_{t+1}^{\intercal} \\[4pt] &= \sum_{i=1}^{t}\frac{\gamma_{t+1}}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal} + \underbrace{\beta_{t+1}\Big(\bm{v}_{t+1} - \sum_{i=1}^{t}\frac{\gamma_{t+1}}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal}\bm{k}_{t+1}\Big)}_{\textstyle \bm{u}_{t+1}}\bm{k}_{t+1}^{\intercal} \\[4pt] &= \sum_{i=1}^{t}\frac{\gamma_{t+1}}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal} + \underbrace{\frac{\gamma_{t+1}}{\gamma_{t+1}}}_{\textstyle 1}\bm{u}_{t+1}\bm{k}_{t+1}^{\intercal} \;=\; \sum_{i=1}^{t+1}\frac{\gamma_{t+1}}{\gamma_i}\bm{u}_i\bm{k}_i^{\intercal} \end{aligned}

\square

第二个等号到第三个等号是全部关键,值得拆开看:

  • αt+1γtγi=γt+1γi\alpha_{t+1}\cdot\frac{\gamma_t}{\gamma_i} = \frac{\gamma_{t+1}}{\gamma_i}–衰减系数被吸收进累积积的比值,这是 γ\gamma 定义为连乘才有的性质,也是四篇一直在用的那条恒等式。
  • 第二项里 αt+1\alpha_{t+1} 同样被吸收进 γt+1γi\frac{\gamma_{t+1}}{\gamma_i},然后整项与第三项合并、共同提出右侧的 kt+1\bm{k}_{t+1}^{\intercal}–括号里剩下的就是 ut+1\bm{u}_{t+1} 的定义。
  • 最后一步只是注意 γt+1/γt+1=1\gamma_{t+1}/\gamma_{t+1} = 1,于是新项能并入求和,下标从 tt 推到 t+1t+1,归纳闭合。

这个证明独立印证了 §2.3 那个坑。 注意第二项的系数是 αt+1βt+1\alpha_{t+1}\beta_{t+1}α\alphaβ\beta 都在,因为删除项 Iβkk\mathbf{I}-\beta\bm{k}\bm{k}^\intercal 作用在已经乘过 αt+1\alpha_{t+1} 的状态上。我在 §2.3 用倒代换推 dr\bm{d}_r 时漏掉这个 αr\alpha_r,refB 就与逐 token 递归差了 2.66×1012.66\times10^{-1}。两条路线在同一个位置要求同一个因子,可以互为校验。

实测验证这个命题(dk=5,dv=4,C=6d_k=5, d_v=4, C=6,fp64,逐步比对 St\mathbf{S}_t):

检验 结果
WY 闭式 vs 逐 token 递推(t=1..6t=1..6 逐步) 最大 1.67×10161.67\times10^{-16}
ut\bm{u}_t 是否等于 §2.3 的修正量 dt\bm{d}_t 1.11×10161.11\times10^{-16}
dt\bm{d}_t 漏掉 αt\alpha_t 偏差 2.99×1022.99\times10^{-2}

第二行说明论文的 ut\bm{u}_t 与我 §2.3 手推的 dt\bm{d}_t 是同一个量,只是推导路径不同:论文归纳法从闭式出发验证,§2.3 从递推倒代换构造。两者都给出同一个下三角系统,下面的 UT 变换对二者通用。

关于跨块:论文只推首块(S0=0\mathbf{S}_0=\mathbf{0})。实际 kernel 里每块的入口状态非零,ut\bm{u}_t 的定义要补上 βtγtS0kt-\beta_t\gamma_t\mathbf{S}_0^{\intercal}\bm{k}_t 这一项,即 §2.3 里那个 γrS[t]kr-\gamma^r\mathbf{S}_{[t]}\bm{k}_r照抄附录 A 而漏掉这项,是我最初 refB 失配(max abs 2.66×1012.66\times10^{-1}、rel L2 6.01×1026.01\times10^{-2})的另一个来源–首块测试全对、第二块开始错,这种症状基本可以直接定位到跨块项。

2.5 UT 变换:把解系统变成矩阵乘

定义(论文 §2.2 与 §3.3):

T[t]=[I+strictLower(diag(β[t])(Γ[t]K[t]K[t]))]1diag(β[t])\mathbf{T}_{[t]} = \Big[\mathbf{I} + \operatorname{strictLower}\big(\operatorname{diag}(\beta_{[t]})(\Gamma_{[t]} \odot \mathbf{K}_{[t]}\mathbf{K}_{[t]}^{\intercal})\big)\Big]^{-1}\operatorname{diag}(\beta_{[t]})

于是修正量一次算出:

Δ=TVU~Tdiag(γ)KWS[t]\Delta = \underbrace{\mathbf{T}\mathbf{V}}_{\widetilde{\mathbf{U}}} - \underbrace{\mathbf{T}\operatorname{diag}(\gamma)\mathbf{K}}_{\overleftarrow{\mathbf{W}}}\mathbf{S}_{[t]}^{\intercal}

这正是论文式 (11)(12) 里的 (U~[t]W[t]S[t])\big(\widetilde{\mathbf{U}}_{[t]} - \overleftarrow{\mathbf{W}_{[t]}}\mathbf{S}_{[t]}^{\intercal}\big)。完整的 chunkwise 算法:

S[t+1]=S[t]+ΔK[t]O[t]=Q[t]S[t]+(Q[t]K[t]Γ[t])Δ\begin{aligned} \mathbf{S}_{[t+1]} &= \overrightarrow{\mathbf{S}_{[t]}} + \Delta^{\intercal}\overrightarrow{\mathbf{K}_{[t]}} \\ \mathbf{O}_{[t]} &= \overleftarrow{\mathbf{Q}_{[t]}}\mathbf{S}_{[t]}^{\intercal} + \big(\mathbf{Q}_{[t]}\mathbf{K}_{[t]}^{\intercal} \odot \Gamma_{[t]}\big)\Delta \end{aligned}

三个箭头量与前两篇完全一致:qr=γrqr\overleftarrow{\bm{q}^r} = \gamma^r\bm{q}^rkr=γCγrkr\overrightarrow{\bm{k}^r} = \frac{\gamma^C}{\gamma^r}\bm{k}^rS=γCS\overrightarrow{\mathbf{S}} = \gamma^C\mathbf{S}

对比上一篇,结构上只有一处变化:块内那一项的右乘对象从 V[t]\mathbf{V}_{[t]} 变成了 Δ\Delta,而 Δ\Delta 需要解一个三角系统才能得到。β0\beta \to 0T0\mathbf{T} \to \mathbf{0}Δ0\Delta \to \mathbf{0},退化为纯衰减;α1\alpha \equiv 1ΓM\Gamma \to \mathbf{M},退化为纯 DeltaNet。

2.6 数值验证

B=2,H=2,N=12,dk=dv=4,C=4B=2, H=2, N=12, d_k=d_v=4, C=4αtU(0.90,0.999)\alpha_t \sim \mathcal{U}(0.90, 0.999)βtU(0.10,0.90)\beta_t \sim \mathcal{U}(0.10, 0.90)k\bm{k} 已 L2 归一化,fp64:

比较 max abs 误差 相对 L2
chunkwise + UT vs 逐 token 递归 8.88×10168.88 \times 10^{-16} 2.49×10162.49 \times 10^{-16}
α1\alpha \equiv 1(退化为纯 DeltaNet) 8.88×10168.88 \times 10^{-16}
β0\beta \to 0(退化为纯衰减) 1.62×10271.62 \times 10^{-27}

两个退化检验都必须做,理由与上一篇同构:α1\alpha \equiv 1Γ\Gamma 退化成 0/1 掩码,Γ\Gamma 里任何指数写错都看不出来;β0\beta \to 0 时整个三角系统消失,UT 部分的错误全部隐身。


3. 三角系统的数值性质

T\mathbf{T} 要求一个 C×CC \times C 矩阵的逆,这是本篇最值得担心的地方。但实际上它比看起来温和得多。

3.1 单位下三角,条件数可控

I+strictLower(A)\mathbf{I} + \operatorname{strictLower}(\mathbf{A})单位下三角矩阵–对角线恒为 1,严格下三角才是 A\mathbf{A}。这有两个直接后果:

  1. 行列式恒为 1,永不奇异。 不存在需要 pivoting 的情形。
  2. 求逆可以用前向替换,不需要通用矩阵求逆。

实测条件数(C=64C = 64k\bm{k} L2 归一化,αtU(0.9,0.999)\alpha_t \sim \mathcal{U}(0.9, 0.999),20 组随机采样):

βt\beta_t 采样区间 cond 中位数 cond 最大
[0.05, 0.2][0.05,\ 0.2] 1.76 1.91
[0.1, 0.9][0.1,\ 0.9] 5.67 6.86
[0.5, 0.99][0.5,\ 0.99] 7.88 9.05
[0.9, 0.999][0.9,\ 0.999] 10.38 11.99
[1.0, 1.99][1.0,\ 1.99] 34.27 40.85

β(0,1)\beta \in (0,1) 时条件数不超过 12,fp16 完全够用。这与上一篇那个「因式分解会溢出」的结论形成有意思的对比:那里是代数变形引入了 1/γ1/\gamma 这种指数增长量,而这里虽然要求逆,但矩阵结构本身保证了良态。

放开到 β(0,2)\beta \in (0,2)(论文脚注提到的负特征值情形)条件数跳到 34,仍可接受,但已需要留意。

3.2 前向替换代替求逆

T\mathbf{T} 从不需要显式求逆。逐行前向替换:

T[r,:]=erj<rA[r,j]T[j,:]\mathbf{T}[r,:] = \bm{e}_r - \sum_{j<r}\mathbf{A}[r,j]\,\mathbf{T}[j,:]

实测与 np.linalg.inv 的差异在 101710^{-17} 量级(三组随机种子分别 7.81,5.55,8.33×10177.81, 5.55, 8.33 \times 10^{-17}),即机器精度。

更进一步,T\mathbf{T} 本身也不必物化。 需要的只是 Δ=TR\Delta = \mathbf{T}\mathbf{R}(其中 R=diag(β)(Vdiag(γ)KS)\mathbf{R} = \operatorname{diag}(\beta)(\mathbf{V} - \operatorname{diag}(\gamma)\mathbf{K}\mathbf{S}^\intercal)),直接对 R\mathbf{R} 做前向替换:

Δ[r,:]=R[r,:]j<rA[r,j]Δ[j,:]\Delta[r,:] = \mathbf{R}[r,:] - \sum_{j<r}\mathbf{A}[r,j]\,\Delta[j,:]

这省掉一个 C×CC \times C 的中间量。代价是引入了 CC 步串行–这是本篇 kernel 与前两篇最本质的区别,下面 §5 会看到它如何限制并行度。


4. 四层参考实现

沿用前两篇的框架:

参考 实现方式 验证目标
A 逐 token 递归 递推式定义 St=St1(αt(Iβtkk))+βtvk\mathbf{S}_t = \mathbf{S}_{t-1}(\alpha_t(\mathbf{I}-\beta_t\bm{k}\bm{k}^\intercal)) + \beta_t\bm{v}\bm{k}^\intercal
B chunkwise + UT 变换 §2.3 的三角系统推导、§2.4 的 WY 表示与式 (11)(12)
C kernel 控制流镜像(切 DV、S\mathbf{S}^\intercal 布局、前向替换) kernel 结构
D 退化检验(α1\alpha \equiv 1 / β0\beta \to 0 与前两篇的一致性

4.0 状态方向

与前两篇一致:论文的状态是 SRdv×dk\mathbf{S} \in \mathbb{R}^{d_v \times d_k},kernel 存转置 SRdk×dv\mathbf{S}^\intercal \in \mathbb{R}^{d_k \times d_v}。本篇多一层要注意的是 ΔRC×dv\Delta \in \mathbb{R}^{C \times d_v}–它的布局与 V\mathbf{V} 相同,所以 ΔK\Delta^\intercal\overrightarrow{\mathbf{K}} 在转置布局下写作 KΔ\overrightarrow{\mathbf{K}}^\intercal\Delta

4.1 参考 A:逐 token 递归

1
2
3
4
5
6
7
8
9
10
11
12
13
14
def ref_A(Q, K, V, alpha, beta):
B, N, H, DK = Q.shape
DV = V.shape[-1]
O = np.zeros((B, N, H, DV))
for b in range(B):
for h in range(H):
S = np.zeros((DV, DK)) # 论文方向
for t in range(N):
k, v, q = K[b, t, h], V[b, t, h], Q[b, t, h]
a, be = alpha[b, t, h], beta[b, t, h]
# 广义 Householder:注意 α 乘在整个括号上
S = S @ (a * (np.eye(DK) - be * np.outer(k, k))) + be * np.outer(v, k)
O[b, t, h] = S @ q
return O

4.2 参考 B:chunkwise + UT 变换

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
def ref_B(Q, K, V, alpha, beta, C):
B, N, H, DK = Q.shape
DV = V.shape[-1]
O = np.zeros((B, N, H, DV))
for b in range(B):
for h in range(H):
S = np.zeros((DV, DK))
for c in range(N // C):
sl = slice(c * C, (c + 1) * C)
Qc, Kc, Vc = Q[b, sl, h], K[b, sl, h], V[b, sl, h]
a, be = alpha[b, sl, h], beta[b, sl, h]

g = np.cumprod(a) # γ^r,chunk 内重置
gC = g[-1]
i = np.arange(C)
Gam = np.where(i[:, None] >= i[None, :],
g[i][:, None] / g[i][None, :], 0.0) # Γ

# UT 变换:T = [I + strictLower(diag(β)(Γ ⊙ K K^T))]^{-1} diag(β)
A = np.diag(be) @ (Gam * (Kc @ Kc.T))
T = np.linalg.inv(np.eye(C) + np.tril(A, -1)) @ np.diag(be)

# Δ = Ũ - \overleftarrow{W} S^T
Delta = T @ Vc - T @ (np.diag(g) @ (Kc @ S.T))

O[b, sl, h] = (Qc * g[:, None]) @ S.T + ((Qc @ Kc.T) * Gam) @ Delta
S = gC * S + Delta.T @ (Kc * (gC / g)[:, None])
return O

与上一篇参考 B 的差别只有两处:多出 T 的构造,以及块内项的右乘对象从 Vc 变成 Delta衰减权重的三处落点完全没变

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
def ref_C(Q, K, V, alpha, beta, C, block_DV):
B, N, H, DK = Q.shape
DV = V.shape[-1]
O = np.zeros((B, N, H, DV))
for bv in range(DV // block_DV): # ← grid.x
dv = slice(bv * block_DV, (bv + 1) * block_DV)
for bbh in range(B * H): # ← grid.y
b, h = bbh // H, bbh % H
S = np.zeros((DK, block_DV)) # 转置布局
for c in range(N // C): # ← T.Pipelined
sl = slice(c * C, (c + 1) * C)
Qc, Kc, Vc = Q[b, sl, h], K[b, sl, h], V[b, sl, h, dv]
a, be = alpha[b, sl, h], beta[b, sl, h]
g = np.cumprod(a); gC = g[-1]
i = np.arange(C)
Gam = np.where(i[:, None] >= i[None, :],
g[i][:, None] / g[i][None, :], 0.0)

A = (Kc * be[:, None]) @ Kc.T * Gam # diag(β)(Γ ⊙ K K^T)
Tm = np.linalg.inv(np.eye(C) + np.tril(A, -1))
R = be[:, None] * Vc - be[:, None] * ((g[:, None] * Kc) @ S)
Delta = Tm @ R # C × block_DV

O[b, sl, h, dv] = (Qc * g[:, None]) @ S + ((Qc @ Kc.T) * Gam) @ Delta
S = gC * S + (Kc * (gC / g)[:, None]).T @ Delta
return O

注意 R 的写法:diag(β) 被展开成逐行乘 be[:, None],这与 kernel 里的做法一致(不物化对角矩阵)。

4.4 一致性验证

B=2,H=2,N=12,dk=dv=4,C=4B=2, H=2, N=12, d_k=d_v=4, C=4blockDV=2\text{block}_{DV}=2,fp64:

比较 max abs 误差 相对 L2
B chunkwise + UT vs A 逐 token 8.88×10168.88 \times 10^{-16} 2.49×10162.49 \times 10^{-16}
C kernel 镜像 vs A 逐 token 8.88×10168.88 \times 10^{-16} 2.61×10162.61 \times 10^{-16}
α1\alpha \equiv 1(纯 DeltaNet)B vs A 8.88×10168.88 \times 10^{-16}
β1012\beta \to 10^{-12}(纯衰减)B vs A 1.62×10271.62 \times 10^{-27}

另验证了 T\mathbf{T} 的两种等价写法(diag(β) @ (Γ * KKᵀ)(K * β) @ Kᵀ * Γ)差异 5.55×10175.55 \times 10^{-17},后者省一次对角矩阵构造。


5. TileLang kernel

grid 划分沿用前两篇:切 (bv, bh)(bv,\ b \cdot h),序列轴走 kernel 内的 T.Pipelined。但本篇多了一个 C×CC \times CA 矩阵和 CC 步串行的前向替换,寄存器与并行度都要重新算。

5.0 寄存器账

dk=dv=128d_k = d_v = 128C=64C = 64、128 线程:

fragment 上一篇 本篇 说明
S_f [dk,bDV][d_k, \text{bDV}] 32 reg 32 reg 不变
acc_o [C,bDV][C, \text{bDV}] 16 reg 16 reg 不变
Dtri / Gam [C,C][C,C] 32 reg 32 reg 衰减掩码
A [C,C][C,C] 32 reg 32 reg QK\mathbf{Q}\mathbf{K}^\intercal 复用
Amat [C,C][C,C] 32 reg 新增:三角系统系数
Delta [C,bDV][C, \text{bDV}] 16 reg 新增:修正量
合计 ≈ 112 reg ≈ 160 reg 上限 255

C=64C = 64 时仍有余量。C=128C = 128Gam + A + Amat 三张 C2C^2 表就是 384 reg,直接超限–本篇的 CC 实际上被限制在 64。

5.1 声明与权重预计算

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
@tilelang.jit(out_idx=[5])
def gdn_chunk(B, H, S, DK, DV, blk=64, block_DV=32,
num_stages=2, threads=128,
dtype=T.float16, accum_dtype=T.float32):
C = blk
NS = T.ceildiv(S, C)

@T.prim_func
def main(
Q: T.Tensor([B, S, H, DK], dtype),
K: T.Tensor([B, S, H, DK], dtype),
V: T.Tensor([B, S, H, DV], dtype),
Alpha: T.Tensor([B, S, H], accum_dtype), # α_t,数据依赖
Beta: T.Tensor([B, S, H], accum_dtype), # β_t,数据依赖
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)
D_s = T.alloc_shared([C, block_DV], dtype) # Δ 的 f16 副本

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) # Q K^T
A_cast = T.alloc_fragment([C, C], dtype)
Amat = T.alloc_fragment([C, C], accum_dtype) # 三角系统系数
Delta = T.alloc_fragment([C, block_DV], accum_dtype)
R = T.alloc_fragment([C, block_DV], accum_dtype) # 右端项
Gam = T.alloc_fragment([C, C], accum_dtype) # Γ
lg = T.alloc_fragment([C], accum_dtype) # log2 γ^r
g_r = T.alloc_fragment([C], accum_dtype) # γ^r
w_decay = T.alloc_fragment([C], accum_dtype) # γ^C/γ^r
be_f = T.alloc_fragment([C], accum_dtype) # β_r

αt\alpha_t 现在是数据依赖的,累积积必须在 kernel 内算。用 log 域 cumsum 而非 cumprod–上一篇实测 fp16 直接连乘在 C=128C = 128 时完全下溢:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
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)

# ── 累积衰减积:log 域前缀和 ──
for r in T.serial(C):
lg[r] = T.log2(Alpha[bb, s0 + r, bh])
for r in T.serial(1, C): # 串行前缀和
lg[r] += lg[r - 1]
for r in T.Parallel(C):
g_r[r] = T.exp2(lg[r]) # γ^r
w_decay[r] = T.exp2(lg[C - 1] - lg[r]) # γ^C/γ^r
be_f[r] = Beta[bb, s0 + r, bh]

# ── Γ[i,j] = γ^i/γ^j (j≤i),log 域相减后取指数 ──
for i, j in T.Parallel(C, C):
Gam[i, j] = T.if_then_else(
j <= i, T.exp2(lg[i] - lg[j]), 0.0)

Γ 的构造沿用上一篇的教训:log 域相减再 exp2,不能先各自取指数再相除,否则 1/γj1/\gamma_j 会溢出。条件求值也不能省成「先全算再掩」,j>ij > i 处指数为正会先溢出成 inf。

5.2 三角系统:构造系数并前向替换

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
# ── Amat = diag(β)(Γ ⊙ K K^T) 的严格下三角 ──
T.gemm(K_s, K_s, Amat, transpose_B=True, clear_accum=True)
for i, j in T.Parallel(C, C):
Amat[i, j] = T.if_then_else(
j < i, be_f[i] * Gam[i, j] * Amat[i, j], 0.0)

# ── 右端项 R = diag(β)(V - diag(γ) K S^T) ──
T.copy(S_f, S_s)
for r, d in T.Parallel(C, DK):
K_s[r, d] *= g_r[r] # diag(γ) K,原地
T.gemm(K_s, S_s, R, clear_accum=True) # (γK) S
for r, d in T.Parallel(C, block_DV):
R[r, d] = be_f[r] * (V_s[r, d] - R[r, d])

# ── 前向替换:Δ[r] = R[r] - Σ_{j<r} Amat[r,j] Δ[j] ──
for r in T.serial(C): # C 步串行,无法避免
for d in T.Parallel(block_DV):
Delta[r, d] = R[r, d]
for j in T.serial(r):
for d in T.Parallel(block_DV):
Delta[r, d] -= Amat[r, j] * Delta[j, d]

前向替换是本篇唯一的串行段。C=64C = 64 时是 64 步,每步内部 block_DV = 32 个元素并行–并行度只有 32,远低于 128 线程。这是 gated delta rule 相比前两篇的固有代价:三角系统的依赖链无法打破。

有两个缓解方向:一是增大 block_DV(但吃寄存器),二是把 CC 切成更小的子块做分块前向替换(blocked forward substitution),用矩阵乘处理子块间的耦合。后者是 fla 库实际采用的做法,但实现复杂度显著上升。

5.3 两项输出与状态更新

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
# ② 跨块项:\overleftarrow{Q} S^T
for r, d in T.Parallel(C, DK):
Q_s[r, d] *= g_r[r] # 位置三:query 侧
T.gemm(Q_s, S_s, acc_o, clear_accum=True)

# ③ 块内项:(Q K^T ⊙ Γ) Δ —— 注意要用未加权的原始 Q、K
T.copy(Q[bb, s0:s0+C, bh, :], Q_s)
T.copy(K[bb, s0:s0+C, bh, :], K_s) # K_s 被 diag(γ) 污染过
T.gemm(Q_s, K_s, A, transpose_B=True, clear_accum=True)
for i, j in T.Parallel(C, C):
A[i, j] *= Gam[i, j]
T.copy(A, A_cast)
T.copy(Delta, D_s)
T.gemm(A_cast, D_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] *= T.exp2(lg[C - 1]) # γ^C
for r, d in T.Parallel(C, DK):
K_s[r, d] *= w_decay[r] # 位置二:γ^C/γ^r
T.gemm(K_s, D_s, S_f, transpose_A=True) # \overrightarrow{K}^T Δ

5.4 五处次序约束

比上一篇多一处,全部写错都不报错、只算错:

  1. 状态更新排在输出写回之后SfS_f 在步骤②③被读时代表进块状态 S[t]\mathbf{S}_{[t]}
  2. S_f *= γ^CT.gemm 之前。先衰减旧状态再累加新贡献。
  3. K_s 被复用三次、污染两次。构造 Amat 用原始 K\mathbf{K},算 R 时被乘上 diag(γ)\operatorname{diag}(\gamma),块内项要用原始 K\mathbf{K}(须重载),状态更新时又要乘 γC/γr\gamma^C/\gamma^r这是本篇最容易错的地方–上一篇 K_s 只污染一次。
  4. Amat 必须只取严格下三角j<ij < i,不含对角)。含对角就变成了求解 (I+Afull)(\mathbf{I} + \mathbf{A}_{\text{full}}),对角上的 βrkr2=βr\beta_r\|\bm{k}_r\|^2 = \beta_r 会被重复计入。
  5. T.clear(S_f) 在流水线循环外,clear_accum=True 在循环内

5.5 与上一篇的改动汇总

位置 上一篇 本篇 新增开销
输入 Q, K, V + αt\alpha_t, βt\beta_t 两个张量 2 次 HBM 读 / chunk
累积积 编译期常量 kernel 内 log 域 cumsum CC 步串行前缀和
Γ\Gamma 编译期可算 运行时 exp2(lg[i]-lg[j]) C2C^2exp2
三角系统 Amat 构造 + 前向替换 C2C^2 reg + CC 步串行
块内右乘 V\mathbf{V} Δ\Delta 一次 C×bDVC \times \text{bDV} f16 转换
K 复用 污染 1 次 污染 2 次,重载 1 次 一次 shared 写入

6. 三篇对照

线性注意力 标量衰减 Gated DeltaNet
递推 S+vk\mathbf{S} + \bm{v}\bm{k}^\intercal αS+vk\alpha\mathbf{S} + \bm{v}\bm{k}^\intercal S(α(Iβkk))+βvk\mathbf{S}(\alpha(\mathbf{I}-\beta\bm{k}\bm{k}^\intercal)) + \beta\bm{v}\bm{k}^\intercal
块内掩码 M\mathbf{M}(0/1) Γ\Gamma(衰减感知) Γ\Gamma + 三角系统
块内右乘 V\mathbf{V} V\mathbf{V} Δ\Delta
串行段 前缀和 + 前向替换(各 CC 步)
C2C^2 fragment 1(A\mathbf{A} 2(+ Γ\Gamma 3(+ Amat
可用 CC 128 128(切 DV 后) 64
对应架构 RetNet / Lightning-Attn Gated DeltaNet / KDA

三篇的衰减权重落点完全一致q\overleftarrow{\bm{q}}k\overrightarrow{\bm{k}}S\overrightarrow{\mathbf{S}} 三处,从第一篇到第三篇没有变过。delta rule 加进来的是块内那一项的内容(VΔ\mathbf{V} \to \Delta),而不是衰减的结构。这是 GDN 论文那套箭头记号的价值:它把「衰减」这件事隔离成了一个可以独立理解的层面。


7. 总结

  1. delta rule 的本质是定向替换Iβtktkt\mathbf{I} - \beta_t\bm{k}_t\bm{k}_t^\intercal 先减掉 kt\bm{k}_t 键上的旧值 St1kt\mathbf{S}_{t-1}\bm{k}_t,再写入新值 βtvt+(1βt)St1kt\beta_t\bm{v}_t + (1-\beta_t)\mathbf{S}_{t-1}\bm{k}_t;与 kt\bm{k}_t 正交的记忆完全不受影响。这与标量衰减的一刀切互补:门控负责快速擦除,delta rule 负责精确修改。
  2. L2 归一化 k\bm{k} 不只是训练技巧kt=1\|\bm{k}_t\| = 1 时 Householder 变换的特征值是 {1βt}{1}dk1\{1-\beta_t\} \cup \{1\}^{d_k-1}βt(0,1)\beta_t \in (0,1) 时它们的绝对值都不超过 1,状态不会被越推越大;若不归一化,βtkt2>2\beta_t\|\bm{k}_t\|^2 > 2 会让特征值翻到 1-1 以下导致发散。
  3. test-time SGD 视角L=12Skv2\mathcal{L} = \frac{1}{2}\|\mathbf{S}\bm{k}-\bm{v}\|^2 的一步梯度下降就是 delta rule,βt\beta_t 是学习率、αt\alpha_t 是 weight decay。
  4. WY 表示的核心事实:CC 个 Householder 的乘积只是秩至多 CC 的修正Pr=Pr1(Iβrkrkr)=Pr1βrPr1krwrkr\mathbf{P}^r = \mathbf{P}^{r-1}(\mathbf{I}-\beta_r\bm{k}_r\bm{k}_r^\intercal) = \mathbf{P}^{r-1} - \underbrace{\beta_r\mathbf{P}^{r-1}\bm{k}_r}_{\bm{w}_r}\bm{k}_r^\intercal,每多一个 Householder 秩只增 1,于是连乘变求和:P=IWK\mathbf{P} = \mathbf{I}-\mathbf{W}^\intercal\mathbf{K}H=UK\mathbf{H} = \mathbf{U}^\intercal\mathbf{K} 同理。
  5. 矩阵值转移矩阵迫使块内计算变成解三角系统。前两篇的转移量是标量,可以查表;Householder 连乘不行。把递推重写成 Sr=αrSr1+drkr\mathbf{S}^r = \alpha_r\mathbf{S}^{r-1} + \bm{d}_r\bm{k}_r^\intercal 后倒代换,dr\bm{d}_r 的依赖关系构成单位下三角系统,系数矩阵恰是 diag(β)(ΓKK)\operatorname{diag}(\beta)(\Gamma \odot \mathbf{K}\mathbf{K}^\intercal)Γ\Gamma 在这里第二次出现
  6. 两条推导路线互为校验。论文附录 A 的归纳法从闭式 St=iγtγiuiki\mathbf{S}_t = \sum_i\frac{\gamma_t}{\gamma_i}\bm{u}_i\bm{k}_i^\intercal 出发验证,§2.3 的倒代换从递推构造,两者给出同一个量(实测 ut=dt\bm{u}_t = \bm{d}_t,差 1.11×10161.11\times10^{-16})。归纳法第二项的系数 αt+1βt+1\alpha_{t+1}\beta_{t+1} 独立确认了下一条那个必须带的 α\alpha。另注意附录 A 只推首块,跨块要补 βtγtS0kt-\beta_t\gamma_t\mathbf{S}_0^\intercal\bm{k}_t——照抄会出现「首块全对、第二块起错」的症状。
  7. dr\bm{d}_r 的定义里带 αr\alpha_rdr=βrvrαrβrSr1kr\bm{d}_r = \beta_r\bm{v}_r - \alpha_r\beta_r\mathbf{S}^{r-1}\bm{k}_r。删除项作用在已衰减的状态上,漏掉 αr\alpha_r 不会报错,α1\alpha \equiv 1 时也看不出来,只在两者都非退化时表现为数值不符。
  8. 三角系统是良态的,这是好消息I+strictLower(A)\mathbf{I} + \operatorname{strictLower}(\mathbf{A}) 是单位下三角,行列式恒为 1、永不奇异、不需 pivoting。实测 C=64C = 64β(0,1)\beta \in (0,1) 时条件数不超过 12,fp16 够用。放开到 β(0,2)\beta \in (0,2) 升到 34。
  9. T\mathbf{T} 与其逆都不必物化,直接对右端项做前向替换即可,实测与 inv 差异在 101710^{-17} 量级。省一个 C×CC \times C 中间量。
  10. 代价是并行度。前向替换的 CC 步依赖链无法打破,每步只有 block_DV 个元素并行(32 < 128 线程)。加上累积积的串行前缀和,本篇有两段 CC 步串行–这是 gated delta rule 相比前两篇的固有开销。
  11. CC 被限制在 64。三张 C2C^2 的 f32 fragment(Γ\GammaQK\mathbf{Q}\mathbf{K}^\intercalAmat)在 C=128C = 128 时合计 384 reg/thread,超过 255 上限。
  12. K_s 污染两次是最容易错的地方:构造 Amat 用原始 K\mathbf{K},算右端项时乘 diag(γ)\operatorname{diag}(\gamma),块内项要重载原始 K\mathbf{K},状态更新再乘 γC/γr\gamma^C/\gamma^r。上一篇只污染一次。
  13. 两个退化检验都必须做α1\alpha \equiv 1 回到纯 DeltaNet(实测 8.88×10168.88\times10^{-16})、β0\beta \to 0 回到纯衰减(1.62×10271.62\times10^{-27})。前者让 Γ\Gamma 的指数错误隐身,后者让整个 UT 部分隐身,缺一不可。
  14. 衰减的三处落点三篇未变q=γrq\overleftarrow{\bm{q}} = \gamma^r\bm{q}k=γCγrk\overrightarrow{\bm{k}} = \frac{\gamma^C}{\gamma^r}\bm{k}S=γCS\overrightarrow{\mathbf{S}} = \gamma^C\mathbf{S} 从第一篇到第三篇完全一致,delta rule 只改变块内项右乘的内容。

可迁移的启示:引入一个"看起来只是多一项"的机制,实际代价往往不在 FLOPs 而在依赖结构。delta rule 的算术开销并不大(一个 C×CC\times C 的三角系统),但它把块内计算从"纯矩阵乘"变成了"带串行依赖的求解",并行度从 128 掉到 32。评估一个改动时,先问它引入了什么依赖,再问它加了多少乘法。

参考

  • Gated Delta Networks(GDN):Yang, Kautz & Hatamizadeh, Gated Delta Networks: Improving Mamba2 with Delta Rule, arXiv:2412.06464,ICLR 2025。§3.1 给出 gated delta rule、§3.3 给出 chunkwise 算法与 UT 变换;§2.2 给出本文 §2.2 复述的无门控 WY/UT 推导(式 3–9)、附录 A(Extended WY Representation for Gated Delta Rule)给出本文 §2.4 复述的归纳法证明,原文只考虑首块(S0=0\mathbf{S}_0=\mathbf{0}
  • DeltaNet 的硬件高效 chunkwise 算法:Yang et al., 2024b(arXiv:2406.06484)
  • WY 表示:Bischof & Van Loan, 1985;UT 变换:Joffrain et al., 2006
  • delta rule 溯源:Widrow & Hoff, 1960;用于线性 Transformer:Schlag et al., 2021a

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 实测数据待补。

TileLang 实战:KDA 从零到一–Chunked 线性注意力

KDA(Kimi Delta Attention)的递推式一行就写完,但直接照着它写 kernel 需要同时处理四个机制:逐通道门控、delta rule 的三角求解、log 域 cumsum、跨 chunk 状态传递。四个机制耦合在一个 kernel 里,数值出错时无法判断误差来自哪一层。 本文采用递进式实现路径,每一级只引入一个新机制、每一级都可独立运行并做数值验证。本文覆盖第一级–移除全部衰减因子与删除因子,只保留 SS+KVS \leftarrow S + K^\top V

这一级的价值不在性能,而在于将 chunkwise 分解恒等式单独隔离验证。本文同时给出一个反直觉的结论:线性注意力的全部优势是把 O(N2D)O(N^2 D) 换成 O(ND2)O(N D^2),但本文采用的逐块独立重算策略会让 FLOPs 退回 O(N2)O(N^2)–与因果 FlashAttention 同量级。§6 会说明这个退化的根源在于把序列轴放进了 grid,并对照 TileLang 官方 chunk_delta_h 的做法:序列轴不进 grid,递推就退回单 block 内的顺序循环,NN 的次数才能保住

数学背景见《KDA 的来龙去脉》,TileLang 语言基础见《TileLang 编程基本知识点》,本文复用的 tile 切分与 fragment 累加模式见《TileLang 实战:FlashAttention 前向 Kernel》。


1. 递进路径:把 KDA 拆成可验证的增量

KDA 的完整递推式(Dt=diag(egt)D_t = \operatorname{diag}(e^{g_t})gtR<0dkg_t \in \mathbb{R}^{d_k}_{<0}):

St=St1Dt(Iβtktkt)+βtvtkt,ot=StqtS_t = S_{t-1} D_t (I - \beta_t k_t k_t^\top) + \beta_t v_t k_t^\top, \qquad o_t = S_t q_t

记号说明:Kimi Linear 论文(arXiv:2510.26692)原文写作 St=(Iβtktkt)Diag(αt)St1+βtktvtS_t = (I - \beta_t k_t k_t^\top)\operatorname{Diag}(\alpha_t) S_{t-1} + \beta_t k_t v_t^\topαt[0,1]dk\alpha_t \in [0,1]^{d_k}。本文取其转置形式以匹配 kernel 中 SS(dv,dk)(d_v, d_k) 内存布局,并沿用 GDN / Mamba 的 log 域记号 αt=egt\alpha_t = e^{g_t},因为实现里门控本来就存在 log 域。

按机制拆解,每一级只放开一个自由度:

级别 递推式 新增机制 新出现的实现结构
第一级(本文) SS+KVS \leftarrow S + K^\top V 分块恒等式本身、块内因果掩码
第二级 SγS+KVS \leftarrow \gamma S + K^\top V 标量衰减 块内权重从 0/1 变 γij\gamma^{i-j} 指数下三角
第三级 Sdiag(egt)S+S \leftarrow \operatorname{diag}(e^{g_t}) \cdot S + \cdots 逐 token 门控 log 域 cumsum、exp2 硬件指令
第四级 SS(Iβkk)+βvkS \leftarrow S(I - \beta k k^\top) + \beta v k^\top delta rule UT 变换、三角求解(wy_fast 雏形)
第五级 门控 + delta rule 二者耦合 五阶段流水线、跨 kernel 状态传递

最后一级即完整 GDN / KDA,对标 flash-linear-attention 中的 chunkwise 实现。

这样拆分的收益是误差定位能力:任何一级数值对不上,怀疑对象只有这一级新引入的那一个机制,前面几级已经验证过了。本文对应第一级,一个机制都还没引入,因此这里能验证的只有分块恒等式本身。


2. 数学推导:分块恒等式

本级递推不含遗忘项,状态单调累加:

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

展开为显式求和,因果且包含当前 token:

oi=jiqi(kjvj)=ji(qikj)vjo_i = \sum_{j \le i} q_i^\top (k_j v_j^\top) = \sum_{j \le i} (q_i \cdot k_j)\, v_j

这是标准线性注意力。与 softmax attention 的唯一区别是分数未经 softmax,因此求和顺序可自由交换–这是下述分块重写成立的前提。

2.1 将求和拆分为跨块与块内两部分

序列按 BCBC 切分为 NC=N/BCNC = N / BC 个 chunk。拆分的依据不是下标落在哪里,而是因果约束是否需要逐 token 判断

  • 当前块之前的 chunk:其内每一个 jj 对当前块里的每一个 ii 都满足 jij \le i,全部可见。因果约束在这些块上恒真、与 ii 无关,所以整块直接相乘就行。
  • 当前块:内部的 jj 才真正受 jij \le i 约束,只能取 kj,vjk_j, v_jjij \le i 的那一半。

oi=qi(c<cKcVc)跨块:全部可见,无需掩码+jchunk cji(qikj)vj块内:需因果掩码o_i = \underbrace{q_i^\top \Big( \sum_{c' < c} K_{c'}^\top V_{c'} \Big)}_{\text{跨块:全部可见,无需掩码}} + \underbrace{\sum_{\substack{j \in \text{chunk } c \\ j \le i}} (q_i \cdot k_j) v_j}_{\text{块内:需因果掩码}}

关键在第一项:正因为它的因果判断与 ii 无关,括号内的量对整个 chunk cc 才能是同一个矩阵,记作 Scprev=c<cKcVcS_c^{\text{prev}} = \sum_{c' < c} K_{c'}^\top V_{c'},形状 dk×dvd_k \times d_v,含义是处理该 chunk 之前已累积的状态。于是跨块贡献退化为一次矩阵乘 QcScprevQ_c S_c^{\text{prev}},块内贡献是一个 BC×BCBC \times BC 的下三角掩码矩阵乘:

Oc=QcScprev+tril(QcKc)VcO_c = Q_c S_c^{\text{prev}} + \operatorname{tril}(Q_c K_c^\top) V_c

图上半部分是前序 chunk 的注意力方阵——它自身也带因果三角,但这些计算在处理 chunk cc 时已经完成,结果被汇总进一个 dk×dvd_k \times d_v 的状态矩阵 ScprevS_c^{\text{prev}},不必再逐 token 展开。下半部分是本文要算的两项:左边整块打满阴影,表示 QcQ_c 的每一行都能无条件地乘上完整的历史状态,没有掩码;右边只有下三角有阴影,表示块内必须逐 token 判断 jij \le i跨块项的历史长度随序列增长,但它被压进固定形状的 ScprevS_c^{\text{prev}}——这正是线性复杂度的来源;块内三角形的边长恒为 BCBC,与序列长度无关。

这条恒等式是后续各级的公共基础,之后每一级都只是在它的两项上分别插入衰减权重,张量形状与 GEMM 调用次序保持不变。

2.2 掩码可以直接置零的原因

FlashAttention 中掩码必须填 -\infty,因为后续要经过 expe=0e^{-\infty} = 0 才能使被掩位置不产生贡献。线性注意力没有 softmax,掩码位置直接写 0 即可:

1
2
for i, j in T.Parallel(BC, BC):
A[i, j] = T.if_then_else(j <= i, A[i, j], 0.0) # 无 exp,掩码直接清零

这一行是线性注意力与 FA 的第一处分岔,也是「无 softmax」在代码上的全部体现–不需要 online 重标定、不需要 mm\ell 两个行状态、不需要输出阶段的最终除法


3. 四层 PyTorch 参考实现

验证 kernel 之前需要建立可信参考。本文构造三份 PyTorch fp64 实现,从「最直白」逐步过渡到「与 kernel 控制流同构」,相邻两层互相验证:

参考 实现方式 验证目标
A 逐 token 递归,双重 for 循环 递推式定义本身,几乎不可能写错
B 分块向量化,cumsum 求前缀状态 §2.1 分块恒等式的正确性
C 逐块独立重算,模拟 grid 与 T.Pipelined kernel 的控制流与访存次序

参考 C 是连接 PyTorch 与 TileLang 的关键一层–它的循环结构与 kernel 逐句对应,若 kernel 与 C 不符,问题必在 TileLang 语法或内存层级使用上,与数学无关。

3.1 参考 A:逐 token 递归

1
2
3
4
5
6
7
8
9
10
11
def ref_recurrent(Q, K, V):
"""最直白的实现:严格照抄递推式 S += k v^T; o = q S"""
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 += torch.outer(K[b, t, h].double(), V[b, t, h].double())
O[b, t, h] = Q[b, t, h].double() @ S
return O

3.2 参考 B:分块向量化

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
def ref_chunked(Q, K, V, BC):
"""分块向量化:用 cumsum 一次算出所有 chunk 的前缀状态"""
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)

outer = torch.einsum("bhcnd,bhcnv->bhcdv", Kc, Vc) # 每块自身的 K^T V
states = torch.cumsum(outer, dim=2) # 含自身的前缀和
states_prev = torch.cat( # 右移一格 → 不含自身
[torch.zeros_like(states[:, :, :1]), states[:, :, :-1]], dim=2)

O_inter = torch.einsum("bhcnd,bhcdv->bhcnv", Qc, states_prev)

A = torch.einsum("bhcnd,bhcmd->bhcnm", Qc, Kc) # 块内分数 [BC, BC]
mask = torch.tril(torch.ones(BC, BC, dtype=torch.bool, device=Q.device))
A = A.masked_fill(~mask, 0.0) # 线性注意力:置 0,非 -inf
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()

states_prev 的右移拼接对应 c<c\sum_{c' < c} 中的严格小于号:cumsum 给出的是含自身的前缀和,右移一格补零才是处理该 chunk 之前的累积状态。这是分块线性注意力最易出错的一行–不右移等价于把当前块的 KVK^\top V 重复计入,块内贡献会被计算两次。该错误的量级实测为相对 L2 误差 7.78×1017.78 \times 10^{-1},属于结构性错误而非精度问题,容易识别。

3.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
def ref_kernel_mimic(Q, K, V, 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)
mask = torch.tril(torch.ones(BC, BC, dtype=torch.float64, device=Q.device))

for bz in range(B): # grid.z ← batch
for by in range(H): # grid.y ← head
for bx in range(NC): # grid.x ← Q 块
sl = slice(bx * BC, (bx + 1) * BC)

# ① T.clear(S_f); for c in T.Pipelined(bx): T.gemm(..., transpose_A=True)
S = torch.zeros(D, D, dtype=torch.float64, device=Q.device)
for c in range(bx):
cs = slice(c * BC, (c + 1) * BC)
S += K[bz, cs, by, :].double().T @ V[bz, cs, by, :].double()

# ② T.gemm(Q_s, S_s, acc_o)
Qb = Q[bz, sl, by, :].double()
acc = Qb @ S

# ③ T.gemm(Q_s,K_s,A,transpose_B=True) → 掩码 → T.gemm(A_cast,V_s,acc_o)
Kb = K[bz, sl, by, :].double()
Vb = V[bz, sl, by, :].double()
acc += ((Qb @ Kb.T) * mask) @ Vb

O[bz, sl, by, :] = acc
return O

对照参考 B 与参考 C 可以看出两种前缀状态求法的区别:B 用 cumsum 一次性算出全部 NCNC 个前缀状态、总代价 O(NC)O(NC);C 的每个 bxbx 独立重算、总代价 O(NC2)O(NC^2)kernel 采用的是 C 的策略,原因与代价见 §6。

3.4 四层参考的一致性验证

B=2,H=2,N=12,D=4,BC=4B=2, H=2, N=12, D=4, BC=4,fp64(numpy 复现):

比较 max abs 误差 相对 L2
B 分块向量化 vs A 逐 token 递归 3.55×10153.55 \times 10^{-15} 1.62×10161.62 \times 10^{-16}
C kernel 结构镜像 vs A 逐 token 递归 7.11×10157.11 \times 10^{-15} 1.86×10161.86 \times 10^{-16}
C vs B 4.44×10154.44 \times 10^{-15} 1.78×10161.78 \times 10^{-16}
参考 B 去掉右移(错误实现) 1.97×1011.97 \times 10^{1} 7.78×1017.78 \times 10^{-1}

前三行误差均在 fp64 机器精度量级(ε2.2×1016\varepsilon \approx 2.2 \times 10^{-16}),恒等式与三份实现均无误。第四行是刻意引入的错误,用于确认该验证流程对结构性错误敏感。

参考 D(§6.4.2 的序列内循环形式)另外验证一件事–换 grid 划分、切 DV 后算的还是同一个恒等式B=2,S=256,H=3,DK=64,DV=128,C=64B{=}2, S{=}256, H{=}3, DK{=}64, DV{=}128, C{=}64,fp64:

对象 相对 L2(vs 参考 A)
参考 D,blockDV=32\text{block}_{DV} = 32 5.60×10165.60 \times 10^{-16}
参考 D,blockDV=64\text{block}_{DV} = 64 5.60×10165.60 \times 10^{-16}
参考 D,blockDV=128\text{block}_{DV} = 128(不切) 5.60×10165.60 \times 10^{-16}
参考 D,状态更新提到写回之前 6.98×1016.98 \times 10^{-1}

三个 blockDV\text{block}_{DV} 误差完全相同,这就是「DV 切块零依赖」的直接证据–切不切、切多细,算出来的是比特级相同的东西,因为每个 dv 竖条的浮点累加顺序本就不受其他竖条影响。对照最后一行:仅仅把状态更新从循环尾部提到开头,误差就跳到 70%–右移语义是硬要求。

以上均为 numpy fp64 实测。TileLang kernel 本身需要 CUDA 设备,本文未给出实测数字–待真卡跑通后单独补充,此处不做性能推测。预期的主导误差源是 T.copy(S_f, S_s) 把 f32 状态降至 f16 落 shared(Tensor Core MMA 的输入必须是低精度);状态 SS 是多个块累加的结果,越靠后的 block 累加项越多,误差随 bxbx 单调增长,影响远大于块内 AA 的那次降精度。

3.5 einsum 输出下标的约束

参考 B 中若把 torch.einsum("bhcnd,bhcnv->bhcdv", ...) 误写为 ->bhcdd(意图表达「输出是 d×dd \times d 方阵」),会直接抛出异常:

1
ValueError: einstein sum subscripts string includes output subscript 'd' multiple times

einsum 的输出下标不允许重复–重复下标在输出侧的语义是取对角线。而 KVK^\top V 的两个维度虽然长度均为 DD语义上分别是 key 维与 value 维,必须使用不同字母 dv。同理 bhcnd,bhcdd->bhcnv 也不合法,因为输出的 v 从未在输入中出现。dk=dvd_k = d_v 时长度相同掩盖了语义差异,一旦 KDA 中 dkdvd_k \ne d_v,该疏忽会立即表现为形状错误。


4. TileLang kernel:三步实现

grid 划分与 FA 一致–按 Q 块切分所有权,每个 block 负责输出一个 BC×dBC \times d 的 tile:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
@tilelang.jit(out_idx=[3])
def linattn_chunk(batch, heads, seq_len, dim, blk, 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) # 状态降精度落地,喂第二个 gemm
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)

4.1 步骤①:流式累加前缀状态

对应参考 C 中的 for c in range(bx)

1
2
3
4
5
6
7
T.clear(S_f)
for c in T.Pipelined(bx, num_stages=num_stages): # 循环上界是 runtime 的 bx
T.copy(K[bz, c * BC:(c + 1) * BC, by, :], K_s)
T.copy(V[bz, c * BC:(c + 1) * BC, by, :], V_s)
T.gemm(K_s, V_s, S_f, transpose_A=True) # S_f += K_s^T @ V_s ← TileLang 的 T.gemm 默认累加

T.copy(S_f, S_s) # f32 → f16,本 kernel 主要精度损失点

代码里看不到 +=,但这个循环确实是在累加:T.clear(S_f) 先把 fragment 清零,之后每次 T.gemm 都在 S_f 原地累加,因此循环等价于 Sf=c<bxKcVcS_f = \sum_{c < bx} K_c^\top V_c。累加语义来自 Tensor Core MMA 的基本形式–MMA 指令算的是 d = a @ b + c,累加器既是输入也是输出,T.gemm 默认沿用这一行为;若要覆盖而非累加,需显式传 clear_accum=True

T.Pipelined(bx) 的上界是 block 索引而非编译期常量–不同 block 的循环次数不同,第 0 块一次都不执行(S=0S = 0),最后一块需执行 NC1NC-1 次。TileLang 支持 runtime 上界的流水线,代价是各 block 负载严重不均衡,尾部 block 构成整个 kernel 的关键路径。

这里每个 block 都在重算自己需要的前缀状态,而递推式本身只需一次加法。为什么本文要这么写、以及官方实现如何避开它,见 §6.2 与 §6.4。

4.2 步骤②③:两项贡献求和

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
# ② 跨块:O = Q_c @ S_prev
T.copy(Q[bz, bx * BC:(bx + 1) * BC, by, :], Q_s)
T.clear(acc_o)
T.gemm(Q_s, S_s, acc_o)

# ③ 块内:O += tril(Q_c K_c^T) @ V_c
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] = T.if_then_else(j <= i, A[i, j], 0.0)
T.copy(A, A_cast) # f32 → f16 喂 MMA
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, :])

步骤③开头必须重新 T.copy 当前块的 K/V。步骤①的流水线循环结束时,K_s / V_s 中残留的是第 bx1bx-1 块的数据,且在 num_stages > 1 时具体残留哪一块取决于流水线的展开方式,不可假设。复用 shared buffer 节省了显存,但必须显式重载–这是 tile 编程中典型的隐式状态陷阱。参考 C 里没有这个问题,因为 Python 每轮都重新切片,不存在缓冲区复用。

T.copy(A, A_cast) 这次 f32→f16 转换与 FA 中 P~\tilde{P} 喂入第二个 GEMM 前的降精度是同一操作–Tensor Core 的 MMA 输入必须是低精度,仅累加器为 f32。

4.3 kernel 与参考 C 的逐句对应

参考 C(PyTorch) kernel(TileLang) 说明
for bz / by / bx 三重循环 T.Kernel(ceildiv(N,BC), heads, batch) 循环变并行 grid
S = torch.zeros(D, D) T.clear(S_f) 状态初始化,fragment 常驻寄存器
for c in range(bx) T.Pipelined(bx, num_stages) 顺序循环变软件流水线
S += K[cs].T @ V[cs] T.gemm(K_s, V_s, S_f, transpose_A=True) 显式 shared 暂存 + Tensor Core
acc = Qb @ S T.gemm(Q_s, S_s, acc_o) 需先 T.copy(S_f, S_s) 降精度
(Qb @ Kb.T) * mask T.gemm(..., transpose_B=True) + T.Parallel 掩码 掩码从广播乘变逐元素条件
acc += (...) @ Vb T.copy(A, A_cast) + T.gemm(A_cast, V_s, acc_o) 多一次 f32→f16 转换
O[bz, sl, by, :] = acc T.copy(acc_o, O_s) + T.copy(O_s, O[...]) 经 shared 中转写回 HBM

两处 PyTorch 中不存在的操作:f32→f16 显式降精度(Tensor Core 输入约束)与shared buffer 中转(内存层级手动管理)。这两项也正是 kernel 与参考 C 数值差异的全部来源。


5. 与 FlashAttention 的三点结构差异

本文实现复用了 FA 的全部 tile 模式,但有三处必须修改:

维度 FlashAttention 线性注意力 原因
归一化 online softmax,维护 mm\ell 两个行状态,每块重标定 无,掩码直接置 0 exp,求和顺序可交换
循环携带的量 无跨块状态,每个 Q 块独立 S[dk,dv]S[d_k, d_v] 是跨迭代累加的状态 递推式本身带状态
T.gemm policy 必须 FullRow–行归约要求整行在同一 warp 默认 Square 即可 无按行归约

第三点值得展开:FA 需对分数块做 rowmax / rowsum,若一行被切分到多个 warp,归约就需跨 warp 通信,因此必须用 FullRow policy 强制整行不拆。线性注意力的块内矩阵 AA 计算完成后仅做逐元素掩码,无任何跨列归约,warp 划分方式不影响正确性,编译器可自由选择寄存器分布最优方案。

第二点需要说明的是,本文实现并没有真正兑现这一行:逐块独立重算把状态依赖藏起来了,每个 block 都从零重算自己需要的前缀状态,块之间不传递任何东西。这不是偷懒,而是本文这种「序列轴占据 grid.x」的划分下,block 之间无法顺序传递状态。换一种 grid 划分就没有这个限制,见 §6.4。


6. 代价账本:线性复杂度从何而来,如何被丢掉,以及如何拿回

6.1 两条路线的规模差异

先把线性注意力的本质说清楚。同样的输出,有两条算法:

路线 计算方式 规模
softmax / FA 先算分数矩阵 QKQK^\topN×NN \times N),再乘 VV O(N2D)O(N^2 D)
线性注意力 先算状态 S=KVS = K^\top VD×DD \times D),再乘 QQ O(ND2)O(N D^2)

区别在结合律往哪边括:(QK)V(QK^\top)V 要物化一个 N×NN \times N 的中间矩阵,Q(KV)Q(K^\top V) 物化的是 D×DD \times DNN 的次数从 2 降到 1,DD 的次数从 1 升到 2–这才是线性注意力唯一的、也是全部的优势来源,代价是 O(D2)O(D^2) 的固定状态取代了 O(N2)O(N^2) 的自由分数矩阵,表达能力随之受限。

chunkwise 形式是这两条路线的混合:跨块走 SS 路线(每块一次 QcScprevQ_c S_c^{\text{prev}},共 N2D2N \cdot 2D^2),块内走 QKQK^\top 路线(每块一个 BC×BCBC \times BC 分数矩阵,共 N4BCDN \cdot 4 BC D),状态更新本身再花 N2D2N \cdot 2D^2。理想总量:

FLOPsideal=4ND2+4NBCD=4ND(D+BC)\text{FLOPs}_{\text{ideal}} = 4N D^2 + 4N\,BC\,D = 4ND(D + BC)

NN 是一次方。与因果 FA 的 2N2D2N^2 D 相比:

FLOPsidealFLOPscausal FA=2(D+BC)N\frac{\text{FLOPs}_{\text{ideal}}}{\text{FLOPs}_{\text{causal FA}}} = \frac{2(D + BC)}{N}

比值按 1/N1/N 衰减–序列越长优势越大,这是线性注意力值得做的全部理由。

6.2 本文为何没有顺着递推式只加一次

理想账本里的 NCNCKVK^\top V,对应的是顺序递推:

Sc=Sc1+Kc1Vc1S_{c} = S_{c-1} + K_{c-1}^\top V_{c-1}

每个 chunk 只做一次 GEMM,读上一步的结果、加上自己这一块。参考 B 就是这么算的(cumsum 一次扫完),CPU 上顺序执行毫无问题。

但 kernel 做不到这一点,原因在 grid 的划分方式。 §4 把所有权按 Q 块切分,bxbxgrid.x 的索引:所有 NCNC 个 block 由硬件并行调度,执行顺序不确定、彼此之间没有同步点、也没有共享的可写缓冲区。block bxbx 若想读 SbxS_{bx},就得等 block bx1bx-1 算完并把结果落到某处——这两件事在单个 kernel 内都不成立:CUDA 不保证 block 间的执行次序(bx1bx-1 可能还没启动),也没有跨 block 的 barrier 可用。

顺序递推要求的是"前一步已完成",而 grid 提供的是"所有步同时开始"。但这个矛盾是本文自己造出来的–它成立的前提是"序列轴必须占据 grid 的一维"。放弃这个前提,矛盾就不存在了,见 §6.4。

本文仍保留逐块独立重算:它让 kernel 保持单文件闭环、控制流与参考 C 严格对应,适合把分块恒等式本身隔离出来验证。下面先算清这个选择的代价。

6.3 冗余重算把 NN 的次数还了回去

bxbx 块执行 bxbxKVK^\top V,全部 block 合计 c=0NC1c=NC(NC1)/2\sum_{c=0}^{NC-1} c = NC(NC-1)/2 次,而顺序递推共 NCNC 次。冗余系数 (NC1)/2(NC-1)/2,且NN 线性增长–正是这个增长把 NN 的次数从 1 顶回 2:

FLOPs=NC(NC1)22BCD2N2D2BC\text{FLOPs}_① = \frac{NC(NC-1)}{2} \cdot 2\, BC\, D^2 \approx \frac{N^2 D^2}{BC}

于是与因果 FA 的比值退化成常数:

FLOPsFLOPscausal FA=D2BC\frac{\text{FLOPs}_①}{\text{FLOPs}_{\text{causal FA}}} = \frac{D}{2\,BC}

这个比值与 NN 无关,恰恰是失败的判据:它说明本文实现与 FA 同属 O(N2)O(N^2)1/N1/N 的衰减优势被完全抹掉了。实测账本(D=64D = 64BC=64BC = 64):

NN NCNC 冗余系数 本文实现(步骤①) 理想线性注意力 因果 FA 本文/FA 理想/FA
512 8 3.5x 0.015 GFLOP 0.017 GFLOP 0.034 GFLOP 0.438 0.500
2048 32 15.5x 0.260 GFLOP 0.067 GFLOP 0.537 GFLOP 0.484 0.125
8192 128 63.5x 4.261 GFLOP 0.268 GFLOP 8.590 GFLOP 0.496 0.031
16384 256 127.5x 17.113 GFLOP 0.537 GFLOP 34.360 GFLOP 0.498 0.016
65536 1024 511.5x 274.609 GFLOP 2.147 GFLOP 549.756 GFLOP 0.500 0.004

看最后两列的走向:本文/FA 收敛到常数 D/(2BC)=0.5D/(2BC) = 0.5,理想/FA 按 1/N1/N 一路衰减到 0.004。 N=65536N = 65536 时理想实现只需 2.1 GFLOP,本文实现要 274.6 GFLOP–差 128 倍,而这个倍数还会随 NN 继续涨。

BCBC 的取舍随之明确:

BCBC D/(2BC)D/(2BC) 含义
32 1.000 与因果 FA 计算量持平,块过小无收益
64 0.500 默认值,寄存器压力可控
128 0.250 冗余减半,但 AA128×128128 \times 128 f32 fragment,每线程约 128 个寄存器,易溢出
256 0.125 理论最省,实际 shared memory 与寄存器均无法容纳

但要注意这张表只是在常数上打折,O(N2)O(N^2) 的量级不变。增大 BCBC 能线性减少冗余,但 AABC×BCBC \times BC 的 f32 fragment,BC=128BC = 128 时寄存器压力已接近溢出边界。真正的解法不是调 BCBC,见下一节。

6.4 更好的方案:把序列轴从 grid 里拿掉

TileLang 官方 examples/gdn 中的 chunk_delta_h 给出了另一种划分。它的 grid 只有两维:

1
2
3
4
5
6
7
8
with T.Kernel(T.ceildiv(DV, block_DV), B * H, threads=threads) as (bv, bbh):
...
for i_s in T.Pipelined(T.ceildiv(S, block_S), num_stages=num_stages):
# 存上一轮的状态快照,供下游 kernel 使用
T.copy(b_h_shared, h[bb, i_s, bh, 0:DK, bv * block_DV:(bv + 1) * block_DV])
...
T.gemm(K_shared, V_new_shared, b_h_fragment, transpose_A=True) # 状态原地累加
T.copy(b_h_fragment, b_h_shared)

序列不在 grid 里,而是 kernel 内部的一个顺序循环。 b_h_fragment 成为 loop-carried 的寄存器变量,跨 chunk 一路累加–这正是 Sc=γcSc1+KcVcS_c = \gamma_c S_{c-1} + K_c^\top V_c 的直接翻译,每个 chunk 只做一次 KVK^\top V,一次都不重算。§6.2 里那个"block 间无法同步"的矛盾根本不会出现,因为递推的顺序性被限制在单个 block 内部,用寄存器解决了,从来不需要跨 block 通信

并行度从另外两个轴补回来:

来源 数量(B=1,H=32,DV=128,blockDV=32B=1,H=32,DV=128,block_{DV}=32
bv DV 切块 128/32=4128/32 = 4
bbh batch ×\times head 融合 1×32=321 \times 32 = 32

合计 128 个 block,足够填满 SM。

6.4.1 为什么只切 DV、不切 DK

这一点容易给出错误的理由。若只看状态更新那一个 GEMM:

Sc+=KcVc,Kc:[DK,C], Vc:[C,DV]S_c \mathrel{+}= K_c^\top V_c,\qquad K_c^\top: [DK, C],\ V_c: [C, DV]

它的收缩维是 chunk 长度 CCDKDK 是 M 维、DVDV 是 N 维。而切 M 维只是输出 tiling,各 block 写各自的行,同样零依赖–所以「DKDK 是 M 维、切开要跨 block 归约」是站不住的,单看这一步切哪边都行。

不对称性来自下游谁消费这个状态。把一个 chunk 里三个 GEMM 的收缩维列出来:

GEMM 形状 收缩维 DKDK 的角色 DVDV 的角色
状态更新 KcVcK_c^\top V_c [DK,C]×[C,DV][DK,C] \times [C,DV] CC M(自由) N(自由)
跨块读出 QcScprevQ_c S_c^{\text{prev}} [C,DK]×[DK,DV][C,DK] \times [DK,DV] DKDK 收缩 N(自由)
块内项 QcKcQ_c K_c^\top [C,DK]×[DK,C][C,DK] \times [DK,C] DKDK 收缩 不出现

结论一句话:DVDV 在三个 GEMM 里始终是自由维,DKDK 在其中两个里是收缩维。

把两种切法的张量尺寸画出来,区别就是一个词:拼接还是相加

  • 切 DV:block (bv,bbh)(bv, bbh) 独占 S[:,dv]S[:, \text{dv}] 这一竖条,自己跑完整条递推,输出 O[:,dv]=QcSprev[:,dv]+tril(QcKc)Vc[:,dv]O[:, \text{dv}] = Q_c S^{\text{prev}}[:, \text{dv}] + \operatorname{tril}(Q_cK_c^\top)V_c[:, \text{dv}]C×DV/nC \times DV/n 的条带,已是最终值,直接写回。状态和输出两步都只有拼接,零累加、零跨块通信。
  • 切 DK:状态更新那步确实也没事(Si=KiVS^i = K_i^\top VDK/n×DVDK/n \times DV 的行条带,也不重叠),但读出就崩了QQ 被迫跟着切成 C×DK/nC \times DK/n,算出的 o~i=QiSi\tilde{o}^i = Q^i S^iC×DVC \times DV 全宽nno~i\tilde{o}^i 盖在同一块 OO 上,各含 1/n 根收缩轴的贡献,Oc=io~iO_c = \sum_i \tilde{o}^i,少加任何一份结果就错。这就是 split-K:每个 chunk 每条序列都要归约一次 [C,DV][C, DV] 的中间结果,摊销不掉,只能选 fp32 atomic add(非确定性 + 带宽)或另开 workspace 起第二个 kernel。

更麻烦的是块内项:QcKcQ_cK_c^\top 也沿 DKDK 收缩,只持有 dk\text{dk} 子块的 block 根本算不出完整的 [C,C][C,C] 转移矩阵(到了 DeltaNet 那级还要求 (I+tril(diag(β)KK))1(I + \operatorname{tril}(\operatorname{diag}(\beta)K K^\top))^{-1},逻辑上无法切)。

图里“QQ 也要切”那一行值得单独指出:DK 是吸收维,一旦切它,所有沿 DK 寻址的张量(QQKKSS 的行)都被动跟随,而产出反而变成全宽部分和–输入维度变小、输出维度不变,这个尺寸上的不匹配就是归约的必然信号。切 DV 恰好相反:输入切窄一根,输出也跟着窄一根。

还有一个纯工程的理由:状态 fragment 是 [DK,blockDV][DK, \text{block}_{DV}] 的 f32,DK=DV=128DK{=}DV{=}128、128 线程时,不切 DV 要 128×128/128=128128\times128/128 = 128 个寄存器/线程,直接溢出;切成 blockDV=32\text{block}_{DV}{=}32 降到 32 个。DV 切块同时解决了 SM 占用率和寄存器压力两件事,且不引入任何归约。

6.4.2 完整例子:按 (b,h,dv)(b, h, dv) 切 tile,沿 SS 走 Pipelined

把上面的结论落成第一级(无衰减、无删除)的可运行形态。与 §4 相比只改一件事:grid 里的序列轴换成 DV 轴,序列退回 kernel 内部的顺序循环

所有权划分:block (bv,bbh)(bv, bbh) 负责 O[b,:,h, dv 竖条]O[b, :, h, \ \text{dv 竖条}]整条序列的一个 value 通道子集。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
@tilelang.jit(out_idx=[3])
def linattn_seq_in_loop(B, H, S, DK, DV, block_S=64, block_DV=32,
num_stages=2, threads=128,
dtype=T.float16, accum_dtype=T.float32):
C = block_S
NS = T.ceildiv(S, C)

@T.prim_func
def main(
Q: T.Tensor([B, S, H, DK], dtype),
K: T.Tensor([B, S, H, DK], dtype),
V: T.Tensor([B, S, H, DV], dtype),
O: T.Tensor([B, S, H, DV], dtype),
):
# grid 两维:DV 切块 x (batch·head) 融合;序列轴不在这里
with T.Kernel(T.ceildiv(DV, block_DV), B * H, threads=threads) as (bv, bbh):
bb = bbh // H # batch 索引
bh = bbh % H # head 索引
dv0 = bv * block_DV # 本 block 负责的 value 通道起点

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) # 状态的 f16 副本,喂 MMA
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)

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 竖条

# ② 跨块项:用「进入本 chunk 前」的状态,必须先读后更新
T.copy(S_f, S_s) # f32 -> f16
T.gemm(Q_s, S_s, acc_o, clear_accum=True) # [C,DK]x[DK,bDV]

# ③ 块内项:tril(Q K^T) V,沿 DK 收缩,本 block 持有完整 DK
T.gemm(Q_s, K_s, A, transpose_B=True, clear_accum=True)
for i, j in T.Parallel(C, C):
A[i, j] = T.if_then_else(j <= i, A[i, j], 0.0)
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])

# ① 状态更新:放在 ②③ 之后 == 递推式里的「右移一格」
T.gemm(K_s, V_s, S_f, transpose_A=True) # [DK,C]x[C,bDV]

return main

四处必须讲清的细节:

  1. 循环体内的顺序就是右移语义。 T.gemm(K_s, V_s, S_f, ...) 必须排在输出写回之后:SfS_f 在读出时代表 Scprev=c<cKcVcS^{\text{prev}}_c = \sum_{c' < c} K_{c'}^\top V_{c'},本 chunk 自己的贡献由 tril\operatorname{tril} 那一项负责。把状态更新提到前面,就变成了 inclusive 前缀,块内项会被重复计入–这是本级唯一的结构性错误点,且数值上表现为「整体偏大」而非 NaN,很容易漏掉。§3.4 里刻意去掉右移的那个错误实现,对应的就是这里的顺序写反。
  2. T.clear(S_f) 在循环外,clear_accum=True 在循环内。 状态要跨迭代累加,所以只能循环外清零一次;而 acc_oA 每个 chunk 都是全新的,进了流水线循环就不能再用 T.clear(清零会被排到流水线的错误阶段),必须靠 clear_accum=True 在 MMA 那一刻覆盖累加器。
  3. num_stages 只能盖住访存,盖不住计算。 S_f -> S_s -> T.gemm -> S_f 构成一条真正的循环依赖,编译器无法把相邻两个 chunk 的计算重叠;流水线的收益全部来自把下一个 chunk 的 Q/K/V 的 HBM→shared 搬运提前发出。所以 num_stages=2 基本够用,继续加只是多占 shared。
  4. Q/K 被 dv 方向重复读。 每个 bvbv 都要读整份 QQKK(各 DKDK 全宽),读放大系数 DV/blockDVDV/\text{block}_{DV}VVOO 则严格切分不重复。这是切 DV 唯一的代价–多读,但不归约。相比之下切 DK 是要归约,这就是取舍的本质区别。

对应的 PyTorch 参考(延续 §3 的写法,可直接跟参考 A/B 对齐):

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
def ref_D_seq_in_loop(Q, K, V, C, block_DV):
B, S, H, DK = Q.shape
DV = V.shape[-1]
O = torch.zeros(B, S, H, DV, dtype=torch.float64)
mask = torch.tril(torch.ones(C, C, dtype=torch.float64))
for b in range(B):
for h in range(H):
for dv0 in range(0, DV, block_DV): # <- grid.x
dv = slice(dv0, dv0 + block_DV)
Sm = torch.zeros(DK, block_DV, dtype=torch.float64)
for i_s in range(S // C): # <- T.Pipelined
sl = slice(i_s * C, (i_s + 1) * C)
Qb, Kb, Vb = Q[b, sl, h, :], K[b, sl, h, :], V[b, sl, h, dv]
O[b, sl, h, dv] = Qb @ Sm + (Qb @ Kb.T * mask) @ Vb
Sm = Sm + Kb.T @ Vb # 右移:更新在写回之后
return O

这份参考与参考 B 的相对 L2 误差在 fp64 下应落在 101610^{-16} 量级:它算的是同一个恒等式,只是把「每块重算前缀」换成了「顺序携带前缀」,数学上完全等价,冗余系数从 (NC1)/2(NC-1)/2 降到 1。

FLOPs 回到 SS 的一次方–每 chunk 每 dv-block 两次 GEMM,求和得 4SDKDV4\,S\,DK\,DV。与本文实现对比(D=64D = 64BC=64BC = 64):

NN 本文步骤① 序列出 grid 倍数
8192 4.261 GFLOP 0.067 GFLOP 63.5x
16384 17.113 GFLOP 0.134 GFLOP 127.5x
32768 68.585 GFLOP 0.268 GFLOP 255.5x
65536 274.609 GFLOP 0.537 GFLOP 511.5x

倍数恰好是冗余系数 (NC1)/2(NC-1)/2,随 NN 线性增长。

6.4.3 融合还是拆成两个 kernel

值得注意的是,§6.4.2 那个 kernel 没有把状态快照写回 HBM–因为它在同一个循环里就把 OO 算完了:block 持有完整的 DKDKQcSprevQ_c S^{\text{prev}}tril(QcKc)Vc\operatorname{tril}(Q_cK_c^\top)V_c 都能就地完成,状态从头到尾只活在寄存器里。

官方 chunk_delta_h 却在循环开头写了 T.copy(b_h_shared, h[...]),把每个 chunk 的状态快照全部落盘。h 的形状是 (B,S/blockS,H,DK,DV)(B, S/block_S, H, DK, DV),按官方 main() 的配置(B=1,S=32768,H=32,DK=DV=128B{=}1, S{=}32768, H{=}32, DK{=}DV{=}128,chunk 64,bf16)达 512 MiB–与 KKVV 输入之和等量。

方案 KVK^\top V 次数 HBM 额外开销 适用
§4 序列进 grid NC(NC1)/2NC(NC-1)/2 隔离验证恒等式
§6.4.2 融合单 kernel NCNC 无门控/无删除的前向、推理
官方两 kernel(chunk_delta_h + chunk_o NCNC hh 快照(示例 512 MiB) 训练(反向要 hh)、DeltaNet 的 UT 变换

拆开的三个真实理由:反向传播需要每个 chunk 的 SprevS^{\text{prev}},重算不如存;DeltaNet 的 WW(I+tril(diag(β)KK))1(I + \operatorname{tril}(\operatorname{diag}(\beta)KK^\top))^{-1} 需要沿 DKDK 收缩的独立阶段,塞不进这个循环;chunk_o 有自己的 tile 划分自由度–读一份现成的快照就能对所有 chunk 完全并行,不再受递推顺序约束。顺序性和并行性被分到两个 kernel 里,各自取所需;融合省 HBM,拆开换灵活性和反向所需的中间量。 第一级用不到后两者,所以融合是更好的起点。

还有两处细节值得对照本文的实现:

  • clear_accum=True 的用途T.gemm(W_shared, b_h_shared, V_new_fragment, clear_accum=True) 显式覆盖而非累加,因为 V_new_fragment 每轮都要重算,不能沿用上一 chunk 的残留。本文 §4.1 靠循环外的 T.clear(S_f) 达到同样目的,两种写法都行;进入流水线循环后就必须用 clear_accum
  • 门控在 log 域相减,从不物化比值。官方实现写作 T.exp2((G_last_local - G_fragment[i_s2, i_v]) * 1.442695),其中 GG 已是 logsigmoid 后的累积和,1.442695=log2e1.442695 = \log_2 eexe^x 转成硬件 exp2GG 单调递减保证 GlastGi0G_{last} - G_i \le 0指数结果恒不大于 1,不存在溢出。本文的第一级没有门控,这里仅作为对照记录。

一句话总结:线性注意力的全部优势是把 O(N2D)O(N^2 D) 换成 O(ND2)O(N D^2),而这个优势能否落地,取决于序列轴放不放进 grid。放进去,block 间无法传递状态,只能各自重算,(NC1)/2(NC-1)/2 的冗余把 NN 的次数顶回 2;不放进去,递推退回单 block 内的顺序循环、用寄存器承载状态,并行度改由 batch/head/DV 提供,代价是把状态快照写回 HBM。本文选前者换取实现的可隔离性,生产实现选后者。


7. 总结

  1. 本级移除了 KDA 的全部衰减因子与删除因子,只保留 SS+KVS \leftarrow S + K^\top V,目的是将 chunkwise 分解恒等式 Oc=QcScprev+tril(QcKc)VcO_c = Q_c S_c^{\text{prev}} + \operatorname{tril}(Q_c K_c^\top) V_c 单独隔离验证。该恒等式是后续各级的公共基础,每级只在它的两项上插入衰减权重,GEMM 的形状与调用次序不变。
  2. 四层 PyTorch 参考构成从数学到 kernel 的完整链条:A 逐 token 递归验证递推式定义,B 分块向量化验证恒等式,C 逐块独立重算镜像 §4 kernel 控制流,D(§6.4.2)按 (b,h,dv)(b,h,dv) 切 tile + 序列内循环镜像生产型 kernel。kernel 与参考的差异仅剩两项–f32→f16 显式降精度与 shared buffer 中转,这也是全部数值差异的来源。
  3. 与 FA 的差异集中在三点:无 online softmax(掩码置 0 而非 -\infty)、循环携带 dk×dvd_k \times d_v 状态、无按行归约(T.gemm policy 用默认 Square 即可)。
  4. 线性注意力的优势与本文实现的退化:优势来自结合律换边,(QK)V(QK^\top)VO(N2D)O(N^2 D) 变成 Q(KV)Q(K^\top V)O(ND2)O(N D^2),理想实现与因果 FA 的比值按 2(D+BC)/N2(D+BC)/N 衰减。但本文的逐块独立重算使冗余系数 (NC1)/2(NC-1)/2NN 线性增长,把 NN 的次数顶回 2–实测比值收敛到常数 D/(2BC)=0.5D/(2BC) = 0.5比值与 NN 无关正是失败的判据,不是中性描述。
  5. 退化的根源是 grid 划分,不是算法:把序列轴放进 grid.x 后,block 间既无执行次序保证也无同步原语,前缀状态只能各自重算。换成 grid =(DV/blockDV, BH)= (DV/block_{DV},\ B \cdot H)序列轴不进 grid,递推退回单 block 内的 T.Pipelined 顺序循环,状态作为 loop-carried fragment 常驻寄存器,每 chunk 只做一次 KVK^\top V。并行度改由 batch/head/DV 提供(示例配置 128 个 block)。N=65536N = 65536 时两者相差 511.5 倍。§6.4.2 给出了完整可运行的融合版本。
  6. 只切 DV、不切 DK 的真正理由不是「KVK^\top V 的 M 维」KVK^\top V 的收缩维是 chunk 长度 CC,DK 和 DV 在这一步都是自由维,切哪边都不需归约。不对称性来自下游:DKDKQcSprevQ_c S^{\text{prev}}QcKcQ_c K_c^\top 两个 GEMM 的收缩维,DVDV 在三个 GEMM 里始终是自由维。切 DK 会把读出变成 split-K(每 chunk 归约一次 [C,DV][C, DV]),并且块内项根本算不出完整的 [C,C][C,C] 转移矩阵;切 DV 只付出Q/K 的读放大,多读而不归约。实测三个 blockDV\text{block}_{DV}(32/64/128)相对 L2 完全相同(5.60×10165.60 \times 10^{-16},见 §3.4)。
  7. 循环体内的顺序就是右移语义:状态更新 T.gemm(K_s, V_s, S_f, transpose_A=True) 必须排在输出写回之后,提前就变成 inclusive 前缀、块内项被重复计入,实测相对 L2 从 101610^{-16} 跳到 6.98×1016.98 \times 10^{-1}(见 §3.4)。还有一条配套规则:T.clear 只能在流水线循环外给跨迭代的状态用,循环内那些每轮重算的累加器必须靠 clear_accum=True
  8. 融合还是拆两个 kernel:本级无门控无删除,单 block 持有完整 DKDK,状态可以全程待在寄存器里,不需要写 hh 快照。官方 chunk_delta_h + chunk_o 拆开的理由是反向传播需要 SprevS^{\text{prev}}、DeltaNet 的 UT 变换需要沿 DKDK 收缩的独立阶段、以及 chunk_o 想要自己的 tile 划分自由度,代价是示例配置下 512 MiB 的 HBM 写回。
  9. 实测数字(均 numpy fp64,见 §3.4):四层参考互验的相对 L2 误差均在 1.65.6×10161.6\text{--}5.6 \times 10^{-16},恒等式无误;刻意去掉右移的错误实现相对 L2 为 7.78×1017.78 \times 10^{-1}(参考 B)与 6.98×1016.98 \times 10^{-1}(参考 D),验证流程对结构性错误敏感。TileLang kernel 本身需 CUDA 设备,实测待补。
  10. 两个实现陷阱:einsum 输出下标不可重复(->bhcdd 会抛异常,key 维与 value 维必须用不同字母);T.Pipelined 之后 shared buffer 的残留内容取决于流水线展开方式,复用前必须显式重载。

参考

本文的分块恒等式验证与 FLOPs 账本在 numpy fp64 上复现;GPU 实测数据待补。

TileLang 实战:FlashAttention 前向 Kernel

FlashAttention 没有发明新的数学。它是 online softmax 递推与两个 GEMM 在同一个 tile 循环里的交织–分数块 SS 算出来就地做 softmax 得到 P~\tilde{P}P~\tilde{P} 就地乘上 VV,中间矩阵全程驻留在 SRAM,HBM 流量从 O(N2)O(N^2) 降到 O(Nd)O(N \cdot d)。本文用 TileLang(v0.1.13)把 FA-2 论文的 Algorithm 1 写成可运行的 kernel:先补齐数学地基,再讲清楚"怎么切",然后逐行拆解主循环,最后做数值验证、测速与调参。只覆盖前向;TileLang 五原语与 T.Pipelined 流水线机制见上一篇《TileLang 编程基本知识点》,本文直接使用其结论。


一、为什么标准 Attention 慢:IO 才是瓶颈

标准 scaled dot-product attention 的流程是 S=QK/dP=softmax(S)O=PVS = QK^\top / \sqrt{d} \to P = \mathrm{softmax}(S) \to O = PV。计算本身没有问题,问题在中间矩阵:SSPP 都是 N×NN \times N,必须写进显存再读回来。N=8192N = 8192 时单个 fp32 矩阵就是 256 MB,长序列下这个开销直接失控。

GPU 内存层级(A100 量级):

层级 容量 带宽 延迟量级
HBM(显存) 40~80 GB ~1.5-2 TB/s ~400-800 cycle
SRAM(每 SM) 164 KB ~19 TB/s ~20-30 cycle
寄存器(每 SM) 256 KB ~19 TB/s 接近 0

SRAM 比 HBM 快 10 倍以上,但只有一百多 KB。attention 的 FLOPs 是 O(N2d)O(N^2 d),标准实现与 FlashAttention 完全相同–快慢差别全部来自中间矩阵走没走 HBM

指标 标准 Attention FlashAttention
中间矩阵内存 O(N2)O(N^2),S/P 落显存 O(N)O(N),S/P̃ 留在 SRAM
HBM 往返 3 次(写 S → 读写 P → 读 P 做 PV) 1 次(读 Q/K/V → 写 O)
FLOPs O(N2d)O(N^2 d) 相同
精确性 精确 精确(非近似)

这就是论文标题里 “IO-aware” 的含义:不省计算,省搬运。

一句话总结:attention 是访存受限的算子,FlashAttention 的全部收益来自让 N×NN \times N 的中间矩阵不落 HBM。


二、数学地基:softmax 的三级台阶

2.1 平移不变性 → 安全 softmax

softmax 对输入加任意常数 cc 不变(分子分母的 ece^c 约掉)。这个自由度是数值安全的救命稻草:fp16 上限只有 65504,s=100s = 100e1001043e^{100} \approx 10^{43} 直接 inf,inf/inf 变 NaN。减去行最大值 mm 后指数上限恰为 0:

softmax(xi)=eximjexjm,m=maxjxj\mathrm{softmax}(x_i) = \frac{e^{x_i - m}}{\sum_j e^{x_j - m}}, \quad m = \max_j x_j

代价是计算顺序被强制:先扫一遍求 mm,再扫一遍求分母与输出–两遍扫描。

2.2 分块场景下,"两遍"成为灾难

SRAM 装不下整行 K/V,只能按块流入。设处理完第 1 块时按基准 m=3m=3 攒好部分和;第 2 块冒出更大分数,基准换成 m=7m=7–旧 exp 全部基准错误。三条路:

方案 做法 代价
存下全部 S 等全局 max 再统一归一化 N×NN \times N 落显存,爆
两遍重扫 pass1 求 mm\ell;pass2 重算加权 K/V 从 HBM 读两次,带宽 ×2
online softmax 边收块边维护"以当前 max 为基准"的部分和 一遍过,数学无损

2.3 online softmax 递推

每个 query 行维护三个状态,初值 m=m = -\infty=0\ell = 0O=0O = 0。新块 SnewS_{new} 到来时:

mnew=max(mold, rowmax(Snew))α=emoldmnew旧成果的打折系数new=αold+rowsum(eSnewmnew)Onew=αOold+eSnewmnewVnew\begin{aligned} m_{new} &= \max(m_{old},\ \mathrm{rowmax}(S_{new})) \\ \alpha &= e^{m_{old} - m_{new}} && \leftarrow \text{旧成果的打折系数} \\ \ell_{new} &= \alpha \cdot \ell_{old} + \mathrm{rowsum}(e^{S_{new} - m_{new}}) \\ O_{new} &= \alpha \cdot O_{old} + e^{S_{new} - m_{new}} \cdot V_{new} \end{aligned}

循环结束后做全程唯一一次归一化:Output=O/\mathrm{Output} = O / \ell

为什么无损:缩放因子连乘时指数项逐项相消,em1m2em2m3=em1m3e^{m_1 - m_2} \cdot e^{m_2 - m_3} = e^{m_1 - m_3},循环结束时每项都被精确换算到最终 max 的基准下。中间任何时刻 OO 都是合法但未归一化的加权和,数值永远有界(当前最大项的 exp 恰为 1)。

数值走一遍(一行 query,三块各来一个分数 1、2、3,正确答案 = softmax(1, 2, 3) 的权重):

1
2
3
4
5
6
7
8
init : m=−∞   ℓ=0      O=0
j=1 : m=1, α=e^(−∞)=0
O = 0×0 + e^0·v₁ = 1.000·v₁ ℓ = 1.000
j=2 : m=2, α=e^(1−2)=0.368
O = 0.368·v₁ + 1·v₂ ℓ = 1×0.368 + 1 = 1.368
j=3 : m=3, α=e^(2−3)=0.368
O = 0.135·v₁ + 0.368·v₂ + 1·v₃ ℓ = 1.503
收尾 : O/ℓ = (0.090, 0.245, 0.665)·v ✓

盯住 v1v_1 的系数:1.000 → 0.368 → 0.135,每来一个更大的 max,历史成果整体打折一次。这就是代码里 acc_o *= scores_scale 那一行的全部含义。

2.4 FlashAttention = online softmax × 两个 GEMM

online softmax(2018,Milakov & Gimelshein)只解决流式归一化,没碰矩阵乘。FlashAttention(2022)的洞察是:attention 恰好是 matmul → softmax → matmul 的三明治,三个部件可以共用同一个 tile 循环:

1
2
3
4
5
for j in 1..Tc:                                  # K/V 方向遍历
S_j = Q_tile @ K_j.T # gemm#1:喂料
m, alpha, P̃_j, ℓ = online_softmax_update(S_j) # 三件套
O_tile = alpha * O_tile + P̃_j @ V_j # gemm#2:消费
Output = O_tile / ℓ # 唯一一次归一化

SSP~\tilde{P} 的生成和消费全部在 SRAM/寄存器内完成,不写入 HBM。

一句话总结:先有 online softmax 这条递推式,才有 FlashAttention 这个 kernel;数学在前,工程在后。


三、切法:block_M 进 grid,block_N 进循环

3.1 都沿 seq 轴切:Q/O 进 grid,K/V 进循环

Q/O 与 K/V 都沿着 seq 轴切块,块大小分别由 block_M 与 block_N 控制,但两者的去向不同:

1
2
3
4
5
6
seq ──->
├── Q₁ ├── Q₂ ├── Q₃ ┤ 按 block_M 切,进 grid:每个 CUDA block 认领一块 Q,
├── O₁ ├── O₂ ├── O₃ ┤ 独立算出对应的 O,块与块之间零依赖
├── K₁ ├── K₂ ├── K₃ ┤ 按 block_N 切,进内层循环:每个 block 沿 K/V 块逐块扫描;
├── V₁ ├── V₂ ├── V₃ ┤ K 与 V 必须按同样的边界切块(P̃ 的列与 V_j 的行须是同一批 key)
└────────────────────┘

为什么这样分?softmax 对分数矩阵做逐行归一化,O 的第 ii 行只由 Q 的第 ii 行决定,与 Q 的其他行无关。因此按行把 Q/O 切成 block_M 大小的块、分给不同的 CUDA block 并行计算,块间不需要任何通信。K/V 则不同:任何一行的 softmax 分母都要对整条 seq 求和,每一块 Q 都需要全部的 K/V。K/V 无法划归某个 block 独占,只能作为内层循环的遍历方向,按 block_N 逐块加载、逐块累积。

对比普通 GEMM 就能看出差别:C=ABC = AB 的输出块 CijC_{ij} 只依赖 AA 的第 ii 个行块与 BB 的第 jj 个列块,M、N 两个方向都可以切块并行;attention 的输出在 seq 方向上对 K/V 有全局依赖(softmax 分母是全行求和),这个方向只能串行遍历。block_M 决定并行度(grid 大小),block_N 决定每个 block 的循环长度

3.2 grid 与布局

张量布局用 BSHD([batch, seq_len, heads, dim]),batch/head 在外层、seq 连续,tile 拷贝才能访存合并。grid 三维:

1
with T.Kernel(T.ceildiv(seq_len, block_M), heads, batch, threads=threads) as (bx, by, bz):

每个 CUDA block 认领一个 Q tile(bx)、一个 head(by)、一个 batch(bz),沿 K/V 方向循环。这正是 FA-2 的循环序:Q 装载一次驻留 SRAM 全程不动。FA-1 是反过来的(外层 K/V、内层 Q),每轮 K/V 块的结果要经 HBM 中转更新各 Q 块,多出大量中间读写–FA-2 把循环反过来之后这条路才彻底堵死。

3.3 块大小的约束

SMEM 需求近似为:

QiBr×d+KjBc×d+VjBc×d+SijBr×BcMSRAM\underbrace{Q_i}_{B_r \times d} + \underbrace{K_j}_{B_c \times d} + \underbrace{V_j}_{B_c \times d} + \underbrace{S_{ij}}_{B_r \times B_c} \le M_{\text{SRAM}}

通常取 Br=BcMSRAM/dB_r = B_c \approx \sqrt{M_{\text{SRAM}} / d}。A100 上 M100M \approx 100 KB、d=128d = 128B128B \approx 128;再考虑 fragment 状态(acc_s、acc_o 等)占的寄存器,128×128 是常见起点,seq 很长或 head_dim 很大时倾向减小 block_N。

3.4 缓冲区清单

缓冲区 位置 形状 角色
Q_shared / K_shared / V_shared SMEM [block_M/N, dim] GEMM 的 A/B 操作数走 SMEM 路径
O_shared SMEM [block_M, dim] 写回前的中转(fragment 布局重排)
acc_s fragment(fp32) [block_M, block_N] S = QK^T 分块累加器
acc_s_cast fragment(fp16) [block_M, block_N] softmax 后的 P̃,cast 给第二个 MMA
acc_o fragment(fp32) [block_M, dim] 输出累加器,全程不归一化
scores_max / _prev / _scale / _sum / ell fragment(fp32) [block_M] online softmax 的逐行状态

一句话总结:所有逐行状态都放 fragment(寄存器),只有 GEMM 操作数和最终输出走 SMEM–这是三级内存抽象在 FA 里的标准分工。


四、主循环逐行拆解

完整代码基于官方 examples/flash_attention/example_mha_fwd_bshd.py(注释为本文所加)。这是仅推理版:不产出 backward 需要的 logsumexp,分母变量直接叫 ell

4.1 三层结构与初始化

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
@autotune(configs=get_configs(), warmup=10, rep=10)
@tilelang.jit(
out_idx=[3], # 第 4 个形参 Output 是输出
pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True},
)
def flashattn(batch, heads, seq_len, dim, is_causal,
block_M=64, block_N=64, num_stages=1, threads=128):
# softmax 的 1/sqrt(dim),预先乘上 log2(e),后面全部用硬件更快的 exp2
scale = (1.0 / dim) ** 0.5 * 1.44269504

@T.prim_func
def main(Q: T.Tensor(shape, dtype), K: T.Tensor(shape, dtype),
V: T.Tensor(shape, dtype), Output: T.Tensor(shape, dtype)):
with T.Kernel(T.ceildiv(seq_len, block_M), heads, batch,
threads=threads) as (bx, by, bz):
... # 缓冲区分配见 3.4

# Q tile 装载一次,全程驻留
T.copy(Q[bz, bx*block_M:(bx+1)*block_M, by, :], Q_shared)
T.fill(acc_o, 0)
T.fill(ell, 0)
T.fill(scores_max, -T.infinity(accum_dtype))

# causal 截断:Q 块最远看到 (bx+1)*block_M - 1,右侧整块跳过
loop_range = (
T.min(T.ceildiv(seq_len, block_N),
T.ceildiv((bx+1)*block_M, block_N))
if is_causal else T.ceildiv(seq_len, block_N)
)

三层结构是 TileLang 的标准混用写法:外层 @autotune 管配置搜索,中间 @tilelang.jit(out_idx=[3]) 声明输出形参,内部 @T.prim_func 管精确签名。

causal 截断值得推一遍:Q 块 bx 的行最远到 (bx+1)blockM1(bx+1) \cdot block_M - 1,key 位置超过它的全被掩码,对应的 K/V 块整块不进循环。N=4096N = 4096、块 128 时对角线右侧一半块直接消失,省一半计算。注意非 causal 时这个 T.min 退化为总块数,掩码换成右边界越界判断(见下节)。

4.2 掩码写进累加器:顺序是精髓

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
for k in T.Pipelined(loop_range, num_stages=num_stages):
T.copy(K[bz, k*block_N:(k+1)*block_N, by, :], K_shared)

# ① 先写掩码,再做 GEMM
if is_causal:
for i, j in T.Parallel(block_M, block_N):
acc_s[i, j] = T.if_then_else(
bx*block_M + i >= k*block_N + j, 0, -T.infinity(acc_s.dtype))
else:
for i, j in T.Parallel(block_M, block_N):
acc_s[i, j] = T.if_then_else(
k*block_N + j >= seq_len, -T.infinity(acc_s.dtype), 0)

# ② S = Q @ K^T(tensor core)
T.gemm(Q_shared, K_shared, acc_s, transpose_B=True,
policy=T.GemmWarpPolicy.FullRow)

掩码写在 GEMM 之前能成立,靠的是 T.gemm 的累加语义(C+=ABC \mathrel{+}= A B,所以纯 GEMM 例子里要先 T.clear):$-\infty + $ 有限值 == -\infty,被掩码的位置在 GEMM 里"存活"下来,之后 exp2(m)=0\mathrm{exp2}(-\infty - m) = 0,对分数和、对输出零贡献。省掉一遍 GEMM 后的掩码 pass。

两个分支的条件不同,处理的边界不同:

  • causal:query 全局位置 ≥ key 全局位置才保留(只许看过去);
  • 非 causal:k*block_N + j >= seq_len 置 −inf,处理的是 seq 不整除 block_N 时最后一块里的"假 key"。

4.3 online softmax 七步

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
# ③ online softmax 更新
T.copy(scores_max, scores_max_prev) # 1. 存历史 max
T.fill(scores_max, -T.infinity(accum_dtype))
T.reduce_max(acc_s, scores_max, dim=1, clear=False) # 2. 本块行 max
for i in T.Parallel(block_M):
scores_max[i] = T.max(scores_max[i], scores_max_prev[i]) # 3. 合并成 m_new
for i in T.Parallel(block_M):
scores_scale[i] = T.exp2( # 4. α = exp2(m_old − m_new)
scores_max_prev[i]*scale - scores_max[i]*scale)
for i, j in T.Parallel(block_M, block_N):
acc_s[i, j] = T.exp2( # 5. P̃ = exp2(S·scale − m·scale)
acc_s[i, j]*scale - scores_max[i]*scale)
T.reduce_sum(acc_s, scores_sum, dim=1) # 6. 本块 rowsum(P̃)
for i in T.Parallel(block_M):
ell[i] = ell[i]*scores_scale[i] + scores_sum[i] # 7. ℓ ← α·ℓ + 本块部分和
T.copy(acc_s, acc_s_cast) # fp32 → fp16,喂第二个 MMA

对应 2.3 节的递推式,逐步可查。两个容易踩的坑:

  • max 与 scale 的次序:先由历史 max 和本块 max 定出 mnewm_{new}由这一对 max 算出 α\alpha。max 是 scale 的来源,不是被更新对象;
  • scores_sumell 别混:前者只是本轮 rowsum(P̃) 的临时量,每次迭代被覆盖;ell 是带折扣链跨轮累积的总分母。归一化除的是 ell,除成某一块的局部和就错了。

4.4 第二个 GEMM 与收尾

1
2
3
4
5
6
7
8
9
10
11
    # ④ 先打折旧成果,再加新项
for i, j in T.Parallel(block_M, dim):
acc_o[i, j] *= scores_scale[i]
T.copy(V[bz, k*block_N:(k+1)*block_N, by, :], V_shared)
T.gemm(acc_s_cast, V_shared, acc_o, policy=T.GemmWarpPolicy.FullRow)

# ⑤ 循环外:唯一一次归一化,然后写回
for i, j in T.Parallel(block_M, dim):
acc_o[i, j] /= ell[i]
T.copy(acc_o, O_shared)
T.copy(O_shared, Output[bz, bx*block_M:(bx+1)*block_M, by, :])

acc_o *= scores_scale 必须在第二个 GEMM 之前:先加新项再打折,新项就被错误地折旧了。归一化只在循环外做一次,中途除会错–分母没定型。

写回走 fragment → O_shared → Output 两跳:fragment 里的数据按 MMA 布局散在各线程寄存器,先落 SMEM 重排布局,再合并写 global。

4.5 工程细节清单

  • exp2 技巧:把 1/d1/\sqrt{d} 预乘 log2e1.4427\log_2 e \approx 1.4427 折进 scale,全部指数运算用 exp2 硬件指令(比 exp 快),FA 实现的标配;
  • P̃ 的 fp16 cast 是全 kernel 唯一降精度点:MMA 只吃低精度操作数,acc_s/acc_o/ell 等累加器与状态全程 fp32;
  • FullRow warp policyT.gemm 把输出 tile 切给 block 内各 warp 的切法。FA 的 gemm#1 后紧跟按行归约(reduce_max/reduce_sum)和按行缩放(acc_o *= scores_scale),FullRow 保证一行完整住在一个 warp 内,这些操作全部退化为 warp 内 shuffle;若用默认 Square,第 i 行的左右两半住在两个 warp 里,归约就得经 SMEM 中转外加 bar.sync–结果仍正确,但每轮迭代多两次往返:
1
2
3
4
5
6
7
8
Square(默认):每 warp 近似方块    FullRow:按 M 横切,每 warp 拿全宽整段行
┌─────────┬─────────┐ ┌──────────────────────┐
│ warp0 │ warp1 │ │ warp0 行 0..31 │
│ 64×64 │ 64×64 │ ├──────────────────────┤
├─────────┼─────────┤ │ warp1 行 32..63 │
│ warp2 │ warp3 │ ├──────────────────────┤
│ 64×64 │ 64×64 │ │ warp2/3 行 64..127│
└─────────┴─────────┘ └──────────────────────┘
  • forward 用 num_stages=1:循环体是两个相互依赖的 GEMM 夹一串 elementwise(S 等 K、P̃ 等 S 和 m),预取重叠窗口小而寄存器压力大,官方权衡后的选择。流水线什么时候真的有用,见上一篇的三因素模型与对照实验;
  • O 不中途归一化是 FA-2 的关键改动之一(见下节)。

一句话总结:主循环五行骨架–搬 K、写掩码、GEMM 喂料、online softmax 三件套、rescale 后 GEMM 消费;每一步都有"顺序不能错"的理由。


五、与 FA-2 论文 Algorithm 1 的对照

这份代码就是 FlashAttention-2 论文 Algorithm 1 的逐行实现。伪代码符号 ↔ 代码变量:

伪代码 内容 代码对应
Br×BcB_r \times B_c 分块大小 block_M × block_N
for i ≤ Tᵣ(外层 Q 块) 外层循环进 grid with T.Kernel(...) as (bx, by, bz)
QiQ_i 进 SRAM 装载后驻留 T.copy(Q[...], Q_shared)
O0=0, 0=0, m0=O^0 = 0,\ \ell^0 = 0,\ m^0 = -\infty 初始化三状态 T.fill(acc_o/ell, 0)T.fill(scores_max, -inf)
for j ≤ T_c(内层 K/V 块) 内层循环 for k in T.Pipelined(loop_range, ...)
Sij=QiKjS_{ij} = Q_i K_j^\top 分数块 T.gemm(Q_shared, K_shared, acc_s, transpose_B=True)
mm 更新、P~\tilde{P}\ell 更新 online softmax 一族 reduce_maxscores_scaleexp2ell 递推
Odiag(eΔm)1O+P~VjO \leftarrow \mathrm{diag}(e^{\Delta m})^{-1} O + \tilde{P} V_j 先打折再加新项 acc_o *= scores_scalegemm(acc_s_cast, V_shared, acc_o)
O/O/\ell(第 12-13 行) 循环外归一化 acc_o /= ell

四处伪代码没写、真实 kernel 必须处理的差异:

  1. 底数:论文 e 底,实现全用 exp2 硬件指令,scale 预乘 log2e\log_2 e
  2. scale 折叠1/d1/\sqrt{d} 不显式乘,折进指数表达式省一遍逐元素乘;
  3. acc_s_cast:P̃ 从 fp32 fragment cast 成 fp16 才能进 MMA;
  4. causal 截断:算法按稠密写,代码用 loop_range 让对角线右侧整块跳过。

顺带把 FA-1 → FA-2 的三个改动列清楚,这份代码全部站在 FA-2 一边:

FA-1(2022) FA-2(2023)
外层循环 K/V 块 Q 块
中间读写 各 Q 块的结果需经 HBM 中转更新 Q 驻留 SRAM,块间零 HBM 往返
O 的归一化 逐块保持已归一化(多两次乘除) 循环外一次性 /= ell
并行度 受 K/V 块数限制 Q 块 × head × batch,并行度更高
warp 分工 均分 K/V 切给不同 warp,减少同步(TileLang 里由 GemmWarpPolicy 表达)

六、验证、测速与调参

写完 kernel 的三个标准动作(数值验证 -> dump 生成码 -> 基准测试),在 FA 上一个不少:

① 数值验证profiler.assert_allclose 直接接受 PyTorch 参考实现做正确性比对,不需要自己写验证框架:

1
2
3
4
5
kernel = flashattn(batch, heads, seq_len, dim, is_causal,
block_M=128, block_N=128, num_stages=1, threads=128)
ref = partial(ref_program, is_causal=is_causal) # einsum 朴素注意力 + tril 掩码
profiler = kernel.get_profiler()
profiler.assert_allclose(ref, rtol=0.01, atol=0.01)

容差 0.01 是 fp16 输出的合理范围;更严的判定方法(与 fp32 参考比、大误差元素计数)见上一篇第九节。

② dump 生成码kernel.get_kernel_source() 打印完整 CUDA 源码。对 FA 值得确认三件事:T.copy 降成了 cp.async 还是 TMA、FullRow 下 warp 怎么切行、exp2 是否真的成了硬件指令。

③ 测速profiler.do_bench(warmup=500) 自动 warmup 多次取统计。TFlops 按 2BHN2d×22 \cdot B \cdot H \cdot N^2 \cdot d \times 2 个 matmul 计算,causal 乘 0.5:

1
2
3
4
flops = 2.0 * batch * heads * seq_len * seq_len * dim   # 单个 matmul
total_flops = 2 * flops * (0.5 if is_causal else 1.0)
latency = profiler.do_bench(warmup=500)
print(f"{total_flops / latency * 1e-9:.2f} TFlops")

调参交给 @autotune:把 block_M/block_N/num_stages/threads 写成带默认值的参数,搜索空间在 get_configs() 里列全(笛卡尔积),机器自己选。注意首跑很慢(每组都要编译 + 测速),调通阶段先裁成单组。


七、总结

  1. 数学在前:softmax 平移不变性 → 安全 softmax → 分块下的换基准问题 → online softmax 递推。FlashAttention 只是这条递推式与两个 GEMM 的交织,精确、非近似;
  2. 切法是结构:Q/O 与 K/V 都沿 seq 轴切块,block_M 控制的 Q 块进 grid(决定并行度),block_N 控制的 K/V 块进内层循环(决定遍历长度);softmax 的逐行归约使 seq 方向对 K/V 产生全局依赖,只能作为循环方向,这是 FA 与 GEMM 的本质结构差异;
  3. 实现靠 TileLang 五原语T.Kernel 定切法、T.copy 管搬运、T.gemm 喂两个 matmul、T.Pipelined 管 K/V 流水(FA 前向官方权衡后用 stages=1)、T.Parallel + reduce_* 承载 online softmax 三件套。掩码写进累加器、exp2 折叠、FullRow 行所有权、O 延迟归一化,四个细节决定了这份代码"像论文"还是"只是能跑"。

可迁移的启示:读 FA 代码的正确顺序是先读递推式再读循环–所有 tile 级 kernel 都是"一条数学递推 + 一个 tile 循环",TileLang 只是把后者写到了 30 行的量级。


参考

  • FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(Dao et al., NeurIPS 2022, arXiv:2205.14135)
  • FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning(Dao, 2023, arXiv:2307.08691)
  • Online normalizer calculation for softmax(Milakov & Gimelshein, 2018, arXiv:1805.02867)
  • TileLang GitHub(v0.1.13):examples/flash_attention/example_mha_fwd_bshd.py 为本文代码底本
  • 本站:《TileLang 编程基本知识点

本文基于 TileLang v0.1.13 官方示例与 FlashAttention-2 论文整理,代码注释为本文所加。

TileLang 的定位可以用一句话概括:把 CUDA kernel 里「怎么切 tile、数据放哪级存储、什么时候预取」这几件事变成显式的 Python 语句,其余的寄存器分配、指令选择、同步插入交给编译器。它不是又一个 Triton–Triton 隐藏 shared memory,TileLang 让你直接写 T.alloc_shared。这个差别决定了两者能碰到的性能天花板不同。

本文整理 TileLang 的编程基本知识点:它是什么、编程模型的 5 个原语、循环原语与 T.Pipelined 的工作原理、同一份代码在不同 GPU 架构上如何 lowering、四个进阶主题(Split-K / warp specialization / autotune / Blackwell 两条路径),以及写完 kernel 之后的三个标准动作(验证 / dump 源码 / bench)。


一、TileLang 是什么

项目 现状(2026-08)
版本 v0.1.13(2026-08-03 发布)
GitHub tile-ai/tilelang,7.2k stars
底层 TVM(IR 已迁移到 TIRX)
Python ≥ 3.10
主力后端 CUDA(SM70~SM120)
其他后端 ROCm/HIP、Apple Metal、LLVM CPU(实验)、CuTe DSL(实验)、WebGPU(实验)
生态后端 华为 Ascend、沐曦 MACA、摩尔线程 MUSA(独立仓库维护)

出身是学术项目:主要由 LeiWang1999、chengyupku、nox-410 在北大杨智教授指导下开发,部分工作在 MSRA 实习期间完成。2025-01 开源。

值得注意的是上游模型厂在用它:TileLang 仓库的 examples/ 里有 deepseek_mladeepseek_v32deepseek_v4deepseek_mhc 四个目录,DeepSeek 系列的 MLA / 稀疏注意力 / mHC 融合 kernel 都有 TileLang 参考实现。这意味着读 TileLang examples 等于读一份最新算子的可执行论文附录–这是它相比 Triton 的一个实际优势。

一句话总结:TileLang = Pythonic 语法 + 显式 tile/memory 层级控制 + TVM 编译基础设施,目标是「写起来像 Triton,控制力接近 CUTLASS」。


二、编程模型:5 个原语撑起全部

TileLang 的 API 面很窄,这是刻意的。一个完整 kernel 基本只用这 5 类原语:

1
2
3
4
5
6
7
8
T.Kernel(grid_x, grid_y, threads=N)     ← ① 定义 grid / block,拿到 block index

├── T.alloc_shared(shape, dtype) ← ② 显式声明 shared memory buffer
├── T.alloc_fragment(shape, dtype) ← ②' 显式声明 register fragment(累加器)

└── for k in T.Pipelined(N, num_stages=3): ← ③ 软件流水(自动双/三缓冲)
T.copy(global_tile, shared) ← ④ 数据搬运(按架构 lower 到 cp.async / TMA)
T.gemm(A_s, B_s, C_frag) ← ⑤ tile 级 MMA(映射到 Tensor Core)

完整的 FP16 GEMM + ReLU(基于官方 quickstart):

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
import torch
import tilelang
import tilelang.language as T


@tilelang.jit
def matmul(A, B, block_M: int, block_N: int, block_K: int):
M, N, K = T.const("M, N, K")
dtype = T.float16
accum_dtype = T.float32
A: T.Tensor((M, K), dtype)
B: T.Tensor((K, N), dtype)
C = T.empty((M, N), dtype)

# grid: (N 方向块数, M 方向块数),每 block 128 线程
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
A_shared = T.alloc_shared((block_M, block_K), dtype) # shared tile
B_shared = T.alloc_shared((block_K, block_N), dtype)
C_local = T.alloc_fragment((block_M, block_N), accum_dtype) # 寄存器累加器,fp32

T.clear(C_local)

# 三级软件流水:搬 k+1 块的同时算 k 块
for ko in T.Pipelined(T.ceildiv(K, block_K), num_stages=3):
T.copy(A[by * block_M, ko * block_K], A_shared) # global -> shared
T.copy(B[ko * block_K, bx * block_N], B_shared)
T.gemm(A_shared, B_shared, C_local) # tensor core,fp32 累加

# epilogue:ReLU
for i, j in T.Parallel(block_M, block_N):
C_local[i, j] = T.max(C_local[i, j], 0)

T.copy(C_local, C[by * block_M, bx * block_N]) # fragment -> global

return C


M = N = K = 1024
# 用静态 shape 编译出可执行 kernel
matmul_kernel = matmul.compile(M=M, N=N, K=K, block_M=128, block_N=128, block_K=32)

a = torch.randn(M, K, device="cuda", dtype=torch.float16)
b = torch.randn(K, N, device="cuda", dtype=torch.float16)
c = matmul_kernel(a, b)

30 行写完一个带 fused epilogue 的 Tensor Core GEMM。几个关键点:

  1. T.const("M, N, K") 声明符号化 shape.compile(...) 按实际 shape 特化出可执行 kernel(也可以直接调用让 @tilelang.jit 在首次调用时惰性编译)。
  2. 累加器 dtype 和存储 dtype 分离C_local 是 fp32 fragment,写回时才降到 fp16。这是数值稳定的标准配方。
  3. epilogue 用 T.Parallel 表达,不需要另起 kernel–省一次 HBM 往返。
  4. 没有一行同步代码__syncthreads() 由编译器根据 T.Pipelined 的依赖关系自动插入。

一句话总结:TileLang 的 API 面窄到只有 5 类原语,但这 5 类恰好覆盖了 GPU kernel 性能的全部决定因素–tile 怎么切、数据放哪级存储、什么时候预取。


三、循环原语:从 T.serial 到 T.Pipelined

GEMM 例子里已经出现过两种循环(T.PipelinedT.Parallel)。TileLang 的循环构造一共就这四个,按暴露给编译器的并行度递进,先顺次过一遍:

原语 语义 典型用途
T.serial 普通 for 循环,迭代之间有依赖 递推、边界处理
T.unroll 要求编译器完全展开 小循环,省掉分支和循环开销
T.Parallel 嵌套并行循环,所有迭代互相独立 elementwise、epilogue
T.Pipelined 软件流水,生产者-消费者跨迭代重叠 GEMM / attention 主循环
1
2
3
4
5
6
7
8
9
10
11
for i in T.serial(N):                        # 串行:下一轮依赖上一轮的结果
...

for k in T.unroll(K_TILE): # 展开:编译期摊平
acc += a[k] * b[k]

for i, j in T.Parallel(M, N): # 并行:迭代独立,映射到线程
C[i, j] = A[i, j] + B[i, j]

for ko in T.Pipelined(iters, num_stages=3): # 流水:copy 与 compute 时间重叠
...

补充几点:

  • T.serial 支持三参数形式 T.serial(0, N, 2)(起点、终点、步长);
  • T.Parallel 可以加 coalesced_width= 提示控制访存合并宽度,loop_layout= 挂 fragment layout 标注;
  • 另有一个高级构造 T.Persistent,表达 persistent thread-block 风格的循环(5.1 提到的 stream-K 变体就靠它);
  • Python 原生的 if/elsewhilebreak/continue 都可用,条件是 TIR 表达式即可;潜在的越界访问由 LegalizeSafeMemoryAccess pass 自动加 guard(见第七节)。

前三个原语都好理解,真正值得单独一节展开的是最后一个。

T.Pipelined 做了什么

上面例子里最「魔法」的一行是 for ko in T.Pipelined(...)num_stages=3 不是「循环展开 3 次」,而是建立 3 级软件流水:

1
2
3
4
5
6
7
时间 ->
iter 0: [copy k=0]
iter 1: [copy k=1] [gemm k=0] ← 进入稳态:搬运与计算重叠
iter 2: [copy k=2] [gemm k=1]
...
iter N-1: [gemm k=N-2] ← epilogue:只剩计算
└─ HBM 延迟被后续 iter 的计算隐藏

手写 CUDA 要实现同样效果,需要自己管理 stage 数组下标、cp.async 的 commit/wait group、以及每级之间的 barrier。TileLang 把这压缩成一个参数。具体地,编译器在这一个循环上做了四件事:

  1. 循环重写:把源代码里的单层循环拆成 prologue(预取前 N-1 轮)/ 稳态 body(搬第 k+N-1 块的同时算第 k 块)/ epilogue(算完尾部) 三段。稳态时 copy 和 compute 在时间上重叠。
  2. 共享内存多缓冲num_stages=3 意味着 A_shared / B_shared 会被自动复制成 3 份(三缓冲),生产者写第 k+2 块、消费者读第 k 块,互不冲突。你在源码里写的是一份 buffer,编译器做的是 buffer 乘法。
  3. 异步拷贝插入:流水线里的 T.copy 会 lower 成异步拷贝(Ampere+ 上是 cp.async,Hopper+ 上是 TMA 的 cp.async.bulk),并自动配好 commit_group / wait_group 的配对。
  4. 同步插入:编译器的 PipelinePlanning / InjectSoftwarePipeline / InjectTmaBarrier 等一系列 pass 负责推导生产者-消费者依赖,在正确的地方插入 __syncthreads()(或 mbarrier)。这就是为什么源代码里一行同步都没有。

注意 T.copy 本身的语义是同步的–语句结束后 dst 就可读,如果 lower 到了异步指令,编译器会补上 wait 保证这一点。想手动控制异步,用 T.async_copy(不自动插 wait,需要自己写 T.ptx_wait_group)。

手动标注 stage / order

常规 GEMM 形态的流水线,num_stages=N 就够了,编译器自己推断谁是生产者谁是消费者。当循环体顺序不寻常(比如想让「下一轮的 copy」排在「这一轮的 compute」之前发射)时,可以显式标注:

1
2
3
4
5
6
7
for ko in T.Pipelined(
num_tiles,
stage=[0, 1], # copy 是 stage 0,gemm 是 stage 1
order=[1, 0], # 发射顺序上 gemm 先、copy 后
):
T.copy(A[ko * BK], A_shared)
T.gemm(A_shared, B_shared, C_local)

规则:

  • stage / order 与循环体内的可调度语句(copy、gemm、reduction、store、同步)按源码顺序一一对齐;
  • 流水线深度由 max(stage) + 1 推断,此时不要再传 num_stages
  • 循环体里的标量别名(base = ko * BK 这类 Bind 语句)不占标注位–它们没有副作用,编译器会在每个消费者处按需重放;
  • 编译器会校验依赖:生产者的 stage 必须不晚于消费者,同 stage 内 order 必须生产者在消费者前。

一句话总结T.Pipelined(num_stages=N) = 循环三分重写 + shared memory N 重缓冲 + 异步拷贝 + 自动同步,这是 GEMM/attention 类 kernel 能打到带宽/算力天花板的全部前提。


四、同一份代码,不同架构的 lowering

TileLang 源码里只有 T.copyT.gemm 这两个「意图」,落到哪条指令由 target 决定。这是它和直接写 CUDA 最大的分工差异–你描述数据流和 tile 结构,编译器按架构选指令

架构 异步拷贝(流水线内的 T.copy MMA 指令(T.gemm
SM70~75(V100/T4) SIMT ld.global + st.shared mma.sync(fp16)
SM80~89(A100/3090/4090/Ada) cp.async 多级流水 mma.sync
SM90a(H100/H200) TMA(cp.async.bulk.tensor)+ mbarrier wgmma(warp group MMA)
SM100a(B100/B200) TMA + mbarrier tcgen05.mma(TMEM 累加)
SM120(RTX 50 / RTX PRO,消费级 Blackwell) 没有 TMA,走 cp.async(LDGSTS) 普通 mma(第五代 tensor core)

这里有个容易踩的坑:不能按 SM 版本号大小推测能力。SM120 数字上大于 SM90,但它没有 TMA、没有 tcgen05,异步拷贝走的是 Ampere 时代引入的 cp.async。在 SM120 卡上写完全相同的 T.Pipelined 代码,编译器会自动退回 cp.async 风格的流水–语义不变,只是底层搬运指令不同。这正是 T.Pipelined 这个抽象的价值:你表达的是「我要 N 级预取流水」这个意图,而不是「我要发 cp.async.bulk.tensor」这个指令。

target 通过三种方式指定:

1
2
3
4
5
6
7
kernel = tilelang.compile(func, target={"kind": "cuda", "arch": "sm_90"})

@tilelang.jit(target="cuda") # 或裸字符串
def factory(...): ...

# 或环境变量(适合整台机器固定 GPU 型号的场景)
# export TILELANG_DEFAULT_TARGET='{"kind": "cuda", "arch": "sm_90"}'

auto(默认)按 CUDA → HIP → Metal 顺序探测。arch 直接对应 NVCC 的 -arch=sm_XX;需要一份代码出多个 SASS 时用 code 列表(fatbin)。跨厂商同理:HIP 配 mcpu="gfx90a",Metal / LLVM CPU / WebGPU 各有对应 kind。

判断「我写的代码在这张卡上到底变成了什么」,最直接的办法还是后面第六节的 dump 源码–看生成的 CUDA 里是 cp.async、TMA descriptor 还是 wgmma,比读文档更可靠。


五、进阶主题

5.1 Split-K:grid 第三维 + atomic_add

M、N 小而 K 很大时,M×N 的 tile 数不够填满 SM,就把 K 切开分给多个 block 并行累加。TileLang 里不需要新原语–T.Kernel 的第三个 grid 维就是 split 因子,最后用 T.atomic_add 归约(来自 examples/gemm_splitk):

1
2
3
4
5
6
7
8
9
10
11
12
splitK = K // split_k

with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), split_k, threads=128) as (bx, by, bz):
...
T.clear(C_local)
for ko in T.Pipelined(T.ceildiv(splitK, block_K), num_stages=0):
T.copy(A[by * block_M, bz * splitK + ko * block_K], A_shared) # K 维偏移带上 bz
T.copy(B[bz * splitK + ko * block_K, bx * block_N], B_shared)
T.gemm(A_shared, B_shared, C_local)

for i, j in T.Parallel(block_M, block_N):
T.atomic_add(C[by * block_M + i, bx * block_N + j], C_local[i, j])

要点:bz 只出现在 K 维索引里;累加结束整体做一次 atomic_add(而不是每个元素多次原子写);输出 C 必须先清零;num_stages=0 是官方示例关掉了自动流水(生产代码里该开的还是要开)。注意 atomic_add 的归约顺序不定,数值不可复现–对复现性有要求的场合要改成两阶段确定性归约(partial 先写回 workspace 再单独 reduce),本博客 DeepSeek-V4 mHC Pre-Block 融合 Kernel 详解 里有完整分析。同一目录下还有 stream-K 变体(gemm_streamk,把尾部 wave 按 K 拆给 peer block 再 fixup),是 persistent kernel 的入门样本。

5.2 warp specialization:T.ws + mbarrier

Hopper 之后,生产者和消费者可以拆成不同 warp group,各自跑各自的「循环」,靠 mbarrier 握手。TileLang 用 T.ws(i) 划分角色、T.alloc_barrier 建握手信号(来自 examples/warp_specialize):

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
with T.Kernel(..., threads=256) as (bx, by):
A_shared = T.alloc_shared((block_M, block_K), dtype)
...
data_is_ready = T.alloc_barrier(arrive_count=128)
compute_is_done = T.alloc_barrier(arrive_count=128)

with T.ws(1): # 消费者 warp group
T.clear(C_local)

for ko in T.Pipelined(T.ceildiv(K, block_K), num_stages=0): # 手动流水,不走自动重写
with T.ws(0): # 生产者:等上一轮算完 → TMA 搬数 → 通知就绪
T.barrier_wait(compute_is_done, (ko + 1) % 2)
T.tma_copy(A[by * block_M, ko * block_K], A_shared, barrier=data_is_ready)
T.tma_copy(B[ko * block_K, bx * block_N], B_shared, barrier=data_is_ready)
T.barrier_arrive(data_is_ready)
with T.ws(1): # 消费者:等数就绪 → gemm → 通知算完
T.barrier_wait(data_is_ready, ko % 2)
T.gemm(A_shared, B_shared, C_local)
T.barrier_arrive(compute_is_done)

with T.ws(1):
T.copy(C_local, C[by * block_M, bx * block_N])

要点:

  • num_stages=0 关掉自动流水–因为流水逻辑已经由两个 warp group 的 barrier 协议手动表达了,双缓冲体现在 barrier_wait% 2 相位翻转上;
  • T.tma_copy(..., barrier=...) 是显式 TMA 入口,比 T.copy 更低一层;
  • 这是 TileLang 里「控制力接近 CUTLASS」的具体形态:mbarrier 的 arrive_count、相位奇偶全部由你负责。examples 目录里同族还有 barrierpipe / softpipe 等多种流水协议写法可以对照。

5.3 autotune:tile 参数交给搜索

block_M / block_N / block_K / num_stages / threads 这组参数的理论最优值依赖具体 GPU 和问题规模,手调不现实。TileLang 内置 autotuner,用法是把可调参数写成带默认值的函数参数,再套一层装饰器:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
def matmul_configs(M, N, K):
return [
dict(block_M=BM, block_N=BN, block_K=BK, num_stages=S, threads=TH)
for BM in [64, 128]
for BN in [64, 128]
for BK in [32, 64]
for S in [2, 3]
for TH in [128, 256]
]

@tilelang.autotune(configs=matmul_configs, warmup=25, rep=100, timeout=60)
@tilelang.jit(out_idx=[-1])
def matmul(M: int, N: int, K: int,
block_M: int = 128, block_N: int = 128, block_K: int = 32,
threads: int = 128, num_stages: int = 3, ...):
...

with set_autotune_inputs(a, b, c): # 固定输入,保证各 config 可比
tuned = matmul(M, N, K) # 编译+验证+benchmark 全部 config,返回最优

值得知道的工程细节:候选 kernel 并行编译、逐个 benchmark;每个 config 会先过正确性检查(ref_prog 或默认的 torch 对比,容差 rtol/atol 默认 1e-2);结果缓存在 ~/.tilelang/cache/autotuner,缓存 key 包含 TileLang 版本 + 函数源码 + config 列表,改代码自动失效。调稳定后建议把最优 config 烘焙成函数默认值写回源码,autotune 只作为开发期工具。

5.4 Blackwell 两条路径

Blackwell 不是一个统一架构,写跟架构相关的 kernel 时必须分开看:

SM100a(B100/B200,数据中心) SM120(RTX PRO / RTX 50,消费级)
异步拷贝 TMA + mbarrier cp.async(LDGSTS),无 TMA
MMA 指令 tcgen05 + TMEM 普通 mma
two-SM(2-CTA)kernel
NVFP4 block-scaled T.mma_gemm_blockscaled T.mma_gemm_blockscaled(2026-07-30 加入)
对应 example blockscaled_gemm_sm100gemm_tcgen05 gemm_sm120

SM100a-- tcgen05 路径。第五代 Tensor Core 的 MMA 指令 tcgen05.mma 从 shared memory 直读操作数、累加到独立的 Tensor Memory(TMEM),还支持两个 CTA 配对发射。对应到 TileLang 是一组新原语:T.alloc_tmem 分配 TMEM 累加器,T.tcgen05_gemm 发射 MMA(不带隐式等待),T.alloc_barrier + T.mbarrier_wait_parity(mbar, k % 2) 手动做相位同步。这条路径目前是实验性 preview:同步协议要自己写,官方 README 明说「manual implementation required」。较新版本提供了半自动入口 T.gemm(..., mbar=...)–发射后自动插入匹配的 mbarrier_wait_parity,并把 fence 插入交给 InjectTcgen05Fence pass。examples/gemm_tcgen05/ 下有从裸 tcgen05_gemm 到 warp-specialized persistent、再到 2-CTA stream-K 的完整梯度。

SM120-- 传统路径。没有 tcgen05、TMEM 和 TMA,走的仍是 SM80 风格的 mma.sync + cp.async 流水线,TileLang 现有代码基本直接可用。换句话说,为 Hopper 写的 kernel 迁到 RTX 50 通常只是换个 arch,迁到 B200 才需要考虑 TMEM 那套新原语。两边唯一真正共享的新能力是 NVFP4 block-scaled MMA(T.mma_gemm_blockscaled,SM120 路径 2026-07-30 加入,可对照 Colfax 那篇 NVFP4 Blockscaled GEMM on RTX Pro Blackwell (sm12x)),但底层指令并不相同。

编译 target 上,数据中心 Blackwell 需要 fatbin 时可以 {"kind": "cuda", "arch": "sm_100f", "code": ["sm_100a", "sm_103a"]} 一份代码出多个 SASS。


六、写完之后的三个动作

quickstart 的官方流程把「kernel 写完之后」固化成了三步,建议形成肌肉记忆:

① 验证正确性–先于一切性能讨论,且以 fp32 参考为准

1
2
ref32 = torch.relu(a.float() @ b.float())
torch.testing.assert_close(c.float(), ref32, rtol=1e-2, atol=0.05)

为什么要跟 fp32 参考比而不是直接 torch.relu(a @ b):kernel 内部用 fp32 累加,a @ b 则是 fp16 累加,拿后者做参考的话误差来源两边不一致–你分不清到底是自己 kernel 错了还是累加精度差异。用 a.float() @ b.float() 做基准,差异才能归因到 kernel 本身。K=1024 累加下 atol=0.05 是合理范围。更复杂的 kernel(attention、MoE)可以保留一个 PyTorch 参考实现专门做这件事;TileLang 的 autotuner 也用同样思路验证每个候选 config。

② dump 生成源码–看编译器到底做了什么:

1
2
cuda_source = matmul_kernel.get_kernel_source()
print(cuda_source)

返回的是最终生成的完整 CUDA 源码。这一步回答所有「lowering 疑问」:T.copy 变成了 cp.async 还是 TMA descriptor?T.gemmmma.sync 还是 wgmma?多缓冲分配了多少 shared memory?__syncthreads() 插在了哪?调性能之前先读一遍生成代码,能省掉大量盲猜。换个 num_stagesblock_K 再 dump 一次、diff 两份源码,比读任何文档都直观。

③ benchmark–拿可信的延迟数字:

1
2
profiler = matmul_kernel.get_profiler(tensor_supply_type=tilelang.TensorSupplyType.Normal)
latency = profiler.do_bench() # ms

get_profiler 会自动生成输入、warmup、多次重复取统计,不用自己写 torch.cuda.synchronize() + time.time() 那套容易测错的东西。注意 TensorSupplyType.Normal 指定用正态分布造输入–对 GEMM 无影响,但对带 exp 的 attention kernel 会影响数值路径,别用全 0。没有 JITKernel 对象时也可以直接 from tilelang.profiler import do_bench 包一个 callable。有了 ① 的正确性和 ③ 的基线数字,后面任何改动(换 tile 尺寸、加 swizzle、上 warp specialization)都是可度量的。


七、学习路径与调试工具

环境与选卡

1
2
pip install tilelang
python -c "import tilelang; print(tilelang.__version__)"

需要 Python ≥ 3.10。想要最新特性走 nightly:pip install tilelang --find-links https://tile-ai.github.io/whl/nightly

没有 GPU 怎么办:TileLang 的 CUDA 后端需要真卡才能编译执行,Apple Silicon 可以走 metal target 但例子覆盖有限。学 TileLang 属于典型的短时高强度用卡–跑几小时 examples 就停,不需要长期占资源,RunPod 按小时租一张卡是最省事的路子。

选卡要看想学什么(原因见第四节):

学习目标 需要的卡
基础 tile / T.Pipelined / autotune 任意 SM80+,一张 4090 就够
TMA、T.tma_copy、warp specialization 必须 SM90a(H100/H200)
tcgen05 MMA、TMEM、two-SM kernel 必须 SM100a(B100/B200)
SM120 NVFP4 block-scaled RTX PRO 6000 / RTX 50 系

别拿 SM120 的卡去学 TMA–消费级 Blackwell 没有这个硬件单元。

调试工具箱

TileLang 在调试工具链上比 Triton 强,别浪费:

  • T.print(buffer, msg=...):kernel 内部打印 shared/fragment buffer,TileLang 自动只从单个线程打印避免刷屏;配合 if i == 0: 谓词用。
  • T.device_assert(cond, msg):device 侧断言,CUDA target 上生效,排查越界和 NaN 比注释掉半段代码快得多。
  • Pass Visualizer(2026-07 加入):结构树浏览器,看每个编译 pass 对 IR 做了什么。
  • IR Lower Trace(2026-07 加入):逐 pass dump IR,定位「我写的 layout 到哪一步被改掉了」。
  • TileLang LSP(2026-08 开源):VSCode 里显示 buffer 的 shape / dtype / scope / 推断 layout 的 inlay hint,写 kernel 时不用反复回头查 shape,第一天就该装上。
  • layout 可视化:把 fragment layout 画出来,检查 bank conflict。
  • get_kernel_source():见第六节,最强的「调试器」其实是读生成的 CUDA。
  • 缓存目录 ~/.tilelang/cache/:autotuner 产物、编译出的 .so / cubin 都在这;改了代码行为不对时先怀疑缓存,TILELANG_DISABLE_CACHE=1 一键排除。
  • 越界防护:LegalizeSafeMemoryAccess pass 会在可能越界的访问处自动插 guard(证明安全则自动消除),所以边界处理很多情况下不用手写 if–但自定义边界逻辑仍建议显式写。

推荐顺序

阶段 材料 目标
TileLang Puzzles 10 题 建立 tile 思维,比读文档有效得多
quickstart.py + elementwise 摸清 5 个原语
语言基础文档 补齐 layout / 内存作用域概念
examples/gemm swizzle、autotune、架构特化
examples/flash_attention online softmax + 双 GEMM 融合
examples/gemm_fp8 / blockscaled_gemm_sm100 量化 GEMM,per-block scale
examples/gemm_splitk / warp_specialize 本文第五节的出处
examples/deepseek_mla / deepseek_v4 / deepseek_mhc 真实生产算子

Puzzles 优先这点要强调:TileLang 的文档偏参考手册风格,直接读容易只学到 API 名字、学不到「为什么这么切」。10 个难度递增的 puzzle 会强迫你自己想清楚 tile 划分。第 ⑧ 步的收益也不在语法,而在于看懂「一个生产级 LLM 算子是如何被拆成 tile 数据流的」。

一句话总结:语法半天就能过完,TileLang 的学习成本在「tile 级思维」–而这只能靠 Puzzles 和 examples 里那些真实 kernel 攒出来。


参考

线性注意力与 SSM:两条技术路线的完整推导

0. 为什么要固定大小的记忆状态

标准 softmax 注意力解码时必须缓存全部历史的 Key/Value,显存和单步延迟都是 O(nd)O(n \cdot d),上下文翻倍就一起翻倍。把历史压缩成一个固定大小的状态 SRd×dS \in \mathbb{R}^{d \times d},显存和单步开销变成 O(d2)O(d^2),与 nn 无关。

代价是这个状态必须会写、会改、会忘。两条路线各给出一半答案:线性注意力造出固定状态,SSM 教它怎么遗忘。


1. 路线一:线性注意力

1.1 去掉 softmax,结合律才可用

裸注意力 (QK)V(QK^\top)V 的两种括号化复杂度不同:先算 n×nn \times nQKQK^\topO(dn2)O(dn^2),先算 KVRd×dK^\top V \in \mathbb{R}^{d\times d} 则是 O(nd2)O(nd^2)。矩阵乘法满足 (AB)C=A(BC)(AB)C = A(BC),先算哪边是自由的。

softmax 挡在中间:它作用在 QKQK^\top 之后、乘 VV 之前,且是行内非线性。两个性质各堵死一条路——非线性使 softmax(QK)VQsoftmax(KV)\mathrm{softmax}(QK^\top)V \ne Q\,\mathrm{softmax}'(K^\top V),中间结果无法先合并;行内归一化的分母 jeqikj\sum_j e^{q_i \cdot k_j} 依赖该行所有 key,没有可以先结合掉的独立块。

所以平方复杂度不是矩阵乘法的错,是 softmax 的作用位置的错。

1.2 把非线性提前到 Q 和 K 各自身上

核技巧:只要相似度非负(Mercer 条件),就存在特征映射 ϕ()\phi(\cdot) 使 sim(q,k)=ϕ(q)ϕ(k)\mathrm{sim}(q, k) = \phi(q)^\top \phi(k)非线性被吸收进 ϕ\phi,分别作用于 q 和 k 各自身上,相似度本身变回线性内积,于是

(ϕ(Q)ϕ(K))V=ϕ(Q)(ϕ(K)V)\big(\phi(Q)\,\phi(K)^\top\big)V = \phi(Q)\big(\phi(K)^\top V\big)

线性注意力取 ϕ(x)=ELU(x)+1\phi(x) = \mathrm{ELU}(x) + 1,取值范围 (0,)(0, \infty) 恒正——这是核分解存在的条件,不是装饰。写成逐 token 形式:

Attn(q)=ϕ(q)Sϕ(q)z,S=iϕ(ki)vi,z=iϕ(ki)\mathrm{Attn}(q) = \frac{\phi(q)^\top S}{\phi(q)^\top z}, \qquad S = \sum_{i} \phi(k_i)\, v_i^\top,\quad z = \sum_i \phi(k_i)

SS 是 KV 外积矩阵(关联记忆本体),zz 是从 softmax 分母继承下来的归一化项。后续 DeltaNet/GDN/KDA 改用 RMSNorm,zz 就退场了,只有 SS 的更新规则一路演进下去。

两者都能写成递推,一步一读出:

St=St1+ϕ(kt)vt,ot=ϕ(qt)Stϕ(qt)ztS_t = S_{t-1} + \phi(k_t)\, v_t^\top, \qquad o_t = \frac{\phi(q_t)^\top S_t}{\phi(q_t)^\top z_t}

这就是 fast weights 视角:SS 是一块随输入不断被改写的「快速权重」,写入靠 ϕ(kt)vt\phi(k_t)v_t^\top,读出靠 ϕ(qt)St\phi(q_t)^\top S_t

1.3 纯加性状态的代价:记忆碰撞

St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top 只有叠加,没有删除。查询定义为右乘 SkS\,k,把递推展开就能看出问题:

Stk=ivi(kik)S_t\,k = \sum_i v_i\,(k_i^\top k)

每个 value 的系数是它存入时的 key 与查询 key 的内积。key 完全匹配则完整取回,正交则不干扰——这是「按 key 相似度加权取回 value」。但同一个 key 先后写入两个不同的 value 时,两个系数都是 1,读出的是两者之和而不是最新那个。这就是记忆碰撞。

修法是 delta rule:写入前先查旧值,只写差值。

St=St1+(vtSt1kt)ktS_t = S_{t-1} + \big(v_t - S_{t-1}k_t\big)k_t^\top

差值中的负项抵消掉旧 key 上的残余,等于先删除再写入,从而保证 StktvtS_t k_t \approx v_t

1.4 手算一个例子:擦除到底发生了什么

dk=dv=2d_k = d_v = 2S0=0S_0 = 0,依次写入三个 token(关键在 k3=k1k_3 = k_1,同一个 key 写入新值):

k1=[1,0], v1=[1,0];k2=[0.6,0.8], v2=[0,1];k3=[1,0], v3=[2,0]k_1 = [1, 0],\ v_1 = [1, 0];\qquad k_2 = [0.6, 0.8],\ v_2 = [0, 1];\qquad k_3 = [1, 0],\ v_3 = [2, 0]

线性注意力St=St1+vtktS_t = S_{t-1} + v_tk_t^\top,直接叠外积):

S1=[1000],S2=[100.60.8],S3=[300.60.8]S_1 = \begin{bmatrix}1&0\\0&0\end{bmatrix},\qquad S_2 = \begin{bmatrix}1&0\\0.6&0.8\end{bmatrix},\qquad S_3 = \begin{bmatrix}3&0\\0.6&0.8\end{bmatrix}

第三步的 v3k3=[2000]v_3k_3^\top = \begin{bmatrix}2&0\\0&0\end{bmatrix} 不清除 S2S_2 里已有的旧值,直接往上叠,第一行变成 [3,0][3, 0]——碰撞就在这一瞬发生。查询 k1k_1S3k1=[3,0.6]S_3k_1 = [3, 0.6]:旧值 v1v_1 与新值 v3v_3 完整叠加(两个系数都是 1),再混进 0.6 份 v2v_2k2k1=0.6k_2^\top k_1 = 0.6)。

delta rule(取 β=1\beta = 1,每步先查后写):

  • t=1t=1S0k1=0S_0k_1 = 0,差值 u1=v1=[1,0]u_1 = v_1 = [1, 0],得 S1=[1000]S_1 = \begin{bmatrix}1&0\\0&0\end{bmatrix}(与线性注意力相同);
  • t=2t=2:先查 S1k2=[0.6,0]S_1k_2 = [0.6, 0]k2k_2 在第一维有 0.6 分量,部分命中旧记忆),差值 u2=v2S1k2=[0.6,1]u_2 = v_2 - S_1k_2 = [-0.6, 1]。这个负分量就是擦除,它扣除 k2k_2 方向上已存的 v1v_1 残余,得 S2=[0.640.480.60.8]S_2 = \begin{bmatrix}0.64&-0.48\\0.6&0.8\end{bmatrix}
  • t=3t=3:先查 S2k3=[0.64,0.6]S_2k_3 = [0.64, 0.6],差值 u3=v3S2k3=[1.36,0.6]u_3 = v_3 - S_2k_3 = [1.36, -0.6],得

S3=[20.4800.8]S_3 = \begin{bmatrix}2&-0.48\\0&0.8\end{bmatrix}

查询 k1k_1S3S_3 第一列,得 [2,0][2, 0]——精确返回最新的 v3v_3。两个分量各自对应一次擦除:第一行的 2 是 u3u_3+1.36+1.36t=2t{=}2 碎掉的 0.64 补回到 2;第二行的 0 是 u3u_30.6-0.6 正好抵消 t=2t{=}2 混进来的 0.6。对比线性注意力的 [3,0.6][3, 0.6]:旧值和残余都还在里面。

DeltaNet 的雏形到此成立:固定大小的外积状态 + 先删旧再写新。


2. 路线二:SSM

2.1 连续 SSM 的定义

h(t)=Ah(t)+Bx(t),y(t)=Ch(t)h'(t) = A\,h(t) + B\,x(t), \qquad y(t) = C\,h(t)

其中 h(t)RNh(t) \in \mathbb{R}^N 是状态(历史的压缩),AA 决定旧记忆如何衰减,BB 是写入强度,CC 是读出权重。要解决的问题:把微分方程改写成递推式 ht=Aˉht1+Bˉxth_t = \bar A h_{t-1} + \bar B x_t,并求出 Aˉ,Bˉ\bar A, \bar B

用到矩阵指数 eM:=k0Mk/k!e^{M} := \sum_{k\ge0} M^k/k! 的两条性质:ddteAt=AeAt\frac{d}{dt}e^{At} = A\,e^{At},以及 (eAt)1=eAt(e^{At})^{-1} = e^{-At}。直觉上 eAΔe^{A\Delta} 就是「让系统按自身动力学自由演化 Δ\Delta 时间」的算子。

2.2 通解:积分因子法

定理h(t)=Ah(t)+Bx(t)h'(t) = A h(t) + B x(t) 满足初值 h(t0)h(t_0) 的解为

h(t)=eA(tt0)h(t0)+t0teA(ts)Bx(s)dsh(t) = e^{A(t-t_0)}\,h(t_0) + \int_{t_0}^{t} e^{A(t-s)}\,B\,x(s)\,ds

证明。移项得 h(t)Ah(t)=Bx(t)h'(t) - A\,h(t) = B\,x(t),两边左乘积分因子 eAte^{-At}

eAth(t)eAtAh(t)=eAtBx(t)e^{-At}h'(t) - e^{-At}A\,h(t) = e^{-At}B\,x(t)

左端恰是一个乘积的全导数,这是整个推导的机关:

ddt[eAth(t)]=eAth(t)eAtAh(t)=eAtBx(t)\frac{d}{dt}\Big[e^{-At}h(t)\Big] = e^{-At}h'(t) - e^{-At}A\,h(t) = e^{-At}B\,x(t)

t0t_0tt 积分,再左乘 eAte^{At}(用 eAteAs=eA(ts)e^{At}e^{-As} = e^{A(t-s)})即得结论。\blacksquare

两项分别是旧记忆自由演化区间内每一瞬输入贡献的叠加

2.3 零阶保持(ZOH)离散化

假设采样区间内输入保持常值 x(tk+τ)=xkx(t_k + \tau) = x_k。在通解中取 t0=tkt_0 = t_kt=tk+Δt = t_k + \Delta,把 xkx_k 提出积分号:

hk+1=eAΔAˉhk+(0ΔeAsBds)Bˉxkh_{k+1} = \underbrace{e^{A\Delta}}_{\bar A}\,h_k + \underbrace{\left(\int_0^{\Delta} e^{As}B\,ds\right)}_{\bar B}\,x_k

对级数逐项积分算出 Bˉ\bar B

0ΔeAsds=k0AkΔk+1(k+1)!=A1(eAΔI)\int_0^{\Delta} e^{As}\,ds = \sum_{k\ge0}\frac{A^k\Delta^{k+1}}{(k+1)!} = A^{-1}\big(e^{A\Delta} - I\big)

 Aˉ=eΔA,Bˉ=A1(eΔAI)B \boxed{\ \bar{A} = e^{\Delta A}, \qquad \bar{B} = A^{-1}\big(e^{\Delta A} - I\big)\,B\ }

ZOH 在「输入确为分段常数」的假设下不是近似而是精确等价Δ\Delta 很小时有 AˉI+ΔA\bar A \approx I + \Delta ABˉΔB\bar B \approx \Delta B(欧拉法),但 S4/Mamba 实现都用精确公式。

2.4 Aˉ=eΔA\bar A = e^{\Delta A} 就是遗忘门

AA 负定(Mamba-2 取 A=aIA = -a\cdot Ia>0a>0),则

Aˉt=eΔta(0,1)\bar A_t = e^{-\Delta_t a} \in (0, 1)

Δt\Delta_t \to \inftyAˉt0\bar A_t \to 0,清空旧记忆;Δt0\Delta_t \to 0Aˉt1\bar A_t \to 1,冻结状态。Mamba 让 Δt=softplus(Linear(xt))\Delta_t = \mathrm{softplus}(\mathrm{Linear}(x_t)),于是「步长」变成了看内容的遗忘门。


3. 两条路线在 Mamba-2 汇合

Mamba-2 把状态写成外积形式,矩阵 AA 退化为标量衰减:

St=αtSt1+vtkt,αt=eΔtaS_t = \alpha_t S_{t-1} + v_tk_t^\top, \qquad \alpha_t = e^{-\Delta_t a}

对照线性注意力的 St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top,只多了一个乘在旧状态上的 αt\alpha_t线性注意力给出了外积状态的样子,SSM 给出了衰减门,从这一步起两者在数学上是同一个东西的两个记法(SSD 框架)。

此后 GDN 在写入侧加 delta rule,KDA 把标量门打开成逐通道门 Diag(αt)\mathrm{Diag}(\alpha_t) 并给衰减加下界。骨架始终是同一条:状态 ×(衰减/删除算子)+(写入项)。后续演化见《KDA 的来龙去脉》


参考

KDA 的来龙去脉:从线性注意力到 Kimi Delta Attention

0. 出发点:固定大小的记忆

KDA(Kimi Delta Attention)是 Kimi K3 里负责长序列的那 69 层(总共 93 层)。它要解决的问题是把随序列线性膨胀的 KV cache 换成一个固定大小的状态矩阵O(nd)O(n \cdot d)O(d2)O(d^2),上下文翻倍时状态大小不变。

代价是这个 dk×dvd_k \times d_v 的矩阵要承载任意长度的序列,记忆必须可写、可改、可遗忘。演进路线就是一步步补齐这三件事:

  • 线性注意力St=St1+ϕ(kt)vtS_t = S_{t-1} + \phi(k_t)v_t^\top,只能累加(同一个 key 写两次读出的是叠加,即记忆碰撞);
  • Mamba-2:加标量衰减门 St=αtSt1+ktvtS_t = \alpha_t S_{t-1} + k_tv_t^\top,会遗忘但衰减全局统一;
  • DeltaNet:改为差值写入 St=(Iβtktkt)St1+βtktvtS_t = (I - \beta_tk_tk_t^\top)S_{t-1} + \beta_tk_tv_t^\top,支持定点改写;
  • GDN:两者结合 St=αt(Iβtktkt)St1+βtktvtS_t = \alpha_t(I - \beta_tk_tk_t^\top)S_{t-1} + \beta_tk_tv_t^\top
  • KDA:把标量门拆成逐通道门 Diag(αt)\operatorname{Diag}(\bm\alpha_t),各维度独立决定衰减速度,并加数值下界。

线性注意力与 SSM 两条路线的推导见前置篇《线性注意力与 SSM:两条技术路线的完整推导》,本文直接用其结论。

记号约定:两套写法互为转置

文献里状态矩阵有两种摆法,内容等价、互为转置,混用会让推导看起来「中途换了个式子」:

主约定(与 KDA 论文一致) 转置约定(chunkwise 推导常用)
状态形状 StRdk×dvS_t \in \mathbb{R}^{d_k \times d_v} S^tRdv×dk\hat S_t \in \mathbb{R}^{d_v \times d_k}
读出 ot=Stqto_t = S_t^\top q_t ot=S^tqto_t = \hat S_t q_t
写入项 ktvtk_t v_t^\top vtktv_t k_t^\top
擦除算子 左乘 (Iβtktkt)St1(I - \beta_t k_tk_t^\top)S_{t-1} 右乘 S^t1(Iβtktkt)\hat S_{t-1}(I - \beta_tk_tk_t^\top)

关系就是一次转置 S^t=St\hat S_t = S_t^\top。由于擦除算子 Ht=IβtktktH_t = I - \beta_tk_tk_t^\top 对称,转置可以直接穿过它。后文若看到 kvkv^\topvkvk^\top 互换、擦除算子从左跑到右,那是切换了约定,不是等式变了。


1. DeltaNet:从加法到差值写入

线性注意力的 SS 只能叠加,修正这一点正是 DeltaNet 的动机。写入差值,不写全值——写入前先用当前 key 把状态里已存的内容读一遍:

vold=St1kt,ut=βt(vtvold),St=St1+ktutv_{\mathrm{old}} = S_{t-1}^\top k_t, \qquad u_t = \beta_t (v_t - v_{\mathrm{old}}), \qquad S_t = S_{t-1} + k_t u_t^\top

其中 ktk_t 经 L2 归一化(kt2=1\|k_t\|_2 = 1),βt(0,1)\beta_t \in (0,1) 是写入强度。

1.1 差值写入 = 先删后写

上式看不出「擦除」在哪,代入 utu_t 展开即可,每一步只用外积的结合律:

St=St1+kt[βt(vtvold)]=St1+βtktvtβtkt(St1kt) ⁣=St1+βtktvtβtktktSt1=(Iβtktkt)St1删除:沿 kt 方向擦除+βtktvt新值\begin{aligned} S_t &= S_{t-1} + k_t\big[\beta_t(v_t - v_{\mathrm{old}})\big]^\top \\[2pt] &= S_{t-1} + \beta_t k_t v_t^\top - \beta_t k_t \big(S_{t-1}^\top k_t\big)^{\!\top} \\[2pt] &= S_{t-1} + \beta_t k_t v_t^\top - \beta_t k_t k_t^\top S_{t-1} \\[2pt] &= \underbrace{\big(I - \beta_t k_t k_t^\top\big) S_{t-1}}_{\text{删除:沿 } k_t \text{ 方向擦除}} + \underbrace{\beta_t k_t v_t^\top}_{\text{新值}} \end{aligned}

关键是倒数第二步:voldv_{\mathrm{old}} 自己就是由 St1S_{t-1} 算出来的,所以「减去旧值」必然能写成一个作用在 St1S_{t-1} 上的线性算子,而不是额外的加项。提公因式后 βtktkt-\beta_tk_tk_t^\top 并入单位阵,变成擦除算子。预测误差 vtvoldv_t - v_{\mathrm{old}} 就是 delta,Delta Rule 由此得名

差值不是选的,是推出来的。把记忆看成一个在线回归问题:希望 SktvtS^\top k_t \approx v_t,写成平方损失 L(S)=12Sktvt2\mathcal{L}(S) = \tfrac12\|S^\top k_t - v_t\|^2。它的梯度是

SL=kt(Sktvt) ⁣\nabla_S \mathcal{L} = k_t\,\big(S^\top k_t - v_t\big)^{\!\top}

即「钥匙 ⊗ 残差」——和一维情形 12(wxy)2(wxy)x\tfrac12(wx-y)^2 \Rightarrow (wx-y)x(误差乘输入)一模一样,只是乘法变外积。以 βt\beta_t 为步长走一步 SGD:

St=St1βtkt(St1ktvt) ⁣=(Iβtktkt)St1+βtktvtS_t = S_{t-1} - \beta_t k_t\big(S_{t-1}^\top k_t - v_t\big)^{\!\top} = \big(I - \beta_t k_tk_t^\top\big)S_{t-1} + \beta_t k_tv_t^\top

与前面逐字相同。于是 βt\beta_t 的角色也明确了:它就是学习率。写完立即读一次可以验证:

Stkt=(1βt)St1kt+βtvtS_t^\top k_t = (1-\beta_t)\,S_{t-1}^\top k_t + \beta_t v_t

读出是旧值与新值的凸组合,βt=1\beta_t = 1 时完全覆写。这也解释了 L2Norm 为何是前提:只有 kt=1\|k_t\| = 1 时系数才是干净的插值,否则变成 1βtkt21-\beta_t\|k_t\|^2,可能跌出 [0,1][0,1]

一句话:记忆 = 在线回归 → 平方损失 → 梯度 = 残差 ⊗ 钥匙 → 一步 SGD = 先删后写

1.2 擦除算子 HtH_t:秩 1 与特征值

Ht=IβtktktH_t = I - \beta_t k_tk_t^\top 是全文频率最高的矩阵。kkkk^\top 的第 jj 列是 kjkk_j \cdot k——换 jj 只换倍率、方向永远是 kk,所以它秩为 1。作为变换,

(kk)x=k(kx)=(kx)k(k k^\top)\, x = k\,(k^\top x) = (k^\top x)\cdot k

不管输入什么,输出永远落在 kk 张成的直线上。

放回 HtH_tII 让所有方向原样保留,减去 βtktkt\beta_tk_tk_t^\top 只在 ktk_t 一个方向上动刀,其余 d1d-1 个方向碰都不碰。特征值分两种:

Htkt=(1βtkt2)kt,Htx=x(xkt)H_t k_t = \big(1 - \beta_t\|k_t\|^2\big)k_t, \qquad H_t x = x \quad (\forall\, x \perp k_t)

配合 L2Norm 就是沿 ktk_t 缩放 1βt1-\beta_t、正交补方向为 1,即特征值落在 [1βt,1](0,1][1-\beta_t, 1] \subset (0,1]——不放大、不翻转、不发散,这就是数值稳定性的特征值表述。反例很直白:不归一化时取 k=(1,2)k = (1,2)^\topβ=0.6\beta = 0.6,则 1βk2=21-\beta\|k\|^2 = -2,反复作用必然发散。

特征值之所以到处出现,是因为它回答了迭代系统最关心的问题:一个变换反复作用很多次之后会怎样——在特征向量方向上就是 λn\lambda^nλ>1|\lambda|>1 爆炸、λ<1|\lambda|<1 衰减。这条线上的约束几乎都在围着它转(SSM 的 Aˉ\bar A 模长 1\le 1、GDN/KDA 的 α(0,1)\alpha \in (0,1)HtH_t[1βt,1][1-\beta_t,1]、以及后面 1/Γ1/\Gamma 的溢出)。

HtH_tHouseholder 变换 I2uuI - 2uu^\top 同型:取 βk2=2\beta\|k\|^2 = 2 时两者相同,β(0,1)\beta \in (0,1) 的 delta rule 相当于「没照到底的半面镜子」,代价是不再正交,好处是擦除强度成了可学习的连续量。Householder 及其 WY 表示的线性代数细节见《线性注意力的线性代数前置知识》

1.3 转移矩阵形式与并行的可能性

定义 Ht=IβtktktH_t = I - \beta_tk_tk_t^\top,更新写成

St=HtSt1+βtktvtS_t = H_t\, S_{t-1} + \beta_t k_t v_t^\top

这暴露了 DeltaNet 与 SSM 的同构:HtH_t 是随输入变化的状态转移矩阵(SSM 里是 Aˉ\bar A),βtktvt\beta_tk_tv_t^\top 是写入项(SSM 里是 Bˉxt\bar Bx_t)。记 Bt=βtktvtB_t = \beta_tk_tv_t^\top 逐层代入,递归可以彻底消掉:

St=(i=t1Hi)S0+it(j=ti+1Hj)BiS_t = \Big(\prod_{i=t}^{1} H_i\Big) S_0 + \sum_{i \le t} \Big(\prod_{j=t}^{i+1} H_j\Big) B_i

只剩矩阵乘法和求和。而矩阵乘法满足结合律,「从左往右扫」只是众多括号化之一——换成二叉树式两两合并,就能在 O(logn)O(\log n) 深度内并行完成。这是 parallel scan 与 chunkwise 并行化共同的理论根基。


2. GDN:把遗忘门与定点改写拼在一起

Gated DeltaNet = DeltaNet 的精确写入 + Mamba-2 的全局遗忘。核心洞察是两者互补:gating 是板擦(大面积擦除,但无法定点修改),delta rule 是铅笔(定点覆写,但无法快速清空)。

St=αt(Iβtktkt)St1+βtktvtS_t = \alpha_t\big(I - \beta_tk_tk_t^\top\big)S_{t-1} + \beta_t k_t v_t^\top

其中 αt(0,1)\alpha_t \in (0,1) 是数据相关的标量门(α=exp(Softplus(Linear(xt)))\alpha = \exp(-\mathrm{Softplus}(\mathrm{Linear}(x_t))),在 log 空间算以保数值稳定)。三种极限:αt1\alpha_t \to 1 退回 DeltaNet;βt1\beta_t \to 1kk 与已有记忆正交时退回 Mamba-2;αt0\alpha_t \to 0 是整表清零再写入——两者单独都做不到的新能力。

几何上,(Iβkk)(I - \beta kk^\top) 沿 kk 方向压缩状态(定向),α\alpha 把整个状态矩阵均匀缩小(全局),作用于不同自由度,所以可以叠加。

注意作用顺序α\alpha 乘的是整个 (Iβkk)(I-\beta kk^\top)。把 α\alpha 只乘到擦除项上,chunkwise 形式会与递归形式对不上——这是实现时最容易踩的坑。

2.1 统一视角:四个模型是同一个在线优化问题

至此四个模型都出现了,它们是同一个在线优化问题的闭式解,差别只在目标函数。把每步看成一次在线学习:已有 St1S_{t-1},新到样本 (kt,vt)(k_t, v_t),求 StS_t

L(St)=StAtF2正则:记忆保留2Stkt, ut拟合:关联学习\mathcal{L}(S_t) = \underbrace{\|S_t - A_t\|_F^2}_{\text{正则:记忆保留}} - \underbrace{2\langle S_tk_t,\ u_t\rangle}_{\text{拟合:关联学习}}

正则项惩罚状态的改动量,锚点 AtA_tSt1S_{t-1} 表示尽量不动、取 αtSt1\alpha_tS_{t-1} 表示容忍遗忘;拟合项要求用 ktk_t 检索的结果朝写入目标 utu_t 对齐。目标对 StS_t 是二次的,=2(StAt)2utkt=0\nabla = 2(S_t - A_t) - 2u_tk_t^\top = 0,闭式解统一为 St=At+utktS_t = A_t + u_tk_t^\top。于是差别完全归结为两个选择:锚点决定怎么遗忘,写入目标决定怎么写入(下表用转置约定)。

模型 锚点 AtA_t 写入目标 utu_t 闭式解
Linear Attn St1S_{t-1} vtv_t St1+vtktS_{t-1} + v_tk_t^\top
Mamba-2 αtSt1\alpha_tS_{t-1} vtv_t αtSt1+vtkt\alpha_tS_{t-1} + v_tk_t^\top
DeltaNet St1S_{t-1} βt(vtSt1kt)\beta_t(v_t - S_{t-1}k_t) St1(Iβtktkt)+βtvtktS_{t-1}(I - \beta_tk_tk_t^\top) + \beta_tv_tk_t^\top
GDN αtSt1\alpha_tS_{t-1} βt(vtαtSt1kt)\beta_t(v_t - \alpha_tS_{t-1}k_t) St1(αt(Iβtktkt))+βtvtktS_{t-1}\big(\alpha_t(I-\beta_tk_tk_t^\top)\big) + \beta_tv_tk_t^\top
KDA St1Diag(αt)S_{t-1}\mathrm{Diag}(\bm\alpha_t) βt(vtSt1Diag(αt)kt)\beta_t(v_t - S_{t-1}\mathrm{Diag}(\bm\alpha_t)k_t) St1Diag(αt)(Iβtktkt)+βtvtktS_{t-1}\mathrm{Diag}(\bm\alpha_t)(I-\beta_tk_tk_t^\top) + \beta_tv_tk_t^\top

说到底,这就是把「kvk \to v 这条记忆是否还对得上」写成损失函数:对不上就修(拟合项),但别为了修这一条把整张表推翻(正则项)。GDN 在两处都取强化版本,KDA 再把锚点的标量收缩换成逐通道的 Diag(αt)\mathrm{Diag}(\bm\alpha_t)

2.2 Chunkwise 并行:WY 表示与 UT 变换

推理用递归(O(1)O(1)/token),但训练和 prefill 必须并行。设 chunk 大小 CC、入口状态 S0S_0,部分展开递归(转置约定):

Sr=γrS0Pr+i=1rγrγiu~iki,γr=jrαj,Pr=ir(Iβikiki)S_r = \gamma_r S_0 P_r + \sum_{i=1}^r\frac{\gamma_r}{\gamma_i}\,\tilde u_i k_i^\top, \qquad \gamma_r = \prod_{j\le r}\alpha_j,\quad P_r = \prod_{i\le r}(I - \beta_ik_ik_i^\top)

PrP_rCCdk×dkd_k\times d_k 矩阵逐个相乘,O(Cdk3)O(C\,d_k^3) 且严格顺序——比原递推还贵,必须处理。

关键观察:这种乘积永远不膨胀。每个因子都是「单位阵减秩 1」,连乘结果仍是「单位阵减一个低秩矩阵」:

Pr:=i=1r(Iβikiki)=IirwikiP_r := \prod_{i=1}^{r}(I - \beta_i k_i k_i^\top) = I - \sum_{i\le r} w_i k_i^\top

证明(对 rr 归纳)P0=IP_0 = I 成立。设 Pr1=Ii<rwikiP_{r-1} = I - \sum_{i<r} w_ik_i^\top,右乘第 rr 个因子:

Pr=(Ii<rwiki)βr(Ii<rwiki)krkrP_r = \Big(I - \sum_{i<r} w_i k_i^\top\Big) - \beta_r\Big(I - \sum_{i<r} w_i k_i^\top\Big)k_r k_r^\top

展开后第三、四项都以 krk_r^\top 结尾,合并同类项:

Pr=Ii<rwikiβr(kri<rwi(kikr))=wrkrP_r = I - \sum_{i<r} w_i k_i^\top - \underbrace{\beta_r\Big(k_r - \sum_{i<r} w_i (k_i^\top k_r)\Big)}_{=\,w_r} k_r^\top

归纳完成。wrw_r 不是发明的技巧,而是「乘积保持低秩」逼出来的——要让结果保持 IiwikiI - \sum_i w_ik_i^\top 的形式,括号里那一坨只能是 wrw_r。它的直觉是修正后的擦除向量:第 rr 步本想擦除 krk_r 方向,但若 krk_r 与之前的 kik_i 有重叠,连乘展开时前面的擦除已经顺带擦过这部分,wrw_r 把已擦的量减掉以避免重复擦除,权重恰是重叠度 kikrk_i^\top k_r。value 侧同理,只多一个衰减比:

u~r=βr(vri<ru~iγrγi(kikr))\tilde u_r = \beta_r\Big(v_r - \sum_{i<r}\tilde u_i\,\tfrac{\gamma_r}{\gamma_i}(k_i^\top k_r)\Big)

γ\gamma 只出现在 u~\tilde u 里而不在 ww 里,因为 GDN 的衰减是标量、与一切矩阵可交换,擦除连乘里的衰减可整体外提成 γr\gamma_r;但 value 侧每次写入的「存活时长」不同,这个相对衰减无法外提。对照 KDA:衰减变成向量后与擦除不可交换,γ\gamma 再也提不出去,只能渗进内积本身。

UT 变换:递归变成一次下三角求解wrw_r 只依赖 wi<rw_{i<r},是严格下三角依赖。把递归移项 wr+i<rβr(kikr)wi=βrkrw_r + \sum_{i<r}\beta_r(k_i^\top k_r)w_i = \beta_rk_r,令 WW 的第 rr 行为 wrw_r^\topL=strictLower(diag(β)KK)L = \mathrm{strictLower}(\mathrm{diag}(\beta)KK^\top),则 CC 个方程一次写成 (I+L)W=diag(β)K(I+L)W = \mathrm{diag}(\beta)K

W=(I+L)1diag(β)KW = (I+L)^{-1}\,\mathrm{diag}(\beta)\,K

求这个逆很便宜,因为严格下三角矩阵幂零LC=0L^C = 0),Neumann 级数有限项精确截断:(I+L)1=IL+L2(I+L)^{-1} = I - L + L^2 - \cdots,一次前代法即可,O(C2)O(C^2)(I+L)(I+L)单位下三角(Unit Triangular,对角为 1 因为 wrw_r 完整依赖自己),这就是「UT」的来源。所以 UT 不是额外发明的东西,它就是这两条递归的矩阵形态

乘法链 → 加法链,这才是并行的真正来源。本质是把擦除矩阵的连乘 r(Iβrkrkr)\prod_r(I - \beta_rk_rk_r^\top) 换成求和 IrwrkrI - \sum_r w_rk_r^\top,代价是求和项带修正。求和为什么就是胜利:加法可交换、可结合,因而可任意分组——树形归约、分块、稠密 matmul,GPU 的全部并行性都在奖励求和结构;而连乘必须一步一步来。同样的手法在这条路线上反复出现:αs=exp(gs)\prod\alpha_s = \exp(\sum g_s)(累积衰减变 cumsum)、SSM 的卷积形式、Mamba-2 的 SSD 半可分矩阵,以及本节的 WY-UT。

辨析:因果掩码 ≠ UT 变换。两者都是下三角,但管的是两件事:因果掩码把 score 的严格上三角置零,管「不许看未来」;UT 变换处理块内历史写入之间的相互影响——delta rule 每次写入都「先读再改」,块内第 2 次写入读到了第 1 次的结果,UT 负责解耦。都是下三角不是巧合,是同一个原因:因果性使干扰系数矩阵天然下三角,对角线天然为 1。因果性决定了它是三角的,但做它的目的是解耦,不是掩码。


3. 从 GDN 到 KDA

维度 GDN KDA(→ K3) 动机
更新式 St1(αt(Iβkk))+βvkS_{t-1}(\alpha_t(I-\beta kk^\top)) + \beta vk^\top (Iβkk)Diag(αt)St1+βvk(I-\beta kk^\top)\mathrm{Diag}(\bm\alpha_t)S_{t-1} + \beta vk^\top α 标量 → 逐通道向量
门粒度 标量(整头同一衰减率) 向量 αt(0,1)dk\bm\alpha_t\in(0,1)^{d_k} 长期记忆通道与短期工作区分离
作用顺序 α 在外乘整个更新 Diag(α) 在内侧先衰减 S 与逐通道参数化配套
α 参数化 负 Softplus,(,0)(-\infty,0) 缩放 sigmoid,下界 gmin=5g_{\min}=-5 1/Γ1/\Gamma 有界 → 全走 Tensor Core
chunkwise Γ 是标量比 Γ 是向量累积比 表达力↑ 数值难度↑

3.1 KDA 的递推公式

回到主约定(StRdk×dv\mathbf{S}_t \in \mathbb{R}^{d_k\times d_v},读出 Stqt\mathbf{S}_t^\top\bm q_t),与论文写法一致:

St=(Iβtktkt)Diag(αt)St1+βtktvt\mathbf{S}_t = \left(\mathbf{I} - \beta_t \bm{k}_t \bm{k}_t^\top\right) \operatorname{Diag}(\bm{\alpha}_t)\, \mathbf{S}_{t-1} + \beta_t \bm{k}_t \bm{v}_t^\top

按从右往左三步:通道衰减 Diag(αt)St1\operatorname{Diag}(\bm\alpha_t)\mathbf{S}_{t-1} 对每一行分别乘对应通道的 α\alpha方向擦除 (Iβtktkt)(\mathbf{I} - \beta_t\bm k_t\bm k_t^\top) 做定点擦除;写入 βtktvt\beta_t\bm k_t\bm v_t^\top

标量门的瓶颈在于所有 key 通道以同一比例衰减——要么一起记住,要么一起忘记。但有的通道在存「长期主题」(希望 α1\alpha \approx 1),有的在存「临时指针」(希望快速腾出容量)。Diag(αt)\mathrm{Diag}(\bm\alpha_t) 相当于给每一列配一个独立的遗忘速度,GDN 是 Diag(αt)=αtI\mathrm{Diag}(\bm\alpha_t) = \alpha_t I 的特例。在线学习视角下,这就是把 §2.1 的正则项换成逐通道版:第 jj 列的信任区域宽度正比于 αt(j)\alpha_t^{(j)}α\alpha 小的通道旧记忆先被压缩、新写入几乎没有阻力。KDA = 逐通道权重衰减 + delta rule。

作用顺序不可颠倒Dt=Diag(αt)D_t = \mathrm{Diag}(\bm\alpha_t)(Iβkk)(I - \beta kk^\top) 不可交换,写成 St1(Iβkk)DtS_{t-1}(I-\beta kk^\top)D_t 或只衰减单位阵部分都是错的。擦除时检索用的也是已衰减的状态。

3.2 逐通道衰减下的 chunkwise

核心记号是累积 log 衰减 γi=sigsRdk\gamma_i = \sum_{s\le i} g_s \in \mathbb{R}^{d_k}。带衰减的 key-key 内积:

Mci=(kceγcγi)ki,1i<cCM_{ci} = \big(k_c \odot e^{\gamma_c - \gamma_i}\big)^\top k_i, \qquad 1 \le i < c \le C

含义是 kik_i 写入的记忆衰减到第 cc 步时与 kck_c 的重叠程度。衰减「长」在 M 内部,无法外提——这是向量门控带来的结构性变化,也是 KDA chunkwise 推导的核心难点。其余步骤与 GDN 同构:构造 Lci=βiMciL_{ci} = \beta_i M_{ci},求 T=(I+L)1T = (I+L)^{-1},得 A^=Tdiag(β)\hat A = T\operatorname{diag}(\beta),然后

W=A^(eγiki)i=1..C,U=A^V,V~=UWS[0]W = \hat{A}\,\big(e^{\gamma_i} \odot k_i\big)_{i=1..C}, \qquad U = \hat{A}\,V, \qquad \tilde{V} = U - W\,S_{[0]}^\top

V~\tilde V伪值:它不是真实的 vv,而是扣除了「历史记忆已有部分 + 块内其他位置已写部分」之后可直接累加、无需再修正的净增量。这与单步 delta rule 的 vtSt1ktv_t - S_{t-1}^\top k_t 一致,V~\tilde V 是它的 chunk 并行版。输出与出口状态:

Acjqk=(qceγcγj)kj,oc=S[0](qceγc)+jcAcjqkv~jA^{qk}_{cj} = \big(q_c \odot e^{\gamma_c - \gamma_j}\big)^\top k_j, \qquad o_c = S_{[0]}\,(q_c \odot e^{\gamma_c}) + \sum_{j \le c} A^{qk}_{cj}\,\tilde v_j

S[C]=S[0]ΓC0+c=1Cv~c(kceγCγc)S_{[C]} = S_{[0]}\,\Gamma_{C\leftarrow 0} + \sum_{c=1}^{C} \tilde v_c \big(k_c \odot e^{\gamma_C - \gamma_c}\big)^\top

所有指数都是 0\le 0 的差,整条链路没有一个 1\ge 1 的因子——这是刻意保持的,原因见 §3.4。

3.3 数值验算(dk=dv=2d_k = d_v = 2C=2C = 2

零初始状态,每步恒定衰减 α=(0.5,0.25)\bm\alpha = (0.5, 0.25)β=1\beta = 1

k1=(11), v1=(20);k2=(12), v2=(02);q1=q2=(11)k_1 = \binom{1}{1},\ v_1 = \binom{2}{0};\quad k_2 = \binom{1}{2},\ v_2 = \binom{0}{2};\quad q_1 = q_2 = \binom{1}{1}

递推式。第 1 步(S0=0S_0 = 0):S1=v1k1=(2200)S_1 = v_1k_1^\top = \begin{pmatrix}2&2\\0&0\end{pmatrix}o1=(4,0)o_1 = (4,0)。第 2 步先衰减 S1D=(10.500)S_1D = \begin{pmatrix}1&0.5\\0&0\end{pmatrix}(第 1 列 ×0.5、第 2 列 ×0.25——逐通道在动),再擦除写入:

S1D(Ik2k2)=(13.500),S2=(13.524),o2=(4.5, 6)S_1D(I - k_2k_2^\top) = \begin{pmatrix}-1&-3.5\\0&0\end{pmatrix}, \qquad S_2 = \begin{pmatrix}-1&-3.5\\2&4\end{pmatrix}, \qquad o_2 = (-4.5,\ 6)

Chunkwiseeγ1=(0.5,0.25)e^{\gamma_1} = (0.5, 0.25)eγ2γ1=(0.5,0.25)e^{\gamma_2-\gamma_1} = (0.5, 0.25),衰减 KKT M21=(0.5,0.5)(1,1)=1M_{21} = (0.5, 0.5)\cdot(1,1) = 1。UT:L=(0010)L = \begin{pmatrix}0&0\\1&0\end{pmatrix}T=ILT = I - LA^=T\hat A = T。WY:

W=A^(0.50.250.250.125)=(0.50.250.250.125),U=A^(2002)=(2022)W = \hat A\begin{pmatrix}0.5&0.25\\0.25&0.125\end{pmatrix} = \begin{pmatrix}0.5&0.25\\-0.25&-0.125\end{pmatrix}, \quad U = \hat A\begin{pmatrix}2&0\\0&2\end{pmatrix} = \begin{pmatrix}2&0\\-2&2\end{pmatrix}

伪值(S[0]=0V~=US_{[0]}=0 \Rightarrow \tilde V = U):v~1=(2,0)\tilde v_1 = (2,0)v~2=(2,2)\tilde v_2 = (-2,2)注意 v~2v2\tilde v_2 \ne v_2:因为 k2k_2 与衰减后的 k1k_1 写入重叠(M21=1M_{21} = 1),WY 把 v2v_2 修正为扣除重叠后真正的新增。衰减注意力与输出:

Aqk=(200.753),o1=2v~1=(4,0) ,o2=0.75v~1+3v~2=(4.5, 6) A^{qk} = \begin{pmatrix}2&0\\0.75&3\end{pmatrix}, \qquad o_1 = 2\tilde v_1 = (4,0)\ \checkmark, \qquad o_2 = 0.75\,\tilde v_1 + 3\,\tilde v_2 = (-4.5,\ 6)\ \checkmark

块末状态 S[2]=v~1(k1eγ2γ1)+v~2k2=(13.524)S_{[2]} = \tilde v_1(k_1 \odot e^{\gamma_2-\gamma_1})^\top + \tilde v_2 k_2^\top = \begin{pmatrix}-1&-3.5\\2&4\end{pmatrix} \checkmark,与递推式逐项一致。

3.4 下界衰减:1/Γ1/\Gamma 溢出

GDN 的标量衰减在 chunkwise 里只以比值出现(s=j+1iαs1\prod_{s=j+1}^i \alpha_s \le 1),天然安全。KDA 把 Γ\Gamma 变成向量累积后,某些等价写法(如论文公式 (4) 的 K/ΓK/\Gamma 因式分解)就得真的算出

Γt1=exp(stgs)(逐通道)\Gamma_t^{-1} = \exp\Big(-\sum_{s \le t} g_s\Big) \quad \text{(逐通道)}

衰减越快的通道,1/Γ1/\Gamma 呈指数膨胀:恒定 α=0.5\alpha = 0.5t=500t=500 已达 1015010^{150}t=2000t=2000 溢出 float64;BF16 训练下数十步即溢出。标量推广为向量后衰减率的动态范围被放大,溢出由个别情形变为普遍现象。

Kimi Linear 的规避方案是在对数空间算相对衰减、并把 chunk 再切成 16 token 二级瓦片:瓦片之间可交给 Tensor Core,但瓦片内部仍需按位置对显式计算,成为块内瓶颈。K3 的解决方式是不改公式、改参数化,用有界缩放 sigmoid 使溢出在数学上不可能发生:

gth=gminSigmoid(eAhzth)(gmin,0),gmin=5g_t^h = g_{\min} \cdot \text{Sigmoid}(e^{A_h} z_t^h) \in (g_{\min}, 0), \qquad g_{\min} = -5

于是 α>e56.7×103\alpha > e^{-5} \approx 6.7\times10^{-3},16 token 块的累积对数衰减落在 (80,0)(-80, 0),重缩放因子小于 e80e^{80},在 BF16 动态范围内。表达力不受损:快通道在瓦片内仍可衰减到 e801035e^{-80} \approx 10^{-35},长期遗忘靠多步连乘即可,不需要单步清零的能力。以一个衰减下界换取全链路 Tensor Core 化,是有利的取舍。


4. 网络里的其他部件

q,k,v,α,βq, k, v, \alpha, \beta 本身怎么来、o~\tilde o 怎么变成层输出:GDN/KDA 采用 Llama 式宏架构,把 attention 换成 gated delta token mixer。

1
2
3
4
5
6
x ─┬─ W_q/W_k ─ ShortConv ─ SiLU ─ L2Norm ──> q, k
├─ W_v ──── ShortConv ─ SiLU ───────────> v
├─ W_α(线性投影, log 空间)──────────────> α
├─ W_β(线性投影 + Sigmoid)─────────────> β
│ 递归/分块计算 gated delta rule ──────> õ
└─ y = W_o( Sigmoid(W_g x) ⊙ RMSNorm(õ) ), 再接 SwiGLU MLP + 残差

三个部件各负责一部分数值稳定性:

  • ShortConv:窗口 4 的因果 depthwise 卷积,ShortConv(x)t=j<Wwjxtj\mathrm{ShortConv}(x)_t = \sum_{j<W} w_j \odot x_{t-j}。补足线性 RNN 缺失的局部建模能力——递归状态只携带压缩后的全局历史,最近几个 token 的精细模式由这层负责。depthwise 使通道间不混,与逐通道衰减语义配套;
  • SiLU(即 β=1\beta=1 的 Swish):xσ(x)x \cdot \sigma(x),自门控、免参数,平滑非单调、梯度处处非零;
  • L2Norm:只用在 Q/K。kt=1\|k_t\|=1 是 delta rule 擦除稳定的前提(§1.2);Q 归一化使 QKQK^\top 的 score 限制在 [1,1][-1,1]。V 不做,因为写入内容的幅值本身携带信息。

K3 另将输出门升级为输入相关的满秩 sigmoid 门 yt=Wo[Sigmoid(Wgxt)RMSNorm(o~t)]y_t = W_o[\mathrm{Sigmoid}(W_gx_t) \odot \mathrm{RMSNorm}(\tilde o_t)],允许每个 token 独立调制从循环状态读取的通道。

整体骨架自始至终没变:

状态×(衰减/删除算子)+(写入项)\text{状态} \times (\text{衰减/删除算子}) + (\text{写入项})


参考

DSpark 的实现和测评

DSpark = DFlash 的并行 backbone forward(1 次)+ N 步轻量 Markov 序列修正,全部在 CUDA Graph 内。 本文结合 vLLM 源码分析 DSpark 的实现细节,并在 Qwen3-8B 上实测 deepseek-ai 官方 draft 和社区 Dogacel draft 的效果差异。


1. 背景:投机解码与并行起草

投机解码(Speculative Decoding, SD)用一个小 draft 模型并行猜测 N 个 token,再由 target 模型一次 verify,通过 rejection sampling 保证输出分布不变。SD 的收益来自把 decode 阶段的 memory-bound 转为 compute-bound——bs=1 时 GPU 利用率极低,draft 的轻量 GEMM + target 的 batched verify 填充了 GPU 空闲。

vLLM v1 的 SD 框架支持多种 method:eagleeagle3dflashdsparkmedusangrammtp 等。DSpark 继承自 DFlash,核心改进是序列马尔可夫采样

继承链:

1
2
3
4
BaseSpeculator (ABC)
└─ DraftModelSpeculator
└─ DFlashSpeculator
└─ DSparkSpeculator ← 本文主角

模型类继承链:

1
2
3
Qwen3ForCausalLM
└─ DFlashQwen3ForCausalLM
└─ Qwen3DSparkForCausalLM

2. DSpark vs DFlash:两个核心差异

DSpark 的 docstring 写得非常清楚,和 DFlash 的差异只有两点。

2.1 Anchor-as-first-prediction(锚位即首预测)

DFlash:每个 request 发 1 + N 个 query token(1 个 anchor/bonus + N 个 mask token)。anchor 是上一步验证通过的 token,只有 N 个 mask 位置做预测:

1
2
3
4
5
DFlash query layout (1+N=9, N=8):
[anchor] [mask] [mask] [mask] [mask] [mask] [mask] [mask] [mask]
↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑
bonus pred pred pred pred pred pred pred pred
(不采样)

DSpark:anchor 本身也是预测位置,每个 request 只发 N 个 query token:

1
2
3
4
5
DSpark query layout (N=8):
[anchor] [noise] [noise] [noise] [noise] [noise] [noise] [noise]
↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑
pred pred pred pred pred pred pred pred
(采样)

代码(DSparkSpeculator.__init__):

1
2
3
4
5
6
7
self.sample_from_anchor = getattr(
self.draft_model_config.hf_config, "sample_from_anchor", True
)
if self.sample_from_anchor:
self.num_query_per_req = self.num_speculative_steps # N
else:
self.num_query_per_req = 1 + self.num_speculative_steps # 1+N (兼容旧格式)

在 Triton kernel _prepare_dflash_inputs_kernel 中,通过 SAMPLE_FROM_ANCHOR 编译常量控制采样行为:

1
2
3
4
# DSpark: 所有 N 个位置都采样,sample_pos = query_pos + 1(标准 next-token)
sample_off = 0 if SAMPLE_FROM_ANCHOR else 1
is_sample = is_query & (query_off >= sample_off)
sample_pos = query_pos + 1 if SAMPLE_FROM_ANCHOR else query_pos

2.2 Sequential Markov Sampling(序列马尔可夫采样)

这是 DSpark 的核心创新。

DFlash:N 个 mask 位置的 hidden states 一次性并行采样,各位置之间无依赖。

DSpark:先并行 forward 得到所有 N 个位置的 hidden states,然后从左到右逐个采样,每步用前一个采样出的 token 注入一个 Markov bias:

1
2
3
4
5
6
7
8
9
10
11
12
并行 backbone forward → [h₀, h₁, h₂, ..., h₇]
│ │ │ │
▼ ▼ ▼ ▼
base_logits[0] base_logits[1] ... base_logits[7]
+ + +
markov_bias( markov_bias( markov_bias(
anchor) sample₀) sample₆)
│ │ │
▼ ▼ ▼
sample₀ sample₁ ... sample₇

└──────────────→ 传给下一步作为 prev

代码在 _sample_sequential

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
def _sample_sequential(self, num_reqs, head_hidden):
n_spec = self.num_speculative_steps
# 1. 一次性算出所有 N 个位置的 base logits
base_logits = self.model.compute_draft_logits(sample_hidden) # [B, N, V]

# 2. anchor token 作为初始 prev
prev = self.input_buffers.input_ids[self._anchor_idx[:num_reqs]]

# 3. 逐位置采样
for i in range(n_spec):
markov_embed = self.model.markov_embed(prev) # [B, r]
bias = self.model.markov_bias(markov_embed) # [B, V]
logits_i = base_logits[:, i] + bias # 加上 Markov 偏置
draft_sampled_i = gumbel_sample(logits_i, ...) # 采样
self.draft_tokens[:num_reqs, i] = draft_sampled_i
prev = draft_sampled_i # 传给下一步

一句话总结:并行 forward 拿到所有位置的 base prediction,再用 N 步轻量 Markov 修正注入序列依赖——把「N 个独立预测」变成「N 个有依赖的预测」。


3. Markov Head 结构

DSparkMarkovHead 是一个 low-rank 转移偏置头:

1
2
3
4
5
6
7
8
9
prev_token_id

│ markov_w1: Embedding(V, r) ← V 是 vocab_size,r 是 markov_rank

markov_embed [B, r]

│ markov_w2: ParallelLMHead(r, V) ← r → V 的线性投影

markov_bias [B, V] ← 加到 base_logits 上

代码(qwen3_dspark.py):

1
2
3
4
5
class DSparkMarkovHead(nn.Module):
def __init__(self, vocab_size, draft_vocab_size, markov_rank, ...):
self.markov_w1 = nn.Embedding(vocab_size, markov_rank) # V×r
self.markov_w2 = ParallelLMHead(
draft_vocab_size, markov_rank, bias=False, disable_tp=True) # r×V

两个权重都是 replicateddisable_tp=True),因为 Markov head 每步都跑,分片会引入 all-reduce 和 full-vocab gather。

参数量 = 2×V×r2 \times V \times r。当 V=151936V=151936(Qwen3 词表)、r=64r=64 时约 19.4M 参数,相比 8B backbone 可以忽略。


4. 完整的 Draft 一步流程

DSparkSpeculator._generate_draft 只有两行:

1
2
3
def _generate_draft(self, num_reqs, num_tokens_padded, ...):
head_hidden = self._run_model(...) # 1. 并行 backbone forward
self._sample_sequential(num_reqs, head_hidden) # 2. 序列 Markov 采样

Step 1:并行 Backbone Forward(继承自 DFlash)

  • 输入:N 个 query token(anchor + mask/noise),position 已对齐
  • 上下文 KV 已在 precompute_and_store_context_kv 中预填充
  • 非因果 attention:N 个 query 位置可以互相 attend
  • 整个 forward 被 CUDA Graph 捕获

Step 2:Sequential Markov Sampling(DSpark 独有)

  • 取出 N 个位置的 hidden states
  • 一次性算出 base logits(compute_draft_logits
  • 逐位置:base_logits[i] + markov_bias(prev) -> gumbel_sample
  • 这个循环也被 CUDA Graph 捕获(所有 buffer 预分配固定地址)

Context KV 预计算(DFlash 的关键优化)

避免逐层跑 target 的 forward 来填充 draft KV cache,而是用 target 的中间层 hidden states 一次性投影:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
target aux hidden states [num_ctx, H_target]

│ fc 层投影到 draft hidden size

context_states [num_ctx, H_draft]

│ ① Fused GEMM(所有层的 KV projection 合成一个矩阵乘法)

all_kv_flat [num_ctx, L×2×kv_size]

│ ② Grouped RMSNorm(所有层的 K-norm 一次算完)

all_k_normed [L, num_ctx, nkv, hd]

│ ③ Fused RoPE(所有层一次应用)

all_k_final [L, num_ctx, nkv, hd] → per-layer 写入 KV cache

代码核心(DFlashQwen3Model.precompute_and_store_context_kv):

1
2
3
4
5
# 融合所有层的 KV 权重做一次大 GEMM
all_kv_flat = F.linear(normed_context_states, self._fused_kv_weight, self._fused_kv_bias)
# 分离 K/V,per-layer 写入 cache
all_kv = all_kv_flat.view(num_ctx, L, 2, nkv, hd).permute(2, 1, 0, 3, 4).contiguous()
all_k, all_v = all_kv[0], all_kv[1]

5. Probabilistic Rejection Sampling 与 Reduced Vocab

DSpark 支持 draft_sample_method="probabilistic"(Gumbel-based rejection sampling)。Draft 采样时把 logits 通过 Gumbel max trick 得到 draft_logits,Target verify 时用相同 Gumbel seed 验证,保证输出分布不变。

支持 reduced draft vocab:draft 在小词表上算 logits,然后 scatter 到 target vocab 位置:

1
2
3
4
if self._d2t_scatter_index is not None:
buf = self._draft_scatter_buf[:num_reqs] # [-inf, -inf, ...]
buf.index_copy_(1, self._d2t_scatter_index, logits_i) # 只填 draft vocab 列
logits_i = buf # 变成 target vocab 大小

6. CUDA Graph 覆盖

DFlash/DSpark 的 CUDA Graph 是 FULL mode,覆盖整个 draft step:

1
2
3
# DFlashSpeculator.init_cudagraph_manager
if wants_full and supports_full:
cudagraph_mode = CUDAGraphMode.FULL_DECODE_ONLY

为了让 Markov 循环能被 CG 捕获,所有 buffer 都是预分配的固定地址:

Buffer 用途 CG 兼容性
draft_tokens 输出 token ✅ 固定地址
draft_logits probabilistic 模式的 processed logits ✅ 固定地址
_draft_scatter_buf reduced vocab scatter buffer ✅ 固定地址
_anchor_idx 每个 request 的 anchor 位置索引 ✅ 固定地址
input_buffers.input_ids anchor token 读取 ✅ 固定地址

7. 模型加载与权重共享

load_dspark_modeldspark/utils.py)做了几件事:

  1. 创建 draft config,设置非因果注意力
  2. 加载 draft 模型
  3. Embed tokens 共享:如果 draft 没有自己的 embedding,用 target 的
  4. LM head 共享:同理
1
2
3
if _should_share(draft_model, "has_own_embed_tokens", draft_embed, target_embed):
del draft_inner.embed_tokens
draft_inner.embed_tokens = target_embed

权重加载(Qwen3DSparkForCausalLM.load_weights):

  • 跳过 t2d(训练用映射,推理不需要)
  • d2t -> draft_id_to_target_id(推理用的 draft→target 映射)
  • 跳过 mask_embedding(DSpark 通过 vocab row 做 mask,不用单独参数)和 confidence_head(未接入推理)
  • 调用 _build_fused_kv_buffers() 构建 fused KV 权重

8. 实验环境

项目 配置
Target Model Qwen/Qwen3-8B
推理引擎 vLLM v0.26.0
conda 环境 dspark-vllm
nsys 版本 2026.1.3(vLLM traces)/ 2025.3.0(DeepSpec trace)
采集参数 -t cuda,nvtx,osrt,cudnn,cublas --python-backtrace=cuda --cudabacktrace=all
Benchmark SPEED-Bench(qualitative split, coding category)

本文的实验都跑在单卡上,8B 级别的 target model + draft model 一张 80GB 卡就够。想复现这套投机解码对比的话,不必自己攒机器——RunPod 上按小时租一张 H100/H200 即可,nsys 采集需要的 --cap-add=SYS_ADMIN 权限它的容器实例也放开了。

Draft Model 配置

配置名 Draft Model 架构 来源
Baseline 无(纯 Qwen3-8B) Qwen3 -
DSpark(deepseek-ai) deepseek-ai/dspark_qwen3_8b_block7 Qwen3DSparkSt deepseek-ai 官方
DSpark(Dogacel) Dogacel/Qwen3-8B-DSpark EAGLE3 社区训练

Dogacel 的 vLLM 启动参数:speculative 开启、acceptance: 0.85num_spec_tokens: 7max_model_len: 2048


9. 性能对比

9.1 端到端性能(3 prompts, 各 64 tokens)

配置 耗时 vs Baseline 加速比
Baseline 0.75s - 1.00x
DSpark(deepseek-ai) 0.50s -33% 1.50x
DSpark(Dogacel) 0.77s +3% 0.97x

9.2 投机解码指标(DeepSpec evaluator trace)

指标 DSpark(deepseek-ai)
verify_steps 12
mean_accept_len 7.1
推测 每次提议 ~7 tokens,几乎全部被接受

Dogacel 的 trace 中未发现 dspark_propose / target_verify 的 NVTX range,推测 acceptance rate 极低。

mean_accept_len=7.1 意味着 N=8 时几乎全部接受——backbone 的并行预测质量极高,Markov head 的序列修正有效。verify_steps=12 表示 12 步验证共接受约 85 个 token(12×7.112 \times 7.1)。


10. Trace 分析

10.1 Trace 文件清单

文件 大小 来源 CUDA Kernel 数据
baseline_trace.nsys-rep 1.5 MB vLLM profile_baseline.py ❌ 无
dogacel_trace.nsys-rep 1.8 MB vLLM profile_dogacel.py ❌ 无
trace.nsys-rep(DeepSpec) 3.5 MB DeepSpec evaluator ✅ 有

10.2 CUDA Kernel 缺失原因

vLLM 的 EngineCore 在子进程中运行,nsys 默认只 trace 主进程。三个 vLLM trace 均无 GPU kernel 数据。

解决方案:重新采集时添加 --trace-fork 参数:

1
2
3
4
5
6
nsys profile -t cuda,nvtx,osrt,cudnn,cublas \
--python-backtrace=cuda --cudabacktrace=all \
--trace-fork \
--force-overwrite=true \
-o baseline_trace_v2 \
bash -c '...'

10.3 DeepSpec Evaluator Kernel 分布

Kernel 耗时占比 Instances 说明
CUTLASS GEMM (16×16) 66.3% 7,177 主要 matmul(Q/K/V/O + MLP)
elementwise_kernel 3.9% 8,522 RoPE、残差等
reduce_kernel (mean) 2.8% 4,298 RMSNorm
CUTLASS GEMM (32×32) 2.7% 382 大块矩阵乘法
Flash Attention 1.6% 864 Attention 计算
Softmax forward 1.0% 168 Softmax

关键观察:GEMM 占 66.3%,但 bs=1 decode 时本质是 memory-bound(M=1 瘦矩阵乘)。kernel launch 开销显著(~20000 次 launch)。vLLM 的 CUDA Graph 会消除大部分 launch 开销,fused kernel 会压缩 elementwise/reduce 占比。预期 vLLM 路径下 GEMM 占比升至 80%+。

10.4 NVTX Range 对比

NVTX Range Baseline Dogacel deepseek-ai dspark
dspark_propose
target_verify
decode_sample
warmup
VLLM::EngineCore

Dogacel 缺少 dspark_propose/target_verify 说明其 draft forward 未走标准 dspark 代码路径


11. Dogacel 无效原因:架构不匹配

维度 deepseek-ai(有效) Dogacel(无效)
Draft 架构 Qwen3DSparkSt EAGLE3
与 vLLM dspark 实现兼容 ✅ 完全对齐 ❌ 不匹配
NVTX range 存在 ✅ propose + verify ❌ 无
Mean accept len 7.1 推测极低
端到端加速 1.50x 0.97x(负优化)

根因:Dogacel 用 EAGLE3 架构训练 draft,中间层 hidden state 接口与 dspark 实现不兼容。即使 draft 能加载运行,acceptance rate 极低,draft 开销 > SD 收益。即使模型本身学得不差,接口不对也白搭。


12. SPEED-Bench 数据集

12.1 整体结构

SPEED-Bench(SPEculative Evaluation Dataset)是 NVIDIA 出的投机解码评测基准。

Split 样本数 用途
qualitative 880(11 类×80) 测 SD 质量(acceptance rate)
throughput_1k/2k/8k/16k/32k 1536×5 测系统吞吐(高并发)

12.2 Qualitative Split(质量评测)

从 18 个公开数据源聚合,分成 11 个 category:Coding、Math、Humanities、STEM、Writing、Summarization、Roleplay、RAG、Multilingual、Reasoning、QA。每类 80 个样本,用 OpenAI text-embedding-3-small 做嵌入,greedy 选择 + swap 优化最大化语义多样性(平均 pairwise cosine similarity 从 SpecBench 的 0.22 降到 0.14)。

12.3 Throughput Split(吞吐评测)

固定输入长度桶(1K/2K/8K/16K/32K),每桶 1536 条(512×3),分 3 个难度类别:low_entropy(coding 类)、high_entropy(creative writing 类)、mixed_entropy。用 tiktoken 精确 pad/truncate,不用 random token(会扭曲 MoE routing 和 acceptance behavior)。

12.4 为什么选 coding 类做 benchmark

  1. Coding 是低熵任务——token 可预测性高,SD 的 acceptance rate 天然高,是 best-case 场景
  2. 语义多样性好——80 条 prompt 覆盖 Python(27)、C++(9)、Java(10)、Go(13)、JS(11)、Rust(3) 等,来自 LiveCodeBench、Code Contests、HumanEvalPack
  3. 固定输出长度--speed-bench-output-len 2048)——隔离 prefill 影响,纯测 decode
  4. 两种并发对比--max-concurrency 32(batched,模拟生产环境)vs --max-concurrency 1(单流,测纯 decode 延迟)
  5. --disable-shuffle 保证可复现,--temperature 1.0 高温采样更反映真实使用场景

13. 接受率与训练效果的关系

Acceptance rate 的天花板由 draft 训练质量决定,工程实现决定能打到多少天花板。

训练侧决定上限

  • Draft 的 hidden state 和 target 的中间层对齐越好,token 分布越接近,accept 越高
  • deepseek-ai 的 block7 专门按 dspark 接口训练,hidden state 严格对齐 Qwen3-8B 第 7 层,所以 mean_accept_len=7.1
  • Dogacel 用 EAGLE3 方式训练,hidden state 映射方式不同,接口不对

工程侧决定下限

  • vLLM dspark 的 propose → verify pipeline 是否正确对接 draft
  • KV cache 的 layout、position ID 对齐、temperature sampling 一致性
  • CUDA Graph 是否覆盖 draft forward(没覆盖的话 launch overhead 会吃掉 SD 收益)
维度 deepseek-ai(训练+工程都对) Dogacel(工程接口不对)
Draft 架构 Qwen3DSparkSt EAGLE3
Hidden state 接口 ✅ 正确对接 ❌ 不匹配
NVTX range ✅ propose + verify ❌ 无
Mean accept len 7.1 推测极低
端到端加速 1.50x 0.97x

一句话总结:训练决定 draft 能不能猜对,工程决定猜对的部分能不能高效用上。Dogacel 的情况是工程接口就不对,猜得再准也走不进去。


14. vLLM 推理引擎优化对 Kernel 分布的影响

无 vLLM 优化的 kernel 分布(DeepSpec evaluator)

Kernel 占比 说明
CUTLASS GEMM (16×16) 66.3% bs=1 时是 memory-bound
elementwise 3.9% RoPE、残差等,未融合
reduce (mean) 2.8% RMSNorm,未融合
Flash Attention 1.6% decode 时计算量小
Softmax 1.0% 未融合
总 kernel launch ~20000 次 launch 开销显著

vLLM 优化后的预期变化

  1. CUDA Graph:20000 次 kernel launch → 1 次 graph launch
  2. Fused kernel:RMSNorm + residual + RoPE 融合为 1 个 kernel
  3. FlashInfer/FlashAttention decode-optimized:attention kernel 更高效
  4. GEMM 占比升至 80%+:其他开销被压缩后,GEMM 成为绝对瓶颈

对投机解码的启示

Baseline 的 decode 在 vLLM 下 GEMM 占 80%+,本质是 memory-bound(M=1 瘦矩阵乘,GPU 利用率低)。SD 的价值在于用 draft 的轻量 GEMM + target 的 batched verify 填充 GPU 空闲。当 batch size 增大(高并发),decode 从 memory-bound 转向 compute-bound,SD 收益下降——这也是 SPEED-Bench throughput split 存在的意义。


15. 总结

  1. DSpark = DFlash 并行 backbone forward + N 步 Markov 序列修正,全部在 CUDA Graph 内,用极小的开销把并行预测的「无依赖」缺陷补上
  2. 实测 deepseek-ai 官方 draft 在 Qwen3-8B 上实现 1.50x 加速mean_accept_len=7.1(N=8 几乎全接受)
  3. Dogacel 社区 draft 因架构不匹配(EAGLE3 vs DSpark)完全无效,0.97x 负优化
  4. 接受率天花板由训练决定,工程决定下限——hidden state 接口对齐是前提
  5. SPEED-Bench coding 类是 SD 的 best-case 场景,低熵任务下 acceptance rate 天然高

参考

  • vLLM 源码:vllm/v1/worker/gpu/spec_decode/dspark/speculator.pyvllm/model_executor/models/qwen3_dspark.py
  • vLLM PR:#50138#50694#50737
  • 模型:deepseek-ai/dspark_qwen3_8b_block7Dogacel/Qwen3-8B-DSpark
  • 数据集:nvidia/SPEED-Bench,arXiv: 2604.09557
  • 复现环境:RunPod 单卡实例(按小时计费,适合这类短时评测)