FlashAttention

1331 字
7 分钟
FlashAttention
2026-08-10

在介绍 FlashAttention 之前,先统一本文使用的符号。

符号含义
NNSequence Length,序列长度
DDHead Dimension,每个 Attention Head 的特征维度
QQQuery 矩阵
KKKey 矩阵
VVValue 矩阵

为什么 Attention 很慢?#

Transformer 中最重要的计算之一就是 Self-Attention, 从公式上看,Attention 主要由两个矩阵乘法和一个 Softmax 组成:

S=QKTP=Softmax(S)O=PV\begin{aligned} S &= Q K^T \\ P &= \operatorname{Softmax}(S) \\ O &= PV \end{aligned}Q,K,VRN×DQ,K,V\in\mathbb{R}^{N\times D}

而:

S,PRN×NS,P\in\mathbb{R}^{N\times N}

最终:

ORN×DO\in\mathbb{R}^{N\times D}
Slide imageSlide imageSlide image
1 / 3
Self-Attention 计算过程

而问题在于这个 N×NN \times N 的中间矩阵,当 Sequence Length NN 增大时,Attention Matrix 的大小按照 N2N^2 增长,当序列长度从 1024 增长到 8192 时,这个中间矩阵的元素数量会增加 64 倍。

FlashAttention 的核心思想并不是减少 Attention 的 FLOPs 计算量,而是:

Important

不要把巨大的 Attention Matrix 写回 GPU 的 HBM,而是把计算拆成小块,让数据尽可能留在 GPU 的片上高速存储中。

在 Standard Attention 实现中,中间矩阵 SSPP 只是为了下一个计算阶段使用,却被反复写入和读取 HBM,后面的 FlashAttention 就是围绕如何减少这些 HBM IO 展开的。

Slide imageSlide imageSlide imageSlide image
1 / 4
Standard Attention 实现

IO-Aware Algorithm:Tiling#

FlashAttention 所指的 IO-Aware 思想实际上就是 Tiling。 在矩阵乘法的优化中,我们介绍了 Tiling,比如矩阵 AA 乘以矩阵 BB,我们将 AA 矩阵和 BB 矩阵进行分块,逐块加载到 shared memory 中,读取 shared memory 完成计算,而不是直接读取 global memory (HBM)。

Slide imageSlide imageSlide imageSlide imageSlide image
1 / 5
矩阵乘法 Tiling

对于 Attention 计算 S=QKTS = QK^T 我们也可以使用 Tiling,一块一块的计算 Score Matrix。

Online Softmax#

先回顾普通的 Safe Softmax,三次遍历分别求最大值、求指数和以及计算结果。

FlashAttention 不希望保存完整的 Score Matrix,因为它很占空间,而标准的 Softmax 实现需要完整的 Score Matrix 的一行,因此我们需要优化 Softmax 的实现让它也可以一块一块地处理。

Safe Softmax 标准实现
Safe Softmax 标准实现

假设已经处理:

x1,,xkx_1,\cdots,x_k

定义两个变量 running max mm 以及 running sum ll

mk=max1ikxim_k=\max_{1\le i\le k}x_ilk=i=1keximkl_k= \sum_{i=1}^{k}e^{x_i-m_k}
Important

lNl_N 就是我们想要的 Softmax 的分母。

现在来了新的 block:

xk+1,,xk+bx_{k+1},\cdots,x_{k+b}

新的最大值:

mnew=mk+b=max(mk,max(xk+1,,xk+b))=max(mk,mblock)m_{new} = m_{k + b} = \max \left( m_k, \max(x_{k+1},\cdots,x_{k+b}) \right) = \max \left( m_k, m_{block} \right)

由于旧的 lkl_k 是按照 mkm_k 计算的,因此需要引入纠正系数:

correction factor=emkmnew\text{correction factor} = e^{m_k-m_{new}}

不难得到:

lnew=lk+b=lkemkmnew+j=k+1k+bexjmnew=lkemkmnew+lblockl_{new} = l_{k + b} = l_k e^{m_k-m_{new}} + \sum_{j = k + 1}^{k + b} e^{x_j-m_{new}} = l_k e^{m_k-m_{new}} + l_{block}

这样,我们只需要维护:

m,l\boxed{m,l}

就可以将 Softmax 的前两次循环,求最大值循环和求和循环,压缩为一次循环。

我们实际上把三次循环缩减到了两次循环,而这就是 Online Softmax

Flash Attention#

Online Softmax 还不够,我们真正需要的是 O=PVO = PV,我们希望 OO 也可以在线更新。

假设我们将 QRN×DQ \in \mathbb{R}^{N \times D} 切分为一个一个的 QblockRtile_q×DQ_{block} \in \mathbb{R}^{\text{tile\_q} \times D},同理将 KRN×DK \in \mathbb{R}^{N \times D}VRN×DV \in \mathbb{R}^{N \times D} 分别切分为一个一个的 KblockRtile_k×DK_{block} \in \mathbb{R}^{\text{tile\_k} \times D}VblockRtile_k×DV_{block} \in \mathbb{R}^{\text{tile\_k} \times D}

