第 4 课:FlashAttention 与 backward FLOPs 的统计方法

当自动工具失灵时,如何手动补算并正确统计反向传播

1. FlopCounterMode 能统计 FlashAttention 的 FLOPs 吗?

直接不能。原因很朴素:

FlopCounterMode 通过拦截 ATen 算子来计数,而 FlashAttention 是一个融合自定义 CUDA kernel,直接绕过标准 PyTorch 算子分发。

结果通常是:

因此,当模型使用 FlashAttention 时,不要依赖自动工具统计 attention 部分,而是用手动公式替代或校准。

2. FlashAttention FLOPs 的手动公式

FlashAttention 的 attention 计算量与普通 attention 的理论值相同(它做的是“计算结果等价”,只是实现上更省内存)。标准 attention FLOPs:

Attention FLOPs = 4 · b · h · t² · d

自回归 causal FlashAttention 要减半:

Causal Attention FLOPs ≈ 2 · b · h · t² · d

注意:这里只算 attention 计算本身,不含 Q/K/V 投影(那部分属于线性层,按 2 × params × tokens 算即可)。

3. backward FLOPs 怎么统计?

不同工具对 backward 的处理差异很大。

方法一:DeepSpeed(估算为 2× forward)

DeepSpeed FlopsProfiler 不直接测量 backward FLOPs,而是按理论近似:

bwd FLOPs ≈ 2 × fwd FLOPs

它实际测量的是 backward latency,然后用这个公式换算成 bwd FLOPS(每秒运算量)。如果你要总训练 FLOPs,DeepSpeed 会输出 fwd + bwd = 3 × fwd FLOPs

方法二:PyTorch FlopCounterMode(跑 actual backward)

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")
注意: 这样跑出来的是实际 autograd 产生的 FLOPs。如果用了 activation checkpointing,二次前向也会被计入,结果会大于“MFU 口径”的 6N。

方法三:calflops(通过参数估算 backward)

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,需要手动调整这个系数。

方法四:PyTorch Profiler(实测,但只给部分算子)

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。

4. 四种方法怎么选?

方法backward 口径适用场景注意点
DeepSpeed 估算 2× forward 训练全流程 profiling、MFU 报告 符合 MFU 理论口径;不含 activation checkpointing 重计算
FlopCounterMode + backward 实测 autograd FLOPs 验证公式、定位反向热点 会包含重计算;自定义 kernel 不计
calflops 参数化估算(默认 2×) 快速估算 HuggingFace 模型训练成本 系数需根据实际模型调整
PyTorch Profiler 实测,但覆盖不全 与 latency 一起分析热点 很多算子 FLOPs 为 None;更适合看耗时而非精确 FLOPs

5. 一个完整示例:BERT + backward + FlashAttention 校准

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")
核心原则: 自动工具算不出时,用解析公式补;工具与公式差距大时,先对齐“到底算的是 forward 还是 forward+backward、包不包含重计算、 causal 有没有砍半”。

小测验

6. 推荐延伸阅读

7. 接下来可以问的问题