为自回归模型实现因果注意力机制,通过跳过上三角计算实现 ~2x 加速。
配套代码
在前三章中,我们实现了一个完整的 Flash Attention Kernel,支持任意长度的序列和多维并行。但我们还没有针对一个非常重要的应用场景进行优化:自回归生成(Autoregressive Generation)。
这正是 GPT、LLaMA 等 Decoder-Only 模型的核心计算模式。通过引入 Causal Masking,我们可以进一步获得约 2倍的加速。
Causal Attention 是自回归语言模型的核心机制:在预测第 𝑖 i 个 token 时,只能看到位置 0 到 i-1 的信息,不能看到未来的 token。
在注意力矩阵中,这通过将上三角部分置为 − ∞ −∞ 实现:
K₀ K₁ K₂ K₃
Q₀ [ · -∞ -∞ -∞ ] ← Q₀ 只能看 K₀
Q₁ [ · · -∞ -∞ ] ← Q₁ 只能看 K₀,K₁
Q₂ [ · · · -∞ ] ← Q₂ 只能看 K₀,K₁,K₂
Q₃ [ · · · · ] ← Q₃ 能看到所有
如果您对 Causal Attention 的概念和原理还不熟悉,建议先阅读 Attention 机制详解 中的详细介绍。
本章重点关注如何在 Flash Attention 中高效实现 Causal Masking 以获得性能提升。
既然上三角的元素最终都会被 Softmax 归零,为什么还要计算它们?
对于一个 ( 𝑁 , 𝑁 ) (N,N) 的注意力矩阵:
1 ) 2 2 N(N+1)
个元素(下三角+对角线)
理论加速比:
当 𝑁 N 很大时,加速比接近 2倍。
几何直觉: