直接不能。原因很朴素:
FlopCounterMode 通过拦截 ATen 算子来计数,而 FlashAttention 是一个融合自定义 CUDA kernel,直接绕过标准 PyTorch 算子分发。
结果通常是:
sdpa 统计,数值偏离真实 FlashAttention 实现;因此,当模型使用 FlashAttention 时,不要依赖自动工具统计 attention 部分,而是用手动公式替代或校准。
FlashAttention 的 attention 计算量与普通 attention 的理论值相同(它做的是“计算结果等价”,只是实现上更省内存)。标准 attention FLOPs:
b:batch sizeh:头数t:序列长度d:每头维度自回归 causal FlashAttention 要减半:
注意:这里只算 attention 计算本身,不含 Q/K/V 投影(那部分属于线性层,按 2 × params × tokens 算即可)。
不同工具对 backward 的处理差异很大。
DeepSpeed FlopsProfiler 不直接测量 backward FLOPs,而是按理论近似:
它实际测量的是 backward latency,然后用这个公式换算成 bwd FLOPS(每秒运算量)。如果你要总训练 FLOPs,DeepSpeed 会输出 fwd + bwd = 3 × fwd FLOPs。
FlopCounterMode 实际上可以统计反向传播,只要在 context 里调用 .backward():
import torch
from torch.utils.flop_counter import FlopCounterMode
from transformers import AutoModel
model = AutoModel.from_pretrained("bert-base-uncased")
model.train()
x = torch.randint(0, 1000, (4, 128))
with FlopCounterMode(model, display=False) as fcm:
out = model(x, output_hidden_states=False)
# 假设做 MLM / classification,构造一个标量 loss
loss = out.last_hidden_state.sum()
loss.backward()
fwd_bwd_flops = fcm.get_total_flops()
print(f"Forward + Backward FLOPs: {fwd_bwd_flops / 1e9:.2f} GFLOPs")
from calflops import calculate_flops
flops, macs, params = calculate_flops(
model=model,
input_shape=(4, 128),
include_backPropagation=True, # 启用 backward 估算
compute_bp_factor=2.0, # backward 按 forward 的 2 倍算
)
print(f"FLOPs (fwd+bwd): {flops}")
compute_bp_factor=2.0 是默认值,对应 dense Transformer 的常规近似。如果你的模型用了 FlashAttention 或 activation checkpointing,需要手动调整这个系数。
import torch
from torch.profiler import profile, ProfilerActivity
model.train()
x = torch.randint(0, 1000, (4, 128))
with profile(
activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
with_flops=True,
record_shapes=True
) as prof:
out = model(x).last_hidden_state.sum()
out.backward()
# 按 FLOPs 排序看热点
print(prof.key_averages().table(sort_by="flops", row_limit=10))
# 简单汇总(注意 None 值)
total = sum(e.flops for e in prof.key_averages() if e.flops is not None)
print(f"Total FLOPs: {total / 1e9:.2f} GFLOPs")
PyTorch Profiler 的 with_flops=True 会尝试给每个算子标注 FLOPs,但覆盖范围比 FlopCounterMode 还有限,很多自定义算子会显示为 None。
| 方法 | backward 口径 | 适用场景 | 注意点 |
|---|---|---|---|
| DeepSpeed | 估算 2× forward | 训练全流程 profiling、MFU 报告 | 符合 MFU 理论口径;不含 activation checkpointing 重计算 |
| FlopCounterMode + backward | 实测 autograd FLOPs | 验证公式、定位反向热点 | 会包含重计算;自定义 kernel 不计 |
| calflops | 参数化估算(默认 2×) | 快速估算 HuggingFace 模型训练成本 | 系数需根据实际模型调整 |
| PyTorch Profiler | 实测,但覆盖不全 | 与 latency 一起分析热点 | 很多算子 FLOPs 为 None;更适合看耗时而非精确 FLOPs |
import torch
from transformers import AutoModel, AutoConfig
from torch.utils.flop_counter import FlopCounterMode
# 假设模型用了 FlashAttention,需要手动加 attention FLOPs
config = AutoConfig.from_pretrained("bert-base-uncased")
model = AutoModel.from_pretrained("bert-base-uncased")
model.train()
b, t = 4, 128
h, d = config.num_attention_heads, config.hidden_size // config.num_attention_heads
x = torch.randint(0, 1000, (b, t))
# 1) 用 FlopCounterMode 跑 forward + backward
with FlopCounterMode(model, display=False) as fcm:
out = model(x).last_hidden_state.sum()
out.backward()
measured = fcm.get_total_flops()
# 2) 手动补 FlashAttention attention 计算量(如果工具漏算)
# causal attention: 2 * b * h * t^2 * d
flash_attn_flops = 2 * b * h * t * t * d
# 3) 总训练 FLOPs = measured + 手动校准项
total = measured + flash_attn_flops
print(f"Measured: {measured/1e9:.2f} GFLOPs")
print(f"FlashAttention add-on: {flash_attn_flops/1e9:.2f} GFLOPs")
print(f"Total: {total/1e9:.2f} GFLOPs")