FlashAttention 笔记
Standard Self-Attention
Safe Softmax
计算 safe-softmax 的时候需要对 [1, N] 重复三次,需要访问 Q 和 K 三次, 并实时重新计算 x,很低效,现在希望将计算进行合并。
现在定义
当 i = N 的时候,有:
递归关系只依赖
FlashAttention V1
online softmax 最多只有一个 2-pass 的算法,不存在 1-pass 算法,但是 Attention 可以有 1-pass 算法。基于上述的 online softmax 可以得到一个 1-pass 的 Attention 算法。
重点在第二个循环:
推导 1-pass 版本的 FlashAttention:
推导
可以看到
FlashAttention V2
主要做了工程上的优化:
- 减少大量非矩阵乘的冗余计算,增加 Tensor Core 的计算比例
- forward pass/backward pass 均增加 seqlen 维度的并行,forward pass 交替 Q,K,V 循环顺序
- 更好的 Warp Partitioning 策略,避免 Split-K
在 Tri Dao 的 FlashAttention 的 Triton 实现中,使用了 LSE(LogSumExp) 来做 smooth maximum。
LSE 可以被用于作为 approximation of max。
Triton 中的算法可以被描述为:

等式 lse_i 被用于作为最大值的一个近似。
在等式
在等式 o_scale 作为 softmax 函数的分母: