跳到主要内容
起居室老虎
返回

Parallelizing Linear Transformers with the Delta Rule over Sequence Length

3,215 字约 12 分钟

原文:arXiv:2406.06484 · Yang et al., NeurIPS 2024

摘要: 带线性注意力的 Transformer(即线性 Transformer)和状态空间模型(SSM)近来被认为是采用 softmax 注意力的 Transformer 的一种线性时间可行替代方案。然而,这些模型仍然不及标准 Transformer,尤其是在需要上下文内检索的任务上。尽管一种更具表达力的线性 Transformer 变体——用 Delta 规则替代线性 Transformer 中的加性更新(DeltaNet)——被发现对关联式回忆更有效,但用于训练此类模型的现有算法无法在序列长度维度并行,因此在现代硬件上训练效率低下。本文提出了一种硬件高效的算法,用于训练采用 Delta 规则的线性 Transformer。该算法利用一种用于计算 Householder 矩阵乘积的内存高效表示。借助该算法,我们得以将 DeltaNet 扩展到标准的语言建模设定。我们在 100B 词元上训练了一个 1.3B 参数模型,发现其在困惑度和下游任务零样本性能方面优于近期的线性时间基线模型,如 Mamba 和 GLA 。我们还尝试了两种混合模型:将 DeltaNet 层与(1)隔层插入的滑动窗口注意力层,或(2)两层全局注意力相结合,结果显示这些混合模型优于强力的 Transformer 基线。

1. Intro

Linear Attention 的提出主要还是希望解决掉 Transformer 中二次复杂度的问题,KV Cache 巨大显存占用的问题。通过将 softmax 注意力中的exe^x处理替换成直接的(或者其他经过变换的)点乘 Kernel,可以将 softmax 注意力的 KV Cache 转换为一固定大小的 Hidden State,实现常数内存大小需求的推理。

之前的如 SSM,直接点乘方法的 Linear Attention LLM 在长序列召回任务上不如 Transformer,但是近期的 DeltaNet 通过 Delta 规则检索并更新,在长序列召回任务上有比较大的潜力。但是,由于 DeltaNet 的写法是完全顺序的,无法跨序列长度并行,导致难以 scale up,硬件利用率低。

2. Background

2.1. Linear Transformer: Transformers with Linear Attention

给定一个dd维的输入[x1,...,xL][x_1, ..., x_L],Transformer 的注意力:

qt,kt,vt=WQxt,WKxt,WVxt,ot=∑i=1texp⁡(ki⊤qt)∑i=1texp⁡(kj⊤qt)viq_t, k_t, v_t = W_Qx_t, W_K x_t, W_V x_t,\quad o_t =\sum^t_{i=1}\frac{\exp(k_i^\top q_t)}{\sum^t_{i=1}\exp(k_j^\top q_t)}v_i

Linear Attention 将exp⁡(ki⊤qt)\exp(k_i^\top q_t)替换为ϕ(kt)⊤ϕ(qt)\phi(k_t)^\top\phi(q_t),ϕ:ℝd→ℝn\phi:\mathbb R^d\rightarrow \mathbb R^n的一个核函数。这意味着计算可以重新排序为:

ot=∑i=1tϕ(ki)⊤ϕ(qt)∑i=1tϕ(kj⊤)ϕ(qt)vi=(∑i=1tviϕ(ki)⊤)ϕ(qt)(∑j=1tϕ(kj)⊤)ϕ(qt)=Stϕ(qt)zt⊤ϕ(qt),o_t=\sum_{i=1}^t\frac{\phi(k_i)^\top\phi(q_t)}{\sum_i=1^t\phi (k_j^\top)\phi(q_t)}v_i=\frac{\left(\sum_{i=1}^tv_i\phi(k_i)^\top\right)\phi(q_t)}{\left(\sum^t_{j=1}\phi(k_j)^\top\right)\phi(q_t)}=\frac{S_t\phi(q_t)}{z_t^\top\phi(q_t)},

其中St=∑i=1tviϕ(ki)⊤∈ℝd×n,zt=∑i=1tϕ(ki)∈ℝnS_t=\sum^t_{i=1}v_i\phi(k_i)^\top\in\mathbb R^{d\times n}, z_t=\sum^t_{i=1}\phi(k_i)\in \mathbb R^n.

当n→∞n\rightarrow\infin时,Linear Attn 选用多项式相关的ϕ(⋅)\phi(\cdot)理论上可以从任意精度逼近 softmax 注意力。一些工作发现zt⊤ϕ(qt)∈ℝz_t^\top\phi(q_t)\in\mathbb R存在数值不稳定的情形,因此将它去除。一种简化的 Linear Transformer:

