线性注意力(如 lightning attention 这类方案)为什么能加速长序列?代价是什么?
先这样答
先讲清标准注意力贵在哪。注意力要对序列里每个 token 和其他所有 token 算相似度,输出是一个长度乘长度的矩阵,计算量和显存都随序列长度平方增长。序列从 4K 到 400K,注意力开销理论上涨一万倍,这就是长上下文推理的天文数字成本来源。
线性注意力的思路是改变计算顺序。标准注意力先算出完整的两两相似度矩阵再加权求和;线性注意力用核函数替换 softmax,利用矩阵乘法的结合律,先对价值向量做聚合、再与查询作用,绕开显式构造的平方大小矩阵。这样计算量随长度线性增长,长序列的加速非常可观,训练和推理两端都受益。
代价是表达力。softmax 注意力的动态归一化和精确的两两交互,是模型容量的一部分;核函数近似和聚合先行都会损失一些细粒度的区分能力,直接全量替换标准注意力,短序列任务的效果会回撤。所以工程落地普遍用混合架构:一部分层保留标准注意力保证容量,一部分层用线性注意力吃下长序列的效率,在效果和成本之间取平衡。MiniMax-01 系列公开采用的就是这类闪电注意力加混合架构的路线。
面试官会怎么追问
- 「为什么不能全部层都换线性注意力?」 效果会掉。标准注意力在捕捉细粒度依赖上仍有优势,混合架构里保留的满注意力层承担关键的信息路由。哪些层保留、比例多少,靠消融实验定。
- 「推理端除了注意力还有什么长序列瓶颈?」 KV 缓存显存随长度线性膨胀,prefill 时间随长度变长拖慢首字延迟。注意力只是一项,完整的长序列优化要同时处理这三件事。
- 「训练端的收益怎么体现?」 长序列训练时注意力的显存峰值从平方级降下来,同样的卡能开更长的序列窗口,长文本数据的训练成本直接可控。
回答的坑
- 只背「线性复杂度」结论,讲不出计算顺序改写的原理。追问一层就露馅。
- 不提表达力代价。任何加速方案都要回答「牺牲了什么」,混合架构的存在本身就是答案。
同系列的题
相关深度笔记
—— 本题完 ——