请添加图片描述

前言

先说结论:同样跑 LLaMA-7B,序列长度 4096,FlashAttention 比标准 Attention 快 3.2 倍,显存从 32GB 降到 6GB。

这个差距不是来自某个算子的优化,而是来自对 Attention 计算公式的重新推导。


标准 Attention 的计算过程是:QK^T → Softmax → ×V。这个公式在 GPU 上实现时,每一步都要把中间结果写回显存。Q、K、V 三个矩阵各 1GB(batch=8, seq=4096, hidden=4096),QK^T 的中间结果又是 1GB,Softmax 之后再算一遍,显存读写总量是矩阵大小的 5-6 倍。

FlashAttention 的核心思路:不把整个矩阵乘出来,而是按 block 切分,每个 block 在 SRAM 里算完再写回 HBM。SRAM 的带宽是 HBM 的 10 倍以上,瓶颈从「显存带宽」变成了「计算」。

具体做法分三步。第一步,把 Q、K、V 按行切分成 block,每个 block 大小是 br × d(Q 的 block)和 bc × d(K、V 的 block)。brbc 的选择取决于 SRAM 大小——昇腾 910 的 SRAM 是 192KB,去掉寄存器占用,实际能用的约 128KB,对应 br=64, bc=64

第二步,外层循环遍历 Q 的 block,内层循环遍历 K/V 的 block,每次只加载一个 br × bc 的小矩阵到 SRAM,算局部注意力,累加到输出 block 上。这个过程里,Q、K、V 的大矩阵全程在 HBM 上,只有当前 block 在 SRAM 里。

第三步,Softmax 的归一化因子要特殊处理。标准 Softmax 需要全局最大值才能算,exp(x - max) 这个操作在分 block 的情况下,max 是局部的,不是全局的。FlashAttention 用了一个递推公式:每读一个新的 K block,更新一次全局 max 和累加和,保证数值正确。这个推导在原始论文的 Appendix B 里有,看懂需要点耐心。

昇腾上的实现比 GPU 多一层:Cube Unit 做矩阵乘,Vector Unit 做 Softmax 和缩放。两个 Unit 可以并行——Cube 在算 Q_block × K_block^T 的时候,Vector 可以同时处理上一个 block 的 Softmax。这个流水线是 FlashAttention 在昇腾上快于 GPU 的关键。

实际使用时,不用手动写这些逻辑。ops-transformer 仓库里的 flash_attention 算子已经封装好了,直接调用:

from ops_transformer import flash_attention

output = flash_attention(
    query=Q,  # [batch, num_heads, seq_len, head_dim]
    key=K,
    value=V,
    attn_mask=mask  # 可选
)

这个算子会自动根据 SRAM 大小选择 brbc,不需要手动调。但要注意输入 shape 的对齐——seq_len 最好是 64 的倍数,否则 SRAM 利用率会掉 10-15%。

瓶颈在哪?当序列长度超过 8192 的时候,FlashAttention 的增益开始缩小。原因是 SRAM 放不下更大的 block,退化为多次小 block 计算,SRAM 和 HBM 之间的搬运次数反而增加了。这时候要用 PagedAttention 或者 Multi-Query Attention 进一步压缩显存。

测试数据(LLaMA-7B,昇腾 910,FP16):

序列长度 标准 Attention (ms) FlashAttention (ms) 显存 (标准) 显存 (FA)
512 12 8 4GB 3GB
2048 68 31 16GB 5GB
4096 OOM 89 >32GB 6GB
8192 OOM OOM >32GB >32GB

4096 是 FlashAttention 在 32GB 显存上的上限。要跑更长序列,需要 PagedAttention 或者激活重计算(activation recomputation)。

Logo

作为“人工智能6S店”的官方数字引擎,为AI开发者与企业提供一个覆盖软硬件全栈、一站式门户。

更多推荐