St=St−1+vtkt⊤,ot=StqtS_t=S_{t-1}+v_tk_t^\top, \quad o_t=S_tq_t

Efficient training

为了实现上面的高效训练,假设Q,K,V∈ℝL×dQ, K, V\in \mathbb R^{L\times d}是堆叠后的q,k,vq, k, v,则可以并行计算O∈ℝL×dO\in \mathbb R^{L\times d}:

O=(QKT⊙ML)VO=(QK^T\odot M_L)V

ML∈ℝL×LM_L\in\mathbb R^{L\times L}是因果 mask。这个模式和前面的递归模式各有取舍:并行模式需要O(L2d+Ld2)O(L^2d+Ld^2)的 FLOPs(多算了一个QK⊤QK^\top,实际上递推的时候只要算LL次d2d^2的隐状态更新),但是可以把 GPU 吃满;递归模式只需要O(Ld2)O(Ld^2),但是无法跨序列长度进行并行,难以利用 Tensor Core 进行加速。

Chunkwise parallel form

为了在上面两个状态中进行权衡,可以进行分块并行。将Q,K,VQ, K, V按照长度CC切分为LCL_C个 chunk。记第tt个块中Q[t]∈ℝC×dQ[t]\in \mathbb R^{C\times d},状态可以重新标记为:Si[t]=StC+1S_i[t]=S_{tC+1},并定义S0[t]=SC[t−1]S_0[t]=S_C[t-1],即初始状态是上一个块的末状态,则:

Sr[t]=S0[t]+∑i=1tvi[t]ki⊤[t],or[t]=S0[t]qr[t]+∑i=1rvi[t](ki⊤[t]qr[t])S_r[t]=S_0[t]+\sum_{i=1}^tv_i[t]k_i^\top[t], \quad o_r[t]=S_0[t]q_r[t]+\sum^r_{i=1}v_i[t](k^\top_i[t] q_r[t])

记S[t]=S0[t]S[t]=S_0[t],块内改写为并行形式:

S[t+1]=S[t]+V⊤[t]K[t]∈ℝd×d,O[t]=Q[t]S⊤[t]+(Q[t]K⊤[t]⊙MC)V[t]S[t+1]=S[t]+V^\top[t]K[t]\in \mathbb R^{d\times d}, \\O[t]=Q[t]S^\top[t]+(Q[t]K^\top[t]\odot M_C)V[t]

这样跨块之间传递的只包括S[t]S[t],中间状态Si[t]S_i[t]不用落盘,复杂度变为O(LCd+Ld2)O(LCd + Ld^2)做O(L/C)O(L/C)步。

2.2. DeltaNet: Linear Transformers with the Delta Update Rule

注意到上面的 Linear Transformer 采用线性递推:

St=St−1+vtkt⊤S_t=S_{t-1}+v_tk^\top_t

纯加性地将新的k−vk-v对写入记忆中。但是这种记忆模式难以 “回收” 过去的关联,在L>dL>d的时候容易出现 key collision 的现象。理想的模型应该能够移除掉不重要的关联,为新的信息腾出空间。

DeltaNet 的 Delta Rule:

St=St−1−βt(St−1kt−vt)kt⊤S_t=S_{t-1}-\beta_t(S_{t-1}k_t-v_t)k_t^\top

其中β\beta是学习绿,St−1ktS_{t-1}k_t是当前预测,vtv_t是目标 value。核心是根据预测与目标之间的 delta 来更新权重,这一过程还可以视为对 online regression loss 做单步 SGD 优化:

ℒt(S)=12||Skt−vt||2,St=St−1−βt∇St−1ℒt(St−1)=St−1−βt(St−1kt−vt)kt⊤\mathcal L_t(S)=\frac{1}{2}||Sk_t-v_t||^2,\\S_t=S_{t-1}-\beta_t\nabla_{S_{t-1}}\mathcal L_t(S_{t-1})=S_{t-1}-\beta_t(S_{t-1}k_t-v_t)k_t^\top

相对的,普通的 Linear Attention(上面的加性版本)则是 online linear (negative inner-product) lossℒt=−⟨Skt,vt⟩3\mathcal L_t=-\langle Sk_t, v_t\rangle^3进行优化。

从 KV 检索的角度来理解,可以看成先用当前的 key 获得旧的 value:

vtold=St−1ktv_t^\text{old}=S_{t-1}k_t

然后将新旧一起进行差值:

