11 · Lightning Attention

注意力 Qin et al. 2024
Qin et al., NUS · arXiv:2401.04658

迭代了什么

上一代的问题:标准注意力 O(n²) 复杂度,线性注意力理论上 O(n) 但实际慢(cumsum 无法利用 GPU 并行)。

这代改了什么:分治策略:块内用标准左乘注意力 O(B²),块间用线性注意力右乘 O(d²)。序列分块,递推更新累积 KV 状态。

效果:序列长度从 1K 到 92K TGS 几乎恒定,显存线性增长 vs FlashAttention O(n²) 到 OOM。

一、解决的问题

线性注意力理论上 O(n) 复杂度,但因果设置中 cumsum 操作无法充分利用 GPU 并行性,实际速度远慢于理论。

二、分治策略

块内用传统左乘注意力 O(B²),块间用线性注意力右乘 O(d²):

Ointra = [(Qi KiT) ⊙ M] Vi
Ointer = Λ Qi (KV)
kvt = λ · kvt-1 + ktT vt

三、效果

序列长度从 1K 到 92K 训练 TGS 几乎恒定,而 FlashAttention-2 随长度增长到 OOM。