先这样答
Flash Attention 是一种精确的注意力计算算法,其输出结果与标准注意力完全一致,没有任何近似。它优化的核心是 GPU 显存层级之间的读写过程。标准注意力计算的时间主要耗在把庞大的中间注意力矩阵于高带宽显存和片上高速缓存之间搬运,这是一个典型的访存受限问题。
为了解决读写瓶颈,该算法引入了两个核心技术。第一个是分块计算,算法把查询、键、值矩阵切成能装进片上高速缓存的小块,中间结果算完直接在缓存里处理,不写回高带宽显存。这把高带宽显存的访问次数从序列长度的平方级降了下来。第二个是在线 Softmax,为了配合分块机制,算法维护一个运行最大值和归一化分母,每算完一块就对历史累积量进行重缩放,从而实现单遍流式的精确计算。
在反向传播时,它不保存庞大的中间注意力矩阵,而是直接从片上高速缓存读取输入数据重新计算。这相当于用大约两倍的浮点运算量换取了空间。从公开测试看,初代算法让 GPT-2 在 1K 长度下实现了 3 倍加速,并在 16K 长度的 Path-X 任务上首次做到了高于随机水平的准确率。
在实际工程中,这种优化对模型训练和推理的预填充阶段收益明显,目前已经成为大模型底层计算的基础配置。
面试官会怎么追问
-
「它的计算量变小了吗?为什么速度变快了?」 它的计算量并没有减少,速度变快是因为标准注意力机制是访存受限的。Flash Attention 减少了中间矩阵在高带宽显存和片上高速缓存之间的读写次数,省下的是通信搬运的时间,而不是矩阵乘法的计算时间。
-
「Flash Attention 2 和 3 分别做了哪些迭代?」 二代主要增加了序列长度维度的并行,改善了线程块分工,将 A100 的理论峰值利用率从一代的不到一半提升到了 50 到 73%。三代则专门针对 H100 硬件,利用异步特性实现了计算和数据搬运的重叠,并在 FP8 精度下达到了接近 1.2 PFLOPs/s 的算力。
-
「推理的 Decode 阶段用 Flash Attention 还有收益吗?」 收益有限。Decode 阶段每次只生成单个 token,主要时间花在加载 KV Cache 上,算术强度很低,属于带宽瓶颈。直接套用该算法无法解决这个问题,通常需要改用按 KV 维度并行的衍生方法来处理。
回答的坑
误以为 Flash Attention 是通过近似计算来加速的。正确方向是强调该算法的计算结果与标准注意力在数学上完全等价,属于精确计算。
说它把显存复杂度降下来是因为不存任何中间结果。正确方向是具体说明前向计算时中间矩阵不写回高带宽显存,以及反向传播时利用输入数据重新计算。
同系列的题