vtnew=βtvt+(1−βt)vtoldv_t^\text{new}=\beta_tv_t+(1-\beta_t)v_t^\text{old}

然后根据这个移除旧的、写入新的:

St=St−1−vtoldkt⊤+vtnewkt⊤S_t=S_{t-1}-v_t^\text{old}k_t^\top + v_t^\text{new}k_t^\top

Schlag 等人证明,在小规模 语言建模与合成 的上下文检索任务上,DeltaNet 优于普通线性 Transformer。然而,他们基于线性 Transformer 内存高效递归实现 的训练算法是严格顺序 的,正如 (下文 3.2) 所指出的那样,对现代硬件并不友好。这促使我们在下文给出一个等价的分块算法,以便在更大规模上训练 DeltaNet。

3. Paralelizingn DeltaNet Across the Sequence Dimension

3.1. A Memory-efficient Reparameterization

首先注意到,StS_t也是可以写成加性的形式的:

St=∑i=1tuiki⊤,ui=vinew−viold=βi(vi−viold)S_t=\sum^t_{i=1}u_ik_i^\top, \quad u_i = v_i^{\text{new}}-v_i^{\text{old}}=\beta_i\,(v_i - v_i^{\text{old}})

因此如果我们能够构造出所有的uiu_i,就能够

O=(QK⊤⊙M)U,U=stackrow(u1,…,uL)O=(QK^\top \odot M)\,U,\quad U=\text{stack}_{\text{row}}(u_1,\dots,u_L)

但是,如果我们朴素地计算utu_t,需要计算每一个St−1S_{t-1}来得到vtoldv_t^\text{old},需要O(d2)O(d^2)的内存。下面通过归纳法,利用 Householder 矩阵乘积证明,实际上我们只需要O(d)O(d)的内存。

显然有:

S1=β1v1k1⊤⟹u1=β1v1S_1=\beta_1v_1k_1^\top\implies u_1=\beta_1v_1

从 DeltaNet 本身出发:

St=St−1−vtoldkt⊤+vtnewkt⊤=St−1−βt(St−1kt)kt⊤+βtvtkt⊤=St−1(I−βktkt⊤)+βvtkt⊤\begin{align*} S_t&=S_{t-1}-v_t^\text{old}k_t^\top+v_t^\text{new}k_t^\top\\&=S_{t-1}-\beta_t(S_{t-1}k_t)k^\top_t+\beta_tv_tk_t^\top\\&=S_{t-1}(I-\beta k_tk^\top_t)+\beta v_tk_t^\top \end{align*}

注意到I−βktkt⊤I-\beta k_tk_t^\top是一广义 Householder 变换。∏j(I−βjkjkj⊤)\prod_j(I-\beta_jk_jk_j^\top)可以写作W−YW-Y表示:Pr=I−∑i=1rwiki⊤P_r=I-\sum_{i=1}^rw_ik_i^\top,其中wiw_i由前面的{w<i,k<i}\{w_{<i}, k_{<i}\}和当前的(βi,ki)(\beta_i, k_i)生成。

归纳展开得到:

St=∑i=1t−1uiki⊤+βt(vt−∑i=1t−1ui(ki⊤kt))kt⊤S_t=\sum_{i=1}^{t-1} u_i k_i^\top \;+\; \beta_t\!\left(v_t - \sum_{i=1}^{t-1} u_i\,(k_i^\top k_t)\right)\!k_t^\top

从而:

ut≜βt(vt−∑i=1t−1ui(ki⊤kt)),St=∑i=1tuiki⊤u_t \triangleq \beta_t\!\left(v_t - \sum_{i=1}^{t-1} u_i\,(k_i^\top k_t)\right),\quad S_t=\sum_{i=1}^{t} u_i k_i^\top

得到了不需要构造St−1S_{t-1}也可以求得的utu_t,只需要O(d)O(d)内存。具体而言,