Flash Attention
Flash Attention

我们先考虑 tile_q=1,tile_k=b\text{tile\_q} = 1, \text{tile\_k} = b 的情况,这个时候我们的 OblockRDO_{block} \in \mathbb{R}^{D}QblockRDQ_{block} \in \mathbb{R}^{D} 变成了一维向量,记作 ooqq。 在 CUDA 实现中,每一个线程块处理到一个 qq,我们后续推导专注一个线程块上的计算,所以默认 qq 是不变的。

假设我们处理到某个 SblockS_{block} 时已经处理了:

x1,,xkx_1, \cdots, x_k

其中 xi=qkiTx_i = q k^T_i,其中 kiTk^T_iKTK^T 矩阵中的一个列向量。

我们在 online softmax 中定义了 running max mm,running row sum ll,也就是:

mk=max1ikxim_k=\max_{1\le i\le k}x_ilk=i=1keximkl_k= \sum_{i=1}^{k}e^{x_i-m_k}

当处理完 SblockS_{block} 时会引入:

xk+1,,xk+bx_{k + 1}, \cdots, x_{k + b}

我们可以得到新的 mmll:

mnew=mk+b=max(mk,mblock)m_{new} = m_{k + b} = \max \left( m_k, m_{block} \right)lnew=lk+b=lkemkmnew+lblockl_{new} = l_{k + b} = l_k e^{m_k-m_{new}} + l_{block}

现在考虑 oo,它是:

o=j=1NexjmNlNvjo = \sum_{j = 1}^N \frac{e^{x_j - m_N}}{l_{N}} v_j

其中 vjRDv_j \in \mathbb{R}^DVV 中的一个行向量。

我们想让 oo 也能在线更新,于是我们定义新的变量 oio_i

oi=j=1iexjmivjo_i = \sum_{j = 1}^i e^{x_j - m_i} v_j

那么自然就有

o=oN/lNo = o_N / l_N

现在我们推导 oko_kok+bo_{k + b} 的关系:

ok+b=j=1k+bexjmk+bvj=j=1kexjmk+bvj+j=k+1k+bexjmk+bvj=(j=1kexjmkvj)emkmk+b+j=k+1k+bexjmk+bvj=ok×emkmk+b+j=k+1k+bexjmk+bvj=ok×emkmnew+[exk+1mnewexk+2mnewexk+bmnew][vk+1vk+2vk+b]\begin{aligned} o_{k + b} &= \sum_{j = 1}^{k + b} e^{x_j - m_{k + b}} v_j \\ &= \sum_{j = 1}^{k} e^{x_j - m_{k + b}} v_j + \sum_{j = k + 1}^{k + b} e^{x_j - m_{k + b}} v_j \\ &= \left(\sum_{j = 1}^{k} e^{x_j - m_{k}} v_j \right) e^{m_{k} - m_{k + b}} + \sum_{j = k + 1}^{k + b} e^{x_j - m_{k + b}} v_j \\ &= o_k \times e^{m_{k} - m_{k + b}} + \sum_{j = k + 1}^{k + b} e^{x_j - m_{k + b}} v_j \\ &= o_k \times e^{m_{k} - m_{new}} + \begin{bmatrix} e^{x_{k+1} - m_{\text{new}}} & e^{x_{k+2} - m_{\text{new}}} & \cdots & e^{x_{k+b} - m_{\text{new}}} \end{bmatrix} \begin{bmatrix} v_{k+1} \\ v_{k+2} \\ \vdots \\ v_{k+b} \end{bmatrix} \end{aligned}

太好了,我们将 ooldo_{old}onewo_{new} 联系到一起了,现在我们可以在线更新 oo 了!

而且我们可以很轻松地扩展到 tile_q1\text{tile\_q} \neq 1 的情况,Flash Attention 的计算过程如下:

Flash Attention 计算过程演示
Flash Attention 计算过程演示

完整的 Flash Attention 实现#

我的实现

参考#

支持与分享

如果这篇文章对你有帮助,欢迎分享给更多人或赞助支持!

赞助
FlashAttention
https://llm-tech.com.cn/posts/flashattention/
作者
Ming
发布于
2026-08-10
许可协议
CC BY-NC-SA 4.0
Profile Image of the Author
Ming
你是来找 Ming 学习的吗
🎉 欢迎来到 Ming 的博客
这里是我的个人博客,分享 AI Infra、LLM 等技术内容。欢迎关注交流!
分类
标签
站点统计
文章
19
分类
8
标签
16
总字数
55,114
运行时长
0
最后活动
0 天前

目录