1. GPU 内存墙的本质与挑战
当我们在谈论 GPU 计算性能时,常常会陷入一个误区:只看浮点运算能力(FLOPs)。但实际上,现代 GPU 面临的最大瓶颈不是计算能力,而是内存带宽。这就是所谓的"内存墙"问题。
以 NVIDIA A100 为例:
- 计算能力:624 TFLOPS(FP16)
- HBM 带宽:1.5 TB/s
- SRAM 带宽:19 TB/s
这意味着即使 GPU 有强大的计算能力,如果数据不能及时供给,计算单元就会处于"饥饿"状态。在标准的注意力计算中,这个问题尤为突出:
- QK^T 计算:需要从 HBM 读取 Q 和 K
- Softmax:需要将中间结果写回 HBM 再读取
- PV 计算:再次读取中间结果和 V
这种反复的数据搬运导致了严重的性能瓶颈。在实际测试中,即使使用最先进的框架,GPU 利用率也常常只有 20-30%。
关键发现:在注意力计算中,数据搬运消耗的时间远超过实际计算时间。这就是 FlashAttention 要解决的核心问题。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. FlashAttention 的架构革新
2.1 存储层次感知计算
FlashAttention 的核心思想是最大化数据复用,最小化 HBM 访问。这需要深入理解 GPU 的存储层次:
- 寄存器(Register):最快,但容量最小(每个线程私有)
- 共享内存(Shared Memory):块内共享,192KB/Block
- 全局内存(Global Memory/HBM):大容量但速度慢
传统实现的问题在于:
- 中间结果(S=QK^T, P=softmax(S))都需要写回 HBM
- 每个操作都是独立的内核(kernel),导致多次数据搬运
FlashAttention 的解决方案:
- 将整个注意力计算融合为单个内核
- 通过分块计算(Tiling)适配 SRAM 容量
- 中间结果保留在寄存器/SRAM 中
2.2 分块计算(Tiling)实现
分块计算的关键在于将大矩阵分解为适合 SRAM 的小块。以序列长度 N=4096,头维度 d=128 为例:
- 将 Q 按行分块:Q = [Q1, Q2, ..., Qb]
- 将 K,V 按行分块:K =