𝐒n=𝐒n−1(𝐈−βn𝒌n𝒌n⊤)+βn𝒗n𝒌n⊤=(∑t=1n−1𝒖t𝒌t⊤)(𝐈−βn𝒌n𝒌n⊤)+βn𝒗n𝒌n⊤=∑t=1n−1𝒖t𝒌t⊤−(∑t=1n−1𝒖t𝒌t⊤)βn𝒌n𝒌n⊤+βn𝒗n𝒌n⊤=∑t=1n−1𝒖t𝒌t⊤+(βn𝒗n−βn∑t=1n−1𝒖t(𝒌t⊤𝒌n))⏟𝒖n𝒌n⊤=∑t=1n𝒖t𝒌n⊤\begin{aligned}\mathbf{S}_{n} & =\mathbf{S}_{n-1}\left(\mathbf{I}-\beta_{n} \boldsymbol{k}_{n} \boldsymbol{k}_{n}^{\top}\right)+\beta_{n} \boldsymbol{v}_{n} \boldsymbol{k}_{n}^{\top} \\& =\left(\sum_{t=1}^{n-1} \boldsymbol{u}_{t} \boldsymbol{k}_{t}^{\top}\right)\left(\mathbf{I}-\beta_{n} \boldsymbol{k}_{n} \boldsymbol{k}_{n}^{\top}\right)+\beta_{n} \boldsymbol{v}_{n} \boldsymbol{k}_{n}^{\top} \\& =\sum_{t=1}^{n-1} \boldsymbol{u}_{t} \boldsymbol{k}_{t}^{\top}-\left(\sum_{t=1}^{n-1} \boldsymbol{u}_{t} \boldsymbol{k}_{t}^{\top}\right) \beta_{n} \boldsymbol{k}_{n} \boldsymbol{k}_{n}^{\top}+\beta_{n} \boldsymbol{v}_{n} \boldsymbol{k}_{n}^{\top} \\& =\sum_{t=1}^{n-1} \boldsymbol{u}_{t} \boldsymbol{k}_{t}^{\top}+\underbrace{\left(\beta_{n} \boldsymbol{v}_{n}-\beta_{n} \sum_{t=1}^{n-1} \boldsymbol{u}_{t}\left(\boldsymbol{k}_{t}^{\top} \boldsymbol{k}_{n}\right)\right)}_{\boldsymbol{u}_{n}} \boldsymbol{k}_{n}^{\top} \\& =\sum_{t=1}^{n} \boldsymbol{u}_{t} \boldsymbol{k}_{n}^{\top}\end{aligned}

然而,直接计算所有utu_t需要O(L2d)O(L^2d)并且无法并行,因此还需要:

3.2. Chunkwise Parallel Form for DeltaNet

首先将递推式展开:

St=St−1(I−βtktkt⊤)+βtvtkt⊤=∑i=1tβi(viki⊤)(∏j=i+1t(I−βjkjkj⊤))S_t= S_{t-1}\bigl(I-\beta_t k_t k_t^\top\bigr) + \beta_t v_t k_t^\top= \sum_{i=1}^{t}\beta_i\,(v_i k_i^\top)\left(\prod_{j=i+1}^{t}(I-\beta_j k_j k_j^\top)\right)

定义:

Pij=∏t=ij(I−βtktkt⊤),Hij=∏t=ijβt(vtkt⊤)Pt+1jP_i^j=\prod_{t=i}^{j}\bigl(I-\beta_t k_t k_t^\top\bigr),\qquad H_i^j=\prod_{t=i}^{j}\beta_t\,(v_t k_t^\top)\,P_{t+1}^j

广义 Householder,化简:

𝐏n=∏t=1n(𝐈−βt𝒌t𝒌t⊤)=𝐏n−1(𝐈−βn𝒌n𝒌n⊤)=(𝐈−∑t=1n−1𝒘t𝒌t⊤)(𝐈−βn𝒌n𝒌n⊤)=𝐈−∑t=1n−1𝒘t𝒌t⊤−βn𝒌n𝒌n⊤+(∑t=1n−1𝒘t𝒌t⊤)βn𝒌n𝒌n⊤=𝐈−∑t=1n−1𝒘t𝒌t⊤−(βn𝒌n−βn∑t=1n−1(𝒘t(𝒌t⊤𝒌n)))⏟𝒘n𝒌n⊤=𝐈−∑t=1n𝒘t𝒌t⊤\begin{aligned}\mathbf{P}_{n} & =\prod_{t=1}^{n}\left(\mathbf{I}-\beta_{t} \boldsymbol{k}_{t} \boldsymbol{k}_{t}^{\top}\right) \\& =\mathbf{P}_{n-1}\left(\mathbf{I}-\beta_{n} \boldsymbol{k}_{n} \boldsymbol{k}_{n}^{\top}\right) \\& =\left(\mathbf{I}-\sum_{t=1}^{n-1} \boldsymbol{w}_{t} \boldsymbol{k}_{t}^{\top}\right)\left(\mathbf{I}-\beta_{n} \boldsymbol{k}_{n} \boldsymbol{k}_{n}^{\top}\right) \\& =\mathbf{I}-\sum_{t=1}^{n-1} \boldsymbol{w}_{t} \boldsymbol{k}_{t}^{\top}-\beta_{n} \boldsymbol{k}_{n} \boldsymbol{k}_{n}^{\top}+\left(\sum_{t=1}^{n-1} \boldsymbol{w}_{t} \boldsymbol{k}_{t}^{\top}\right) \beta_{n} \boldsymbol{k}_{n} \boldsymbol{k}_{n}^{\top} \\& =\mathbf{I}-\sum_{t=1}^{n-1} \boldsymbol{w}_{t} \boldsymbol{k}_{t}^{\top}-\underbrace{\left(\beta_{n} \boldsymbol{k}_{n}-\beta_{n} \sum_{t=1}^{n-1}\left(\boldsymbol{w}_{t}\left(\boldsymbol{k}_{t}^{\top} \boldsymbol{k}_{n}\right)\right)\right)}_{\boldsymbol{w}_{n}} \boldsymbol{k}_{n}^{\top} \\& =\mathbf{I}-\sum_{t=1}^{n} \boldsymbol{w}_{t} \boldsymbol{k}_{t}^{\top}\end{aligned}

