引言

经典的Transformer架构,在目前的大模型时代取得了巨大成功。但其核心组件Self-Attention的时空复杂度都为为输入序列的长度。
Flash-attenion技术利用online-softmax,分块计算等技巧,将标准的self-attention的空间复杂度降低到,时间复杂度仍是,大幅降低长序列训练、推理的显存消耗。
而本文想要探究的Linear Attention对self-Attention算法进行了重新设计。
下面来看具体的做法。

回顾经典的Self-Attention

对于self-attention,从vector view可以写作下面的形式
notion image
对于
若为full-attention
此时
若为causal attention
此时
在generation阶段,时刻下,可以复用过去时间步计算的key,value
也就是常说的KV-cache
从上面的式子可见,随着生成序列的增加,需要存储的KV-cache线性增加,既有显存存储瓶颈,也有带宽瓶颈。
为了降低原生CausalTransformer的KV-cache,业内设计了很多优秀的算法,如GQA,MQA,MLA等

Classic Linear Attention的做法

本文将介绍另一类技术路线:LinearAttention
回顾经典的softmax attention,在时刻的输出为
其可以抽象为
对于标准的softmax attention,
Linear Attention 用另一种方法设计这个相似度函数,其定义为
这样设计有一个很大的好处,此时可以将移出求和符号
论文将定义为:
以保证相似度权重为正。
从上式可见,在generation阶段,时刻下,可以复用过去时间步的状态
式子12可写作
从上可见,LinearAttention的“KV-cache”是恒定的,不会随着序列的增加而增加。
注意,论文中的公式与本文差一个转置,这是因为本文将定义成行向量。
通常为了数值稳定性,分母会加一个
在训练阶段,当然也可以类似softmax attention,写成矩阵乘法的形式并行计算,如:
,上式可以简写为
其中是causal mask
这里需要注意与softmax based attention的区别,softmax based attention是加性的mask,linear attention是乘性的mask。这是因为:
softmax有指数运算,因此设置,从而
式16的算式时空复杂度为。因此,实际实现中,还是会采用递推的方式计算,此时
如果考虑整个transformer block (忽略layernorm,multi-head)
这里需要注意softmax attention利用online softmax的技巧,也能写成递推形式。二者attention的核心区别点是如何定义“相似”,这个定义方法决定了“kv-cache”的形式。linear attention是固定维度的前缀和,softmax attention会随序列增长而不断拼接 online softmax相关的内容可见: http://myhz0606.com/article/flash_attn 或者论文:https://arxiv.org/abs/1805.02867
算法流程20用RNN的形式,重写了transformer block的计算流程。也映射了论文的标题《Transformers Are RNNs》,即带 causal mask 的 Linear Attention 可以等价写成一个拥有矩阵状态的 RNN。
注意⚠️,虽然写成了上述递归形式,但并不意味着训练时需要串行计算(也无需用式16的矩阵形式)。显然从计算上本质都是求前缀和,因此可以用并行扫描来实现并行化。前缀扫描相关的知识可以参考我之前的blog:http://myhz0606.com/article/mini_rnn

Classic Linear Attention的局限性

1 短序列场景下,缓存和计算未必占优
下表为linear attention与softmax attention在时刻下的cache
cache
元素个数
Linear Attention
Softmax Attention
临界点:
从中可见,若序列较短(约小于)时,linear attention需要的cache反而更多。
2 长序列表达能力受限
随生成序列长度的增加,经典的linear attention会持续将外积用累加的方式更新到state。这至少有2个问题:
  1. Linear attention的memory是有限的,随着序列的增长,这个有限的memory会成为瓶颈。
  1. linear attention的memory更新机制太过“粗暴”。直接用累加的方式对整个memory进行更新,新旧信息很难无损叠加,易产生memory collision.
针对第一点,目前较多的解决方案是采用linear attention和softmax attention交替的hybrid transformer架构。目前较多采用3:1的配比(如qwen3.5),即3层linear attention后,插入1层full attention。相关研究可以参考《A Systematic Analysis of Hybrid Linear Attention》
针对第二点,后续很多linear attention的改进工作设计了更精细的状态更新机制。如mamba2给state update的更新机制引入gate,使其能够“遗忘”历史信息;DeltaNet引入delta-rule,根据当前key的读取误差定向修正状态;gated deltanet 引入decay gate和delta rule,同时支持遗忘和精确写入。详细的原理后续博文介绍。

小结

本文详细介绍了经典的linear attention提出的motivation和算法原理。并对其局限性做了讨论。若有问题,欢迎指出~

参考文献

给身边考研的小伙伴diffusion model(一):DDPM技术小结 (denoising diffusion probabilistic)
Loading...