FlashAttention:分块计算与在线 Softmax
理解不落盘完整注意力矩阵的计算方式,以及在线 Softmax 如何跨块合并。
FlashAttention:分块计算与在线 Softmax
长序列注意力不仅需要算力,也需要搬运大量中间数据。本文根据注意力优化笔记整理 FlashAttention 的核心直觉,讨论标准稠密注意力的精确计算。
标准注意力的中间矩阵
单个头、忽略掩码时:
\[O=\operatorname{softmax}\left(\frac{QK^T}{\sqrt d}\right)V\]序列长度为 \(N\),头维度为 \(d\)。直接实现会生成 \(N\times N\) 的分数矩阵及概率矩阵,给显存容量和读写带来压力。
点积相关性不是因为 Q、K 必须经过某种 LayerNorm 才成立;不同模型的归一化与投影结构需要分别确认。
稳定 Softmax
对一行分数,令 \(m=\max_j s_j\),则:
\[p_j=\frac{e^{s_j-m}}{\sum_k e^{s_k-m}}\]减去最大值保持比值不变,同时避免直接对很大的正数取指数。分块时,难点是此前块使用的最大值可能不是全行最大值。
在线合并的量
保留当前最大值 \(m\)、归一化和 \(\ell\),以及未归一化输出累计量 \(z\)。新块的对应量为 \(m_b,\ell_b,z_b\)。更新:
\[m'=\max(m,m_b)\] \[\ell'=e^{m-m'}\ell+e^{m_b-m'}\ell_b\] \[z'=e^{m-m'}z+e^{m_b-m'}z_b\]最后输出 \(z'/\ell'\)。缩放旧块与新块的累计量,让它们使用同一个最大值基准。
分块降低的是数据搬运
FlashAttention 原论文将这种计算组织为 GPU 分块内核,减少 HBM 与片上存储之间的读写,避免保存完整注意力矩阵。它不是通过删去注意力项来实现近似。
“精确”指计算目标仍是原来的注意力,浮点执行顺序仍可能带来数值差异。标准稠密注意力的算术规模也不会因此简单变成线性。
与缓存管理分开理解
FlashAttention 关注如何计算,PagedAttention 关注服务时 KV cache 如何组织。两者解决的问题不同,可以在系统中形成互补。实际速度还取决于序列长度、头维度、硬件和实现。
This post is licensed under CC BY 4.0 by the author.