分块:

Si[t]=StC+i,Pr[t]=PtC+1tC+r,Hr[t]=HtC+1tC+r S_i[t]=S_{tC+i},\quad P_r[t]=P_{tC+1}^{\,tC+r},\quad H_r[t]=H_{tC+1}^{\,tC+r}

块内递推:

Sr[t]=S0[t],Pr[t]+Hr[t]S_r[t]=S_0[t],\qquad P_r[t]+H_r[t]

与 3.1 中类似地,

Pr[t]=I−∑i=1rwi[t]ki[t]⊤,Hr[t]=∑i=1rui[t]ki[t]⊤P_r[t]=I-\sum_{i=1}^{r} w_i[t]\,k_i[t]^\top,\qquad H_r[t]=\sum_{i=1}^{r} u_i[t]\,k_i[t]^\top

且wr[t],ur[t]∈ℝdw_r[t], u_r[t]\in \mathbb R^d的递推:

wr[t]=βr[t](kr[t]−∑i=1r−1wi[t](ki[t]⊤kr[t])),ur[t]=βr[t](vr[t]−∑i=1r−1ui[t](ki[t]⊤kr[t]))w_r[t]=\beta_r[t]\!\left(k_r[t]-\sum_{i=1}^{r-1} w_i[t]\,(k_i[t]^\top k_r[t])\right),\\u_r[t]=\beta_r[t]\!\left(v_r[t]-\sum_{i=1}^{r-1} u_i[t]\,(k_i[t]^\top k_r[t])\right)

故:

Sr[t]=S0[t]−S0[t](∑i=1rwi[t]ki[t]⊤)+∑i=1rui[t]ki[t]⊤=S0[t]+∑i=1r(ui[t]−S0[t]wi[t])ki[t]⊤or[t]=Sr[t]qr[t]=S0[t]qr[t]+∑i=1r(ui[t]−S0[t]wi[t])(ki[t]⊤qi[t])S_r[t]= S_0[t]-S_0[t]\!\left(\sum_{i=1}^{r} w_i[t]\,k_i[t]^\top\right)+\sum_{i=1}^{r} u_i[t]\,k_i[t]^\top= S_0[t]+\sum_{i=1}^{r}\Bigl(u_i[t]-S_0[t]\,w_i[t]\Bigr)\,k_i[t]^\top\\o_r[t]=S_r[t]\,q_r[t] = S_0[t]\,q_r[t]+\sum_{i=1}^{r}\Bigl(u_i[t]-S_0[t]\,w_i[t]\Bigr)\,\bigl(k_i[t]^\top q_i[t]\bigr)

同样地,将S[t]=S0[t]S[t]=S_0[t],有:

S[t+1]=S[t]+(U[t]−W[t]S[t]⊤)⊤K[t]O[t]=Q[t]S[t]⊤+(Q[t]K[t]⊤⊙M)(U[t]−W[t]S[t]⊤)S[t+1]=S[t]+\Bigl(U[t]-W[t]\,S[t]^\top\Bigr)^\top K[t]\\O[t]=Q[t]\,S[t]^\top+\bigl(Q[t]K[t]^\top\odot M\bigr)\,\Bigl(U[t]-W[t]\,S[t]^\top\Bigr)

Practical considerations

考虑到块内的or[t]o_r[t]的写法仍然是递推的,难以高效利用 Tensor Core。引入 UT 变换11:

T[t]=(I+tril⁡(diag⁡(β[t])K[t]K[t]⊤,−1))−1diag⁡(β[t])W[t]=T[t]K[t],U[t]=T[t]V[t]T[t]=\Bigl(I+\operatorname{tril}\bigl(\operatorname{diag}(\beta[t])\,K[t]K[t]^\top,\,-1\bigr)\Bigr)^{-1}\operatorname{diag}(\beta[t])\\W[t]=T[t]\,K[t],\quad U[t]=T[t]\,V[t]

从而将大多数运算改写为 matmul,使得状态更新得到了几乎与一般 Linear Attention 一样的计算流程。在 BP 中,通过重计算 Hidden State 节省显存。

Speed comparison

image.png

基于 Triton 实现了纯递归和 Chunkwise 两种,Chunckwise 由于可以更好地利用 Tensor Core 往往能有很多倍的提升。

Fully Parallel Form for DeltaNet

注意力矩阵:

Aij={kj⊤Pj+1iqi,j≤i,0,j>i,A_{ij}=\begin{cases}k_j^\top\,P_{j+1}^{\,i}\,q_i,& j\le i,\\0,& j>i,\end{cases}

同样可以写成A=(QK⊤⊙M)TA=(QK^\top\odot M)\,T的形式,但是计算TT需要对上面的 UT 变换得到的内容求逆,如果写成完全并行的模式复杂度变成三次方,因此在训练中不采用完全并行形式;但该 “注意力” 矩阵对 RNN 可解释性研究可能有用。

3.3. DeltaNet Transformer

image.png

结构类似 LLaMA/Transformer++,将自注意力的地方换成上面的 DeltaNet,输入前添加 RMSNorm 增强训练稳定性。参数和 Transformer++ 类似,DeltaNet 层约4d24d^2,SwiGLU 层约8d28d^2。

Feature map and normalization

将 key & query 定义为:

kt=SiLU⁡(WKxt)‖SiLU⁡(WKxt)‖2,qt=SiLU⁡(WQxt)‖SiLU⁡(WQxt)‖2k_t=\frac{\operatorname{SiLU}(W_K x_t)}{\left\|\operatorname{SiLU}(W_K x_t)\right\|_2},\quad q_t=\frac{\operatorname{SiLU}(W_Q x_t)}{\left\|\operatorname{SiLU}(W_Q x_t)\right\|_2}

为了稳定,需要保证状态转移矩阵的特征值模长不超过 1。对于:

I−βtktkt⊤I-\beta_tk_tk_t^\top

有d−1d-1个特征值为11,11个为1−βt||kt||221-\beta_t||k_t||^2_2。本文采用 L2 归一化:当βt=1\beta_t=1的时候I−ktkt⊤I-k_tk_t^\top变成投影矩阵,在一个子空间中 “抹除” 信息,同时保留其他d−1d-1个子空间的信息。

3.4. Hybrid Models

Convolutional layers

在Q/K/VQ/K/V矩阵投影之后接一个比较小的 Depthwise separable 卷积,本文采用。

Local sliding window and global attention

Linear Attention 比较依赖于 “content-based addressing”,缺乏 ”positional information”,因此在强检索(retrieval-intensive)任务上受限。因此:

4. Empirical Study

4.1. Synthetic Benchmark

image.png

image.png

4.2. Language Modeling

image.png

image.png

Ablations

主要是 L1 Norm vs. L2 Norm, ReLU vs. 1+ELU

Training throughput

image.png

长序列上能超过 Transformer++,短序列应该是因为 Tensor Core 用不满?Training 确实是 Compute Bound。

5.1. DeltaNet vs. State Space Models/Linear RNNs

考虑一类 “具有 Hidden State Matrix、assiciative“的 RNN,可以写作;

St=St−1 ⋅ Mt+vtkt⊤(recurrence)ot=Stqt(memory read-out)S_t = S_{t-1}\ \boldsymbol{\cdot}\ M_t + v_t k_t^\top\qquad\text{(recurrence)}\\ o_t = S_t q_t\qquad \text{(memory read-out)}

其中 ⋅\boldsymbol\cdot 是一个可结合的算子(矩阵乘,Hadamard 乘法,…)。

这种 RNN 可以通过并行扫描的方法,以O(log⁡L)O(\log L)步,O(L)O(L)工作量计算[S1,...,SL][S_1, ..., S_L]。因此,只要 ⋅\boldsymbol \cdot 本身的开销不打,训练就会是高效的。然而,很多 ⋅\boldsymbol \cdot 就是有很大的计算量,最近的工作如 Mamba, Gated Linear Attention Transformer 就用⊙\odot逐元素乘法作为 ⋅\boldsymbol \cdot 。

另一方面,matmul 的表达能力当然比 ⊙\odot 强,但是如果对MtM_t不施加任何限制/结构先验,每步的更新量从之前的O(dn)O(dn)变成O(dn2)O(dn^2),代价过高。DeltaNet 采用

Mt=I−βtktkt⊤M_t=I-\beta_tk_tk_t^\top

相当于引入了结构化的先验,是上面两者的一个折中。

本文提出的 Chunkwise 算法还可以推广到更加一般的 “对角加低秩 “(Diagnoal-Plus-Low-Rank)形式:

Mt=D−atbt⊤M_t=D-a_tb_t^\top

S4 中讨论过这个形式,但它的做法让MtM_t不是 input dependent 的22。本工作则需要 input dependent。

image.png

5.2. Towards a Unifying Framework for Efficient Autoregressive Sequence Transformations

尽管上述类模型有助于统一 近期方法,我们并不声称它就是观察自回归序列变换的 “唯一正确层次 ”。

序列变换形如:

{xt}t=1L→{ot}t=1L\{x_t\}^L_{t=1}\to\{o_t\}^L_{t=1}

且oto_t不得依赖于xj,j>tx_j, j>t(就是因果)。

例如,这一框架不易 整齐地涵盖其他已被证明有效的次二次 方法。另一种统一思路是把上述序列变换看作连续状态空间模型 的离散化,或视作与掩码结构矩阵 相乘。 更关键的是,一个好的框架应当能催生高效的训练算法,并且对硬件友好——在现代 GPU 上,这通常意味着富含矩阵乘。

5.3. Limitations and Future Work

本工作仍有若干局限。首先在计算 方面,尽管我们提出了新的硬件高效算法,其训练速度仍落后 于 GLA。这是因为我们在内核中对状态间依赖 的建模引入了开销,需要在 head 维度上做 “边缘化 ”,这与 softmax 注意力 的情形类似;而对 GLA 而言,不存在状态内依赖(全为逐元素),因此易于通过 Tiling 支持任意 head 维度 。这一限制可能会限制 DeltaNet 的状态容量,从而降低强召回任务的表现(与 §4.2 的观察一致)。一种潜在改进是采用块对角的广义 Householder 过渡矩阵,使块大小适配 GPU SRAM(如 128),在保持整体较大 head 维度(即较大循环状态)的同时降低内核负担。 其次,我们发现 DeltaNet 的长度泛化 有限;相对地,GLA 与 RetNet(以及一定程度上的 Mamba)在外推长于训练序列 方面更好 。我们推测这是因为 DeltaNet 缺少显式衰减因子。可考虑在递推中引入门控项 进行改进。

线性 Transformer 可被视为一种迭代 Hopfield 网络;这一联系有助于理解线性注意力的局限与改进方向。譬如,朴素线性注意力 使用类似 Hebbian 的更新,已被证明容量有限。后续 Hopfield 工作通过高阶多项式与指数核提升记忆容量,这也与多项式核的线性注意力相关。另一方面,Delta Rule 被证明具有更好的记忆容量。在固定大小的循环状态下,Delta Rule 能在 “召回‑记忆” 权衡上取得更优前沿,并已用于增强现实世界的检索任务;它在多个领域也优于线性 Transformer 的加性规则。尽管如此,Irie 等指出 delta 更新在表达力上存在理论局限。

普通的 Linear RNN 中,维护一个矩阵态SS:

S←S+vk⊤,o=Sq,S \leftarrow S + v\,k^\top,\qquad o = S\,q,

如果两个kk不正交,会出现碰撞的情况:

o=vi⟨ki,q⟩⏟想要+∑j≠ivj⟨kj,q⟩噪声.o \;=\; \underbrace{v_i\langle k_i,q\rangle}_{\text{想要}} \;+\; \sum_{j\neq i} v_j\langle k_j,q\rangle_{\text{噪声}}.

Delta Rule 做一次 SGD,得到:

S←S−β(Sk−v)k⊤,S \;\leftarrow\; S \;-\; \beta\,(S k - v)\,k^\top,

能够收敛到最小二乘解:

S⋆=arg⁡minS⁡∑i‖Ski−vi‖2=VK⊤(KK⊤)−1(KK⊤ 可逆时),S^\star \;=\; \arg\min_S \sum_i \| S k_i - v_i\|^2\;=\; V K^\top (K K^\top)^{-1} \quad (\text{$K K^\top$ 可逆时}),

而原版 softmax 注意力:

A=softmax(QK⊤dk),O=AV.A \;=\; \mathrm{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}}\right),\qquad O \;=\; A\,V.

