FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
大名鼎鼎的算子融合技术,将attention融合到一个算子完成,但具体是怎么做的?
Background
自注意力(Self-Attention)机制的计算复杂度和内存消耗均随序列长度 \(N\) 呈二次方 \(O(N^2)\) 增长,但随着硬件算力的发展,计算不是瓶颈,真正的瓶颈在于GPU HBM(高带宽内存)读写带来的巨大时延。
经典算法中(算子不融合)
- \(S = QK^\top\) 计算:读取 \(Q, K \in \mathbb{R}^{N \times d}\) 时,虽然底层通过 SRAM 局部块完成乘积累加计算,但必须将输出结果 \(S \in \mathbb{R}^{N \times N}\) 完整写回 HBM,此时有 \(\Theta(Nd + N^2)\) 的直接访存开销。
- \(P = \text{softmax}(S)\) 计算:由于 Softmax 算子是独立的,硬件必须重新从 HBM 读入庞大的 \(S\),计算后再次将 \(P \in \mathbb{R}^{N \times N}\) 写入 HBM,引发额外的 \(\Theta(N^2)\) HBM 访问
- \(O = PV\) 计算:最后再次读取 \(P\) 和 \(V\),计算完毕将输出 \(O \in \mathbb{R}^{N \times d}\) 写入 HBM,访存同样为 \(\Theta(Nd + N^2)\)。
而FlashAttention做了将这些计算步骤一次性完成,减少其中的数据搬运。具体而言使得\(K\) 和 \(V\) 的每个元素仅从 HBM 加载一次,使得性能得到提升。
核心思路是:允许更多的计算以减少内存访问的瓶颈(计算换带宽)
Methods
这里我们只聚焦推理的计算过程(Forward),训练需要的Backward也做了对应的优化。 标准的Softmax中,归一化的分母是整个行的注意力分数的指数运算之和,这需要一整行Q.K的结果,同时SRAM的容量是有限的(容量为M),所以标准的做法是算完Q.K保存回HBM,再读取计算下一步。 如果想在有限的SRAM上一次性算完Attention,必须对QKV进行分块读取,FlashAttention设计了一个精妙的计算,克服了分块的同时又能维护整个行的注意力分数的指数运算之和。 具体步骤是 1. Q矩阵常驻SRAM,分块读取K、V矩阵,与Q运算求出并维护局部的和li,同时根据局部的权重与分块后的V计算局部的输出Oi 2. 计算下一块时:会重缩放局部指数和并累加,根据上一块的局部和进行数学形式的放缩以及新块局部和的聚合,能得到当前的局部和li+1,这不是估计而是无损的。相同的思路,对之前的局部输出Oi进行放缩,并加入新块的作用,得到Oi+1。
重缩放局部指数和并累加:$\(l^{(j+1)} = e^{m^{(j)} - m^{(j+1)}} l^{(j)} + e^{\tilde{m} - m^{(j+1)}} \tilde{l}\)$
重缩放输出 \(O\)(消除旧的分母,应用新的统一分母):$\(O^{(j+1)} = \text{diag}(l^{(j+1)})^{-1} ( \text{diag}(l^{(j)}) e^{m^{(j)} - m^{(j+1)}} O^{(j)} + e^{\tilde{m} - m^{(j+1)}} \tilde{P} V_j )\)$
其他算子
Mask 直接作用在QK分块计算后的结果,不再整体MASK。Dropout 按照概率随机置0,并存下伪随机数生成器状态(不知道是什么,这个在Backward使用。不过多介绍了)
Evaluation
测试平台硬件: 实验基准主要在 NVIDIA A100 GPU (40GB/80GB HBM, 高达 1.5-2.0 TB/s 带宽) 上进行。
速度大幅提升(到现在都在使用)