引言

Vanilla Linear Attention将softmax kernel替换为可分解kernel
随后利用矩阵乘法结合律,将自回归推理中随序列增长的KV Cache压缩为固定大小的累积状态 ,其递归形式为:
,用的是行向量的记法
正如Linear Attention系列解读1 所言,累积外积和的方式虽然解决了softmax-based attention在长序列面临的随序列线性增长的显存承压,但是存在以下问题:
  • 短序列缓存计算未必占优
  • 长序列表达受限。主要是因为,对于vanilla linear attention直接用累加的方法更新memories。随着context增大,新旧association不断叠,使得不同信息之间的crosstalk不断增大,最终导致memory collision。同时,模型缺少修改或删除旧 association 的机制,因此固定容量 memory 的利用效率较低。
基于以上考虑,DeltaNet给原始的linear attention设计了一种error-correcting delta-rule (Widrow & Hoff, 1960)的增量更新机制来代替key,value外积和的累积更新,以此提升其在长序列的表达能力。
注意:
为了方便,先把这个kernel function 记为identity。即
时刻的linear attention输出

DeltaNet具体方法

Delta-rule的核心思想:
delta-rule的思路分为以下几个步骤:注意后面的公式采用row-vector convention,并假定key,value的维度一致,即
在时刻, 输入通过投影可以得到当前时刻的
首先从历史的memories,找到当前相关的,并作weight sum,记作,计算方式如下:
可以理解为memories中所存储的与相关的信息。
随后通过线性插值的方式融合当前时刻的和历史相关的,将其称为
这个也称为writing strength。计算方式为, 为sigmoid函数,因此能保证
  • :保留已有映射;
  • :完全采用当前 value;
  • :部分修正。
同时蕴含了当前新增的信息,和历史已有的信息,但memories依旧包含了的相关的信息,因此不需要写入完整的,只需写入残差,即
可见,delta-rule更新的是信息的增量
整理一下,上式也可以进一步写作:
上面就是delta-rule的核心思路。相比vanilla的累积和,deltanet做的是增量更新。

Delta Rule可以视作Online SGD

这个视角源于Fast weight programming (FWP)。核心是,将memories 视作是fast weight(linear model),输入,预测value
期望预测的与真实的的误差尽可能的小,定义L2 loss:
两边对求gradient
假定learning rate为,那么根据SGD的更新策略:
这一项正好是delta-rule memories的更新方式。
因此,从这个视角看,delta rule的memory看成一个从key到value的映射。写入新的key-value association时,先读取memory对当前key的已有预测,再只写入目标value与已有预测之间的差值。

DeltaNet的并行化初探

根据式8,令,
对每一个时间步,都有pair
从式13可见,这是经典的仿射递推形式(affine recurrence)。
考虑两个时间步
代入得到
定义
类比14,展开不难发现
从上式可见,只要能并行计算就能并行计算
下面来看,能否并行计算
定义combine运算
显然这个运算满足结合律。证明如下:
左边:
右边:
显然左边等于右边,因此满足结合律。
因此可以通过parallel scan计算

DeltaNet的局限性

1 短序列场景下,缓存和计算未必占优
cache
元素个数
Linear Attention
Softmax Attention
临界点
从中可见,若序列较短(约小于)时,linear attention需要的cache反而更多。
只考虑单个attention head。key value的维度相同
2 长序列的表达依旧受限
相比vanilla linear attention,deltanet明显解决了memories的利用率问题。但是,它依然无法提升memories的容量。因此在长序列下,依然存在memories collision。
3 parallel scan并行化io开销太大
虽然单步的都具有rank-one structure,无须直接构建为dense matrix。但在 naive parallel scan中,经过多次affine composition得到的通常是dense的。如果为每个prefix显式保存这两个矩阵,会产生 的memory I/O;因此,DeltaNet 虽然在数学上可以使用parallel scan,但naive scan的方式会有较大的I/O瓶颈。后续工作《Parallelizing Linear Transformers with the Delta Rule over Sequence Length》对此做了深入优化。

小结

本文相对系统讨论了linear attention中的deltanet变体。本文着重介绍了该方法的核心motivation和insight:将固定state看成一个从key到value的fast-weight model,并根据prediction error对已有映射做增量修正。文本有一些细节并没有详细阐述,如q/k的额外做了sum normalization,通过DPFP扩大kernel function的输出维度等,读者有兴趣可以详细阅读原文。
如有疏漏之处,欢迎指出~
 
Kimi K3 技术解析diffusion model(一):DDPM技术小结 (denoising diffusion probabilistic)
Loading...