直接显示保留了整个K,VK, V作为记忆,记忆能力非常强但n2n^2开销。

为增强 DeltaNet 的递归能力,提出了 Recurrent DeltaNet、Modern Self‑Referential Weight Matrix、mesa‑layer 等,并显示更优。但这些模型超出了线性 RNN 范畴,无法跨序列并行,提示存在并行性与表达力的根本权衡。如何在不牺牲并行性的前提下进一步增强 DeltaNet 仍是开放问题;TTT 采用 “跨块非线性 + 块内线性” 的混合策略,或许能提供折中路径。最后,Delta Rule 与经由梯度下降的元/在线学习密切相关,近期如 Longhorn、TTT 予以重访;Titans进一步加入动量与权重衰减。

7. Conclusion

我们给出了一个沿序列长度维度 并行化 DeltaNet 训练的算法,在现代硬件上相较既有实现取得了显著加速。这使得将 DeltaNet 扩展到中等规模语言建模成为可能;在该设定下,我们发现其相较近期的线性递归基线表现良好。

对于一个长度为CC的块内,定义K[t]∈ℝC×d,V[t]K[t]\in \mathbb R^{C\times d}, V[t]同理,定义:

T[t]=(I+tril(diag(β[t])K[t]K[t]⊤,−1))−1diag(β[t])𝐖[t]=𝐓[t]𝐊[t]𝐔[t]=𝐓[t]𝐕[t]T[t] = (I + \text{tril}(\text{diag}(\beta[t]) K[t] K[t]^\top,-1))^{-1}\text{diag}(\beta[t])\\\begin{aligned} \mathbf{W}_{[t]} & =\mathbf{T}_{[t]} \mathbf{K}_{[t]} \\ \mathbf{U}_{[t]} & =\mathbf{T}_{[t]} \mathbf{V}_{[t]} \end{aligned}

