引言
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
期望预测的与真实的的误差尽可能的小,定义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 短序列场景下,缓存和计算未必占优
这和前文Linear Attention系列解读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的输出维度等,读者有兴趣可以详细阅读原文。
如有疏漏之处,欢迎指出~
- 作者:莫叶何竹🍀
- 链接:http://www.myhz0606.com/article/linear_attention_p2?_sm_nck=1
- 声明:本文采用 CC BY-NC-SA 4.0 许可协议,转载请注明出处。
相关文章










.png?table=block&id=1f63c18f-f81c-808d-8829-cd2c12c40b3e&t=1f63c18f-f81c-808d-8829-cd2c12c40b3e)