PyTorch 2.x 深度学习专题【左扬精讲】—— Flash Attention原理:分块计算与重计算,如何降低O(n^2)的显存占用
PyTorch 2.x 深度学习专题【左扬精讲】—— Flash Attention原理:分块计算与重计算,如何降低O(n^2)的显存占用
Flash Attention v1 (Tri Dao, 2022) ← 分块计算 + 在线softmax
Flash Attention v2 (Tri Dao, 2023) ← 优化thread block与warp调度
Flash Attention v3 (Tri Dao, 2024) ← FP8量化 + 异步pipeline
torch.nn.functional.scaled_dot_product_attention ← PyTorch 2.0 内置实现
xformers library ← memory_efficient_attention
Flash Attention 分块计算 在线softmax SRAM/HBM 重计算 显存优化
本文学习重点
- 必须掌握
- 标准Attention的显存瓶颈在哪里,为什么S矩阵是罪魁祸首
- Flash Attention如何通过分块(Tile)避免物化完整S矩阵
- 在线softmax的数学推导:如何分块增量计算softmax
- 重计算(Recomputation)的思想:在反向传播时重新计算Q/K/V而非存储它们
- 需要理解
- SRAM与HBM的带宽差异如何驱动算法设计
- Flash Attention 1/2/3的核心差异
- PyTorch 2.0中scaled_dot_product_attention如何自动选择最优实现
目录导航
一、标准Attention的显存困境:O(n^2)到底困在哪里
What — 标准Attention的显存占用是什么?
回顾标准Scaled Dot-Product Attention的计算流程:
- 输入:Q、K、V,形状均为 (batch, num_heads, seq_len, head_dim)
- Step 1:计算 S = Q * K^T,形状 (seq_len, seq_len)
- Step 2:softmax归一化:P = softmax(S / sqrt(d))
- Step 3:O = P * V,形状 (seq_len, head_dim)
Why — 为什么说O(n^2)显存是真正的瓶颈?
问题一:S矩阵的显存随序列长度平方增长。
以GPT-3级别的模型为例:seq_len = 4096,head_dim = 64,batch = 1,num_heads = 96。单层Attention中,S矩阵的大小为:
seq_len = 4096
head_dim = 64
# S矩阵形状:(seq_len, seq_len) = (4096, 4096)
s_matrix_bytes = seq_len * seq_len * 4 # FP32: 4字节
s_matrix_bytes / 1024 / 1024 # 转换为MB
# 结果:64 MB(单层,单个head)
# 考虑96个head、40层,总显存:
total_s = 64 * 96 * 40 # MB
total_s / 1024 # 转换为GB
# 结果:约 240 GB(远超A100 80GB显存上限)
这还仅仅是S矩阵,P矩阵(softmax后的结果)同样需要存储,加上Q/K/V本身,实际单层Attention的中间结果就轻松突破数百GB。
问题二:反向传播时需要存储P矩阵,显存翻倍。
反向传播需要 dO 梯度,以及对P矩阵的依赖(S矩阵或P矩阵需要保留),否则无法计算 dS 和 dQ、dK、dV。
没有分块计算会发生什么?
- 序列长度翻倍 → 中间激活值显存翻4倍(n^2关系)
- 超过GPU HBM容量 → 不得不使用CPU主存 → 速度断崖式下降
- 长上下文训练(32K~128K token)几乎不可行
- 即使能装下,多次HBM读写也导致带宽成为瓶颈(见下节)
以下代码展示标准Attention的实现,其中中间矩阵S和P都需要完整物化到HBM:
def attention_forward(Q, K, V, scale=None):
if scale is None:
scale = (Q.size(-1) ** -0.5) # 缩放因子 1/sqrt(d_k)
# Step 1: Q @ K^T,形状 (seq_len, seq_len),显存占用 O(n^2)
# 完整S矩阵必须存储在HBM中,供后续softmax和反向传播使用
S = torch.matmul(Q, K.transpose(-2, -1)) * scale # S[i,j] = Q[i] dot K[j]
# Step 2: softmax归一化,输出P矩阵,同样需要存储(反向传播需要)
P = F.softmax(S, dim=-1) # 按列(key维度)归一化
# Step 3: P @ V,形状 (seq_len, head_dim)
O = torch.matmul(P, V)
# 返回O和中间结果S/P,供反向传播使用
return O, {"S": S, "P": P} # S和P都存显存,这是瓶颈所在
关键观察:上述实现中,S 和 P 这两个 (n, n) 的矩阵都需要存储。如果序列长度达到4096,float32下每个矩阵就是64MB,96个head就是6GB+。
第一节小结
- S矩阵物化:标准Attention的S矩阵必须完整存储,显存随n^2增长
- 反向传播加倍:P矩阵也需要存储用于梯度计算,显存再翻倍
- 长序列致命:n=4096时单层S+P矩阵就需要约128MB(单head),乘以多层多头后轻松爆显存
二、存储层级视角:SRAM是加速的物理基础
What — GPU上存在哪些存储层级,各有什么特点?
现代GPU(以A100/H100为例)有两级关键存储:
- HBM(High Bandwidth Memory):即显存,容量大(80GB)但带宽相对低(约2TB/s A100 / 3.35TB/s H100)。读写延迟高。
- SRAM(Static RAM):片上共享内存,容量极小(A100每SM 192KB,共约19MB)但带宽极高(每个SM 19.5TB/s,聚合后远超HBM)。读写延迟极低。
两者带宽差距在10~20倍量级,容量差距则超过4000倍。
Why — 为什么SRAM带宽高却容量小,这如何驱动算法设计?
问题一:SRAM容量限制了单次能处理的数据量。
A100每个SM(流多处理器)有192KB SRAM,128个SM总计约24.5MB。这远小于4096x4096的S矩阵(64MB FP32)。因此必须将数据分块(Tile)放入SRAM。
问题二:HBM带宽是系统瓶颈,但容量大。
HBM的2TB/s带宽看起来很高,但相对于计算单元(Tensor Core)的峰值吞吐量(约 312 TFLOPS BF16 / 989 TFLOPS FP8,A100为例)来说远远不足。GPU计算速度远快于数据供给速度,这被称为"访存墙"(Memory Wall)。
没有SRAM分块会发生什么?
- 每次Q@K^T都要从HBM读取Q和K:读取量 = 2 * n * d * n = 2 * n^2 * d
- HBM带宽成为瓶颈:计算单元等待数据,大量算力被浪费
- SRAM分块让数据一次性加载到片上,多次复用,大幅减少HBM读写次数
以下数据帮助理解两级存储的特性差异:
# A100 GPU 存储层级参数
HBM_capacity_gb = 80 # 显存容量 80 GB
HBM_bandwidth_tb = 2.0 # 带宽 2.0 TB/s
SRAM_per_sm_kb = 192 # 每个SM的SRAM大小 192 KB
num_sm = 128 # SM数量
SRAM_total_mb = SRAM_per_sm_kb * num_sm / 1024 # 约 24 MB
# A100 每 SM 的 shared memory 单向带宽约 19.5 GB/s(双向 39 GB/s)
# 128 个 SM 聚合后 shared memory 总带宽约 2.5 TB/s(不是 PB/s)
# 这个带宽远高于 HBM 的 2.0 TB/s,因此分块计算可以大幅加速
print(f"HBM容量: {HBM_capacity_gb} GB") # 80 GB
print(f"HBM带宽: {HBM_bandwidth_tb} TB/s") # 2.0 TB/s
print(f"SRAM总量: {SRAM_total_mb:.1f} MB") # 约 24 MB
print(f"容量比: {HBM_capacity_gb * 1024 / SRAM_total_mb:.0f}x") # 约3400x
Flash Attention的核心思想就是:利用SRAM的小容量、高带宽特性,将数据切块分次送入SRAM计算,每次计算完毕后立刻释放,最终结果直接写回HBM。中间过程不物化大矩阵。
第二节小结
- HBM:容量大(80GB)但带宽低(2TB/s),适合存储最终结果
- SRAM:容量小(24MB A100)但带宽极高(PB/s级),适合中间计算
- 分块计算的本质:将大矩阵切分为能在SRAM中放下的小块,减少HBM读写
- 容量-带宽权衡:Flash Attention用算力和算法复杂度换取存储层级间的带宽优势
三、分块计算(Tile):把大矩阵切碎喂入SRAM
What — 分块计算的核心思想是什么?
分块计算(Tile-based Computation)的基本思路:将 Q、K、V 按行(或列)切分成若干小块(Tile),每次只将一对块加载到SRAM中,计算它们之间的局部注意力,然后增量聚合到最终结果中。
以序列长度n=4096、块大小B=64为例:
- Q被切分为 4096/64 = 64 个块,每块形状 (64, d)
- K和V同样被切分为64个块
- 每次只将 Q_i 块和 K_j 块加载到SRAM,计算局部 S_ij = Q_i @ K_j^T
Why — 为什么简单的切分思路并不能直接工作?
问题一:softmax必须在全局做,不能在块上独立做。
softmax的计算依赖所有元素的指数和归一化因子:
softmax(x_i) = exp(x_i) / sum_j(exp(x_j))
如果把序列分成 [a, b] 两块分别计算 softmax,则局部 softmax_a(x_i) 不等于全局 softmax(x_i)。这就是分块计算最大的数学障碍。
没有在线softmax会发生什么?
- 只能把所有S_ij块累加后统一做softmax → 仍需要物化完整S矩阵 → 退化为标准Attention
- 或者对每个块独立做softmax → 结果错误,模型无法收敛
以下代码展示分块后单次加载的数据量,以及如何控制SRAM使用:
import torch
# 分块计算参数
seq_len = 4096 # 序列长度
head_dim = 64 # 注意力头维度
block_size = 64 # 块大小(能放入SRAM的关键参数)
num_heads = 96 # 注意力头数量
# 标准实现:一次性加载全部Q和K
standard_load = seq_len * head_dim * 2 # Q和K全部加载
standard_s_matrix = seq_len * seq_len # S矩阵大小
print(f"标准实现 - 加载量: {standard_load} 元素, S矩阵: {standard_s_matrix} 元素")
# 分块实现:每次只加载 B_r x d 的Q块 和 B_c x d 的K块
# 块数量
num_blocks = seq_len // block_size # 64个块
# 每次计算加载量
per_block_load = block_size * head_dim * 2 # Q块 + K块
# 局部S矩阵大小
local_s = block_size * block_size
print(f"分块实现 - 每次加载: {per_block_load} 元素, 局部S矩阵: {local_s} 元素")
print(f"SRAM节省比例: {standard_s_matrix / local_s:.0f}x (以块为单位)")
# FP32下的SRAM占用:local_s * 4 bytes ≈ 16 KB(远小于A100 24MB SRAM上限)
# 分块后的HBM读写次数分析
# 标准:Q@K^T需要1次完整读写(实际上Q和K已经在HBM,但S矩阵需要写回并存储)
# 分块:每个Q块需要遍历所有K块,每次产生局部S矩阵立即参与softmax计算
# 总共 num_blocks^2 = 4096 次局部矩阵乘法,但每次只占用极小SRAM
注意:块大小(Block Size)的选择是性能的关键。太小则分块数量过多(增加循环开销),太大则无法放入SRAM。Flash Attention论文建议根据GPU型号和 SRAM 大小自动选择最优块大小。
第三节小结
- 分块矩阵乘法:Q/K/V按块加载到SRAM,每次只计算局部注意力
- 数学障碍:softmax必须在全局做,分块不能独立softmax
- 解决思路:在线softmax(Online Softmax)算法,允许分块增量计算
四、在线softmax:分块增量计算的核心数学
What — 在线softmax(Online Softmax)是什么?
在线softmax是一种将softmax拆解为可增量计算形式的算法。它的核心思想是利用指数运算的恒等变换,将全局softmax的归一化因子(分母)拆分成多个块的累加——每处理一个新块,只需更新当前的 m_max(全局最大值)和 l(指数和),而不需要物化完整的S矩阵。
Why — 在线softmax的数学原理是什么?
核心洞察:指数函数的稳定性要求。
softmax中 exp(x_i) 当 x_i 很大时会产生数值溢出(float32的上限约 10^38)。标准实现通过减去行最大值来保证数值稳定:
softmax(x_i) = exp(x_i - max(x)) / sum_j(exp(x_j - max(x)))
在线softmax将此性质推广到分块场景。假设分两批处理数据:第一批 x^{(1)},第二批 x^{(2)}。令 m = max(x),m_1 = max(x^{(1)}),m_2 = max(x^{(2)})。
关键恒等变换(归一化因子的链式更新):
l = sum_j(exp(x_j - m)) = exp(m_1 - m) * sum_{i in block1}(exp(x_i - m_1)) + exp(m_2 - m) * sum_{i in block2}(exp(x_i - m_2))
即:只需要维护当前全局最大值 m 和归一化因子 l,就可以逐步合并新块。合并时需要"重新缩放"旧块的贡献(乘以 exp(m_old - m_new)),这是在线softmax的核心操作。
没有在线softmax会发生什么?
- 无法在分块场景下正确计算softmax
- 要么每次都加载完整S矩阵(显存瓶颈),要么接受错误的数值结果
- 所有基于分块的高效注意力算法都以在线softmax为理论基础
以下代码展示从标准softmax到在线softmax的完整推导过程:
import torch
import torch.nn.functional as F
def standard_softmax(x):
# 标准softmax:沿最后一维(key维度)归一化,需要一次性看到所有x
return F.softmax(x, dim=-1) # x形状 (batch, seq_len, seq_len),沿最后一维做softmax
def online_softmax_two_pass(x):
# 分两批演示在线softmax的基本原理
# 假设 x 是注意力分数矩阵 S,形状为 (batch, q_len, kv_len)
B, Q, KV = x.shape
mid = KV // 2
# 按 key 维度分成两块
x1 = x[:, :, :mid] # 块1:包含前mid个key
x2 = x[:, :, mid:] # 块2:包含后mid个key
# 第一步:计算各块的局部统计量(局部最大值和指数和)
m1 = x1.max(dim=-1, keepdim=True)[0] # 块1最大值,形状 (batch, q_len, 1)
m2 = x2.max(dim=-1, keepdim=True)[0] # 块2最大值
# 第二步:全局最大值(取两块的最大值中的较大者)
m = torch.maximum(m1, m2) # 全局最大值(广播到两块的形状)
# 第三步:计算各块相对于自身局部最大值的指数和
e1 = (x1 - m1).exp().sum(dim=-1, keepdim=True) # 块1相对于m1的指数和
e2 = (x2 - m2).exp().sum(dim=-1, keepdim=True) # 块2相对于m2的指数和
# 第四步:重新缩放后累加,得到全局指数和
# 块1的贡献需要对齐到全局最大值m:乘以 exp(m1 - m)
# 块2的贡献需要对齐到全局最大值m:乘以 exp(m2 - m)
l = e1 * (m1 - m).exp() + e2 * (m2 - m).exp() # 全局指数和,形状 (batch, q_len, 1)
# 第五步:用全局m和l计算两块的最终softmax
s1 = (x1 - m).exp() / l # 第一块所有key的注意力权重
s2 = (x2 - m).exp() / l # 第二块所有key的注意力权重
return torch.cat([s1, s2], dim=-1) # 合并:还原为 (batch, q_len, kv_len)
# 验证在线softmax与标准softmax的数值等价性
B, Q, KV = 2, 512, 4096
x = torch.randn(B, Q, KV) * 2 # 模拟注意力分数分布
result_standard = standard_softmax(x)
result_online = online_softmax_two_pass(x)
max_diff = (result_standard - result_online).abs().max()
print(f"在线softmax与标准softmax的最大误差: {max_diff:.2e}") # 应接近机器精度(1e-6级别)
这个代码演示的是"在线softmax"的核心原理:分块处理注意力分数矩阵(S),通过维护全局最大值 m 和归一化因子 l,实现逐块增量计算 softmax。注意这演示的是通用原理,Flash Attention 在此基础上增加了分块矩阵乘法和 SRAM 优化。
第四节小结
- 在线softmax:通过维护全局最大值 m 和归一化因子 l,实现分块增量计算
- 重新缩放:新块到来时,旧块的贡献需要乘以 exp(m_old - m_new) 进行对齐
- Flash Attention:将在线softmax与分块矩阵乘法结合,实现端到端的高效注意力计算
五、重计算(Recomputation):用算力换显存
What — 反向传播时为什么需要存储中间激活值?
在标准Attention的反向传播中,需要计算 dQ、dK、dV。这要求知道前向传播时的 Q、K、V 和 softmax结果 P。标准实现将这些全部存储在HBM中用于反向计算。
Why — 为什么重计算能降低显存占用?
问题一:Q/K/V是输入,它们的体积有多大?
Q/K/V的形状都是 (batch, num_heads, seq_len, head_dim)。以 GPT-3 配置为例:
# GPT-3级别配置下的激活值显存占用
batch = 1
num_heads = 96
seq_len = 4096
head_dim = 64
dtype = torch.float32 # 4 bytes
# Q/K/V的显存(每个)
qkv_size = batch * num_heads * seq_len * head_dim * 4 # bytes
qkv_gb = qkv_size / 1024**3
print(f"单个Q(或K或V): {qkv_gb:.2f} GB") # 约 0.25 GB
# 存储Q/K/V/P四个矩阵的显存
total_activations = qkv_size * 4 # Q, K, V, P
total_activations_gb = total_activations / 1024**3
print(f"存储Q/K/V/P总显存: {total_activations_gb:.2f} GB") # 约 1 GB(单层)
这还仅仅是单层。40层就是40GB。加上梯度、优化器状态、模型参数,单层Transformer的激活值显存很快就会爆掉。
重计算的核心思想:不存储Q/K/V/P,而是在反向传播时从HBM中读取原始参数,重新计算出需要的中间结果。虽然增加了额外的FWD/BWD计算量(通常约10~20%的额外算力),但激活值显存大幅降低。
没有重计算会发生什么?
- 长序列训练时激活值显存占用远超模型参数本身
- batch_size被迫降低到1甚至更小,训练效率极低
- 部分序列不得不使用梯度检查点(Gradient Checkpointing)的手动实现,但实现复杂且容易出错
Flash Attention的反向传播重计算策略:前向时不存储S和P,只存储必要的统计量(每行的最大值m和归一化因子l)。反向时重新计算Q/K/V及其softmax值,再求梯度。这样激活值显存降为 O(n * d) 而非 O(n^2):
# 反向传播重计算示意
# 标准Attention(存储激活值):
# 前向:存储 S(n,n) + P(n,n) + Q/V(n,d) 显存 O(n^2)
# 反向:直接使用存储的S和P计算梯度
# Flash Attention(重计算):
# 前向:只存储 m(n,) + l(n,) 显存 O(n)(远小于n^2)
# 反向:
# Step 1: 从HBM加载模型权重(W_Q, W_K, W_V)
# Step 2: 用权重和输入X重新计算Q/K/V(前向的中间结果)
# Step 3: 用重计算得到的Q/K/V计算注意力输出梯度 dO -> dQ/dK/dV
# Step 4: 释放临时Q/K/V
# 显存节省估算:
seq_len = 4096
head_dim = 64
dtype_bytes = 4
standard_activations = seq_len * seq_len * 2 * dtype_bytes # S + P
flash_activations = seq_len * 2 * dtype_bytes # m + l
saving_ratio = standard_activations / flash_activations
print(f"激活值显存节省倍数: {saving_ratio:.0f}x")
# n=4096, d=64时:节省约4096/2=2048倍
# 注意:重计算增加了前向计算量(需要再跑一次Q/K/V计算)
# 但这部分开销相比HBM带宽瓶颈带来的减速是值得的
# 实际测试:Flash Attention前向+反向总时间仍比标准Attention快2~4倍
Flash Attention将重计算与分块计算结合:前向传播时用SRAM分块处理,每块计算完毕后直接丢弃(不写回HBM),只把最终结果写入。反向传播时用重计算恢复中间状态。这是其能将显存从 O(n^2) 降到 O(n) 的关键。
第五节小结
- 重计算的目的:用额外的FWD计算量换取激活值显存的大幅降低
- 存储内容变化:标准实现存S(n,n)和P(n,n),Flash Attention只存m(n,)和l(n,)
- 显存收益:激活值从 O(n^2) 降到 O(n),节省约 n/2 倍(n=4096时约2000倍)
- 计算代价:约10~20%的额外FWD计算量,但总时间仍因HBM带宽节省而更优
六、Flash Attention v1/v2/v3演进路径
What — Flash Attention三个版本的区别是什么?
Flash Attention自2022年提出以来,已经历三次主要迭代。三个版本的演进主线是:算法固化(v1)→ 硬件适配优化(v2)→ 低精度量化支持(v3)。
Why — 为什么需要持续优化,每次优化的核心瓶颈是什么?
v1 -> v2 的驱动力:Warp级并行度不足。
Flash Attention v1将注意力按行(Query维度)并行化,每个thread block处理一个或多个Query块。但当序列较短或batch较大时,并行度不足,导致GPU资源利用率低。v2重新设计了block-level并行策略,使warp(32个threads为一组)更高效地协作处理块间归约。
v2 -> v3 的驱动力:FP8计算的低精度加速空间。
H100 GPU引入了FP8 Tensor Core(吞吐量是FP16的两倍)。v3针对FP8进行了专门优化,并引入了异步执行(asynchronous pipeline)来隐藏内存访问延迟。
# Flash Attention各版本性能对比概览
# (基于A100/H100实测数据,非精确数值)
# v1 (2022, "FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness")
# - 核心创新:分块计算 + 在线softmax + 重计算
# - 显存:O(n) vs 标准 O(n^2)
# - 速度:比标准Attention快2~4倍
# - 限制:对GPU硬件利用不够充分,长序列时仍有优化空间
# v2 (2023, "FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning")
# - 核心改进:
# 1. 更好的thread block调度:每个thread block处理更多行,减少同步开销
# 2. Warp级归约优化:减少warp间的shuffle操作
# 3. 支持更长的序列(可达64K+)
# - 速度:比v1再快1.5~2倍
# v3 (2024, "FlashAttention-3: Fast and Accurate Attention with Softmax Tuning and FP8-Interleaved Pipeline")
# - 核心改进:
# 1. FP8量化支持:利用H100 FP8 Tensor Core,理论吞吐翻倍
# 2. Softmax Tuning:对softmax的不同输入分布做数值优化
# 3. 异步Pipeline:重叠内存访问与计算
# - 速度:比v2再快1.5~2倍(H100 FP8场景)
# - 注意:v3依赖H100的新硬件特性,在A100等旧卡上不可用
# 显存收益总结(以seq_len=4096, d=64, FP32为例)
print("显存占用对比(单层Attention, 1个head):")
print(f" 标准实现: S矩阵 + P矩阵 = 2 x 4096^2 x 4B = {2*4096**2*4/1024**2:.0f} MB")
print(f" Flash Attn v2: m + l 统计量 = 2 x 4096 x 4B = {2*4096*4/1024:.1f} KB")
print(f" 节省比例: {(2*4096**2*4) / (2*4096*4):.0f}x")
print(" 加上Q/K/V(必须存储的输入):")
qkv_mem = 3 * 4096 * 64 * 4
print(f" Q/K/V总显存: {qkv_mem/1024**2:.1f} MB(这是无法省掉的理论下限)")
值得注意的是,Flash Attention的精度(数值正确性)与标准Attention完全等价(误差在机器精度范围内),不会引入近似Attention(如Linformer、Performer)中的近似误差。
第六节小结
- v1:奠定理论基础(分块+在线softmax+重计算),显存从O(n^2)降到O(n)
- v2:优化thread block调度和warp级并行,速度比v1快1.5~2倍
- v3:引入FP8和异步Pipeline,进一步提速(依赖H100硬件)
- 精度保证:Flash Attention是精确注意力,与标准实现数学等价,无近似误差
七、PyTorch集成:如何正确使用scaled_dot_product_attention
What — PyTorch 2.0提供的统一Attention API是什么?
PyTorch 2.0引入了 torch.nn.functional.scaled_dot_product_attention(简称SDPA),它是一个统一的注意力计算接口,会根据硬件和上下文自动选择最优实现(标准math实现、Flash Attention、或内存高效Attention)。
Why — 为什么应该优先使用SDPA而非手写注意力?
问题一:手写注意力无法利用硬件加速。
标准 Attention 的实现通常为:
# 手写实现(问题代码)
S = Q @ K.transpose(-2, -1) / math.sqrt(d)
P = F.softmax(S, dim=-1)
O = P @ V
return O
这种实现会产生完整的 S 和 P 矩阵,无论在显存还是速度上都是次优的。
问题二:SDPA会探测硬件能力并自动选择最优路径。
SDPA内部维护了一个选择器,会按优先级尝试:
- Flash Attention(如果可用且合适):最快,显存最优
- Memory Efficient Attention(xformers风格):显存最优,精度略低
- Math实现(fallback):通用但低效
不使用SDPA会发生什么?
- 长序列训练时显存爆炸,被迫降低batch_size
- 即使手动实现了Flash Attention逻辑,也难以匹敌PyTorch内核级别的优化
- PyTorch升级后无法自动享受新的硬件优化
import torch
import torch.nn.functional as F
import math
# 写法1:最简洁,直接使用(自动选择最优实现)
class AttentionModuleV1(torch.nn.Module):
def __init__(self, embed_dim, num_heads):
super().__init__()
self.q_proj = torch.nn.Linear(embed_dim, embed_dim)
self.k_proj = torch.nn.Linear(embed_dim, embed_dim)
self.v_proj = torch.nn.Linear(embed_dim, embed_dim)
self.out_proj = torch.nn.Linear(embed_dim, embed_dim)
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
def forward(self, query, key, value, attn_mask=None, dropout_p=0.0):
# Q/K/V投影
Q = self.q_proj(query) # (B, T, D)
K = self.k_proj(key)
V = self.v_proj(value)
# 适配多头格式:(B, T, D) -> (B, num_heads, T, head_dim)
B, T, _ = Q.shape
Q = Q.view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
K = K.view(B, -1, self.num_heads, self.head_dim).transpose(1, 2)
V = V.view(B, -1, self.num_heads, self.head_dim).transpose(1, 2)
# scaled_dot_product_attention自动选择最优实现
# 支持 is_causal=True 自动生成causal mask,无需手动创建
O = F.scaled_dot_product_attention(
Q, K, V,
attn_mask=attn_mask,
dropout_p=dropout_p if self.training else 0.0,
is_causal=True # 自动生成因果mask(A100上用Flash Attention实现)
)
# 合并多头维度
O = O.transpose(1, 2).contiguous().view(B, T, -1)
return self.out_proj(O)
# 写法2:手动指定实现后端(用于调试或特定场景)
class AttentionModuleV2(torch.nn.Module):
def forward(self, Q, K, V, attn_mask=None):
# 设置只使用Flash Attention(不存在则报错)
with torch.backends.cuda.sdp_kernel(enable_flash=True,
enable_math=False,
enable_mem_efficient=False):
return F.scaled_dot_product_attention(Q, K, V, attn_mask=attn_mask)
# 后端选项说明:
# enable_flash=True: CUDA Flash Attention(基于Flash Attention算法)
# enable_math=True: 标准数学实现(CPU fallback 或调试用)
# enable_mem_efficient=True: xformers风格的内存高效Attention
# 验证两种写法的数值一致性
B, T, D = 2, 512, 64
Q = torch.randn(B, T, D)
K = torch.randn(B, T, D)
V = torch.randn(B, T, D)
model = AttentionModuleV2(D, 1)
result = model(Q.unsqueeze(1), K.unsqueeze(1), V.unsqueeze(1))
print(f"SDPA输出形状: {result.shape}") # torch.Size([2, 1, 512, 64])
print("SDPA已正确执行(具体后端由PyTorch自动选择)")
使用注意事项:
- is_causal=True 会自动生成causal mask,禁用时应传 attn_mask 而非手动在外面加mask
- dropout_p 只在 training=True 时生效,eval模式下传0
- PyTorch 2.0.0+ 支持SDPA,低于此版本需要安装flash-attn包并使用 torch.nn.attention
- HuggingFace Transformers库默认已使用SDPA,替换手动实现后速度普遍提升1.5~3倍
第七节小结
- SDPA:PyTorch 2.0的统一注意力API,自动选择Flash/Math/Efficient最优路径
- is_causal:自动生成causal mask,替代手动的attention_mask
- HuggingFace:已默认使用SDPA,老代码升级后显存和速度双优化
- 后端选择:通过 torch.backends.cuda.sdp_kernel 可手动指定后端(调试用)
关于Flash Attention的20个高频问题
Q1. Flash Attention和标准Attention的结果完全一致吗?
是,Flash Attention是精确注意力算法,与标准Attention数学等价。它通过分块+在线softmax+重计算降低了显存和提升了速度,但数值结果与标准Attention在机器精度范围内完全一致(误差通常 < 1e-6)。与Sparse Attention、Linear Attention等近似方法有本质区别。
Q2. Flash Attention为什么能降低显存占用?
因为它通过三个机制避免了完整S矩阵和P矩阵的存储:分块计算(每次只加载块到SRAM)、在线softmax(增量更新统计量)、重计算(反向时重新算Q/K/V而非存储它们)。存储内容从 S(n,n)+P(n,n) 变为 m(n)+l(n),激活值显存从 O(n^2) 降到 O(n)。
Q3. 在线softmax(Online Softmax)是什么,解决了什么问题?
在线softmax通过维护全局最大值m和归一化因子l,使得softmax可以分块增量计算。新块到来时,旧块的贡献需要乘以 exp(m_old - m_new) 进行对齐,然后累加新块的指数和。这解决了分块计算中无法对独立块做softmax的数学障碍。
Q4. Flash Attention的块大小(Block Size)是怎么确定的?
根据GPU的SRAM大小自动确定(编译时hardcoded,或运行时计算)。块大小需要满足:每次加载到SRAM的Q块(Br*d)和K/V块(Bc*d)加上中间计算结果不超过SRAM容量。A100的建议值约为 Br=128, Bc=64(不同实现略有差异)。
Q5. SRAM和HBM的区别是什么?
SRAM是GPU片上存储,容量小(约20~50MB)但带宽极高(PB/s级);HBM是显存,容量大(80GB)但带宽相对低(TB/s级)。Flash Attention利用SRAM的高带宽进行分块计算,减少对HBM的频繁读写,从而突破访存瓶颈。
Q6. 重计算(Recomputation)会增加多少计算量?
约增加10~20%的计算量(主要是重新计算Q/K/V投影)。但由于避免了O(n^2)级别的HBM读写(这对速度影响更大),总的前向+反向时间仍比标准Attention快2~4倍。
Q7. Flash Attention 1和Flash Attention 2的核心区别是什么?
v2改进了thread block调度和warp级并行度,充分利用GPU硬件资源。v1按行并行度不够高(特别是短序列场景),v2重新设计了工作分区策略,使每个thread block处理更多行,减少了同步开销,速度比v1快1.5~2倍。
Q8. Flash Attention 3有哪些新特性?
v3引入FP8量化支持和异步执行pipeline,充分利用H100的新硬件特性。FP8 Tensor Core的吞吐量是FP16的两倍,异步执行可以重叠内存访问与计算。但v3依赖H100硬件,在A100等旧卡上不可用,速度提升约1.5~2倍(相比v2)。
Q9. 使用torch.nn.functional.scaled_dot_product_attention有什么前提条件?
需要PyTorch 2.0以上版本,且需要对应的CUDA扩展(Flash Attention后端需要CUDA 11.6+ 和支持compute capability 8.0+的GPU)。安装 flash-attn 包可以获得更完整的Flash Attention支持,但PyTorch内置的SDPA会优先使用已有的最优实现。
Q10. Flash Attention支持变长序列(Variable Length)吗?
支持,Flash Attention原生支持pack多个序列变长batch(通过cu_seqlens参数)。这避免了padding mask带来的无效计算,比手动padding后计算Attention更高效。HuggingFace DataLoader默认返回padded batch,但配合SDPA可以有效处理变长序列。
Q11. Flash Attention支持自定义mask(如padding mask、causal mask)吗?
支持。PyTorch SDPA通过attn_mask参数接受任意布尔/二进制mask,通过is_causal参数自动生成causal mask。Flash Attention的kernel内部处理mask的效率远高于Python层面的手工mask(后者会产生完整mask矩阵并物化到显存)。
Q12. Flash Attention能用于推理(Inference)阶段吗?
能,推理阶段Flash Attention主要用于长上下文场景(如长文档摘要、代码补全)。推理的特殊性在于自回归生成时每次只计算一个token的query与已生成KV的attention。Flash Attention配合KV Cache(PagedAttention等)使用效果更佳。
Q13. Flash Attention和KV Cache是什么关系?
互补关系,共同用于加速推理。Flash Attention优化单次attention计算;KV Cache存储历史token的K和V以避免重复计算。PagedAttention(vLLM的核心技术)进一步优化了KV Cache的显存管理,两者结合可以实现超长上下文的高效推理。
Q14. Flash Attention对多头注意力(Multi-Head Attention)有什么特殊要求吗?
没有特殊要求,Flash Attention按head维度独立计算,最终拼接效果与标准MHA一致。每个head的attention可以独立并行,Flash Attention kernel通常以head作为thread block的并行维度,因此head数越多(如96头)并行度越高。
Q15. Flash Attention支持Cross-Attention吗?
支持,Flash Attention天然支持Q与K/V来源不同的Cross-Attention(如Encoder-Decoder架构)。只需要在调用时传入不同的K/V矩阵即可。Flash Attention的tile机制对Cross-Attention同样有效(只是Q的序列长度与K/V的序列长度可以不同)。
Q16. 显存足够的情况下还需要Flash Attention吗?
即使显存充足,Flash Attention仍能显著加速训练和推理。原因:HBM带宽是瓶颈,Flash Attention减少了HBM读写次数,即使完整S矩阵能装下,分块计算仍然更快。实测中Flash Attention通常比标准Attention快2~4倍,且这个加速在所有序列长度下都成立。
Q17. Flash Attention和Memory Efficient Attention有什么区别?
两者目标相同(降低attention显存),但实现路径不同。Memory Efficient Attention通过近似手段(低秩分解等)降低计算复杂度;Flash Attention保持精确性。PyTorch SDPA会同时支持两者,按需选择。Memory Efficient Attention显存更低(理论最优),但精度可能有损失。
Q18. Flash Attention在分布式训练(多卡)中如何使用?
分布式训练中,每个GPU负责处理不同的sequence或不同的head,Flash Attention在每个GPU内部独立运行。对于Tensor Parallel(张量并行)场景,每个GPU只有部分Q/K/V,需要特别注意attention的并行分区设计(FasterTransformer、DeepSpeed等框架已集成)。
Q19. Flash Attention的反向传播梯度正确性如何保证?
Flash Attention的反向传播通过重计算恢复中间状态,严格按照链式法则计算梯度。重计算保证了 dQ/dK/dV 与标准实现数学等价。Flash Attention论文中提供了完整的反向传播公式推导,PyTorch的SDPA在底层kernel中已实现正确的反向逻辑。
Q20. Flash Attention的下一步演进方向是什么?
主要方向是更激进的低精度(FP4/NF4)、更长的上下文支持(1M+ token)、以及稀疏/混合专家变体。Flash Attention-v3的FP8实践已经证明低精度+数值调优的可行性。Ring Attention将Flash Attention扩展到多节点协同计算超长序列。Flash Attention的core思想(IO感知设计)也被应用到Linear Attention、Mamba等状态空间模型中。
全篇总纲
- 显存瓶颈的根源:标准Attention的S矩阵和P矩阵都需要完整存储,O(n^2)显存
- Flash Attention三板斧:分块计算(避免物化S)、在线softmax(分块增量归一化)、重计算(用算力换显存)
- 存储层级设计:SRAM小容量高带宽,HBM大容量低带宽,分块让数据尽量在SRAM中流转
- v1/v2/v3演进:v1建立算法基础,v2优化并行度,v3引入FP8(依赖H100)
- PyTorch集成:SDPA API自动选择最优实现,应作为默认选择替代手写attention
下篇预告:RoPE旋转位置编码——从绝对位置到相对位置的突破
位置编码是Transformer理解序列顺序的关键。RoPE(Rotary Position Embedding)通过旋转矩阵将位置信息编码进Q/K的相对相位,让Attention天然具备相对位置感知能力,且无需额外的位置偏置。本系列后续将深入讲解RoPE的数学推导、与Sinusoidal PE的对比、以及在LLaMA等主流大模型中的应用。

浙公网安备 33010602011771号