那么就有:

𝐖[t][r,:]=β[t]r𝐊[t][r,:]−β[t]r∑i=1r−1𝐖[t][i,:](𝐊[t][i,:]𝐊[t][r,:]⊤)\mathbf{W}_{[t]}[r,:]=\beta_{[t]}^{r} \mathbf{K}_{[t]}[r,:]-\beta_{[t]}^{r} \sum_{i=1}^{r-1} \mathbf{W}_{[t]}[i,:]\left(\mathbf{K}_{[t]}[i,:] \mathbf{K}_{[t]}[r,:]^{\top}\right)

定义:

𝐁[t]=diag⁡(β[t])𝐋[t]=tril⁡(𝐁[t]𝐊[t]𝐊[t]⊤,−1)\begin{array}{c}\mathbf{B}_{[t]}=\operatorname{diag}\left(\beta_{[t]}\right) \\\mathbf{L}_{[t]}=\operatorname{tril}\left(\mathbf{B}_{[t]} \mathbf{K}_{[t]} \mathbf{K}_{[t]}^{\top},-1\right)\end{array}

则:

𝐖[t]+𝐋[t]𝐖[t]=𝐁[t]𝐊[t]\mathbf{W}_{[t]}+\mathbf{L}_{[t]} \mathbf{W}_{[t]}=\mathbf{B}_{[t]} \mathbf{K}_{[t]}

因此:

𝐖[t]=(𝐈+𝐋[t])−1𝐁[t]𝐊[t]=𝐓[t]𝐊[t]𝐓[t]=(𝐈+𝐋[t])−1𝐁[t]\mathbf{W}_{[t]}=\left(\mathbf{I}+\mathbf{L}_{[t]}\right)^{-1} \mathbf{B}_{[t]} \mathbf{K}_{[t]}=\mathbf{T}_{[t]} \mathbf{K}_{[t]} \qquad \mathbf{T}_{[t]}=\left(\mathbf{I}+\mathbf{L}_{[t]}\right)^{-1} \mathbf{B}_{[t]}

另一个也同理。

U[t]=T[t]V[t],W[t]=T[t]K[t]U[t]=T[t]V[t], \qquad W[t]=T[t]K[t]

等价于一次把ur=βr(vr−∑i<rui(ki⊤kr))u_r=\beta_r(v_r-\sum_{i<r}u_i\left(k_i^\top k_r)\right)中所有的ur,wru_r, w_r递推,一次性通过单位下三角阵 - 对角缩放求出来。

脚注

  1. 将逐元素递推写成一次下三角线性方程的求解。 ↩

  2. 所以 S6 就改了。 ↩