MFU 与 FLOPs 计算速查表

训练效率评估的核心公式、架构差异与工具选型

1. 核心定义

MFU = (模型每秒所需 FLOPs) ÷ (GPU 理论峰值 FLOPs/s)

等价地,按单步计算:

MFU = (每步模型 FLOPs) ÷ (单卡峰值 × 卡数 × 每步耗时)

2. Transformer / LLM 训练 FLOPs

PaLM 论文给出的 dense Transformer 每 token 公式(未考虑 causal mask):

FLOPs/token = 6N + 12 · l · h · q · t

自回归语言模型使用 causal mask,attention 计算量可近似减半:

FLOPs/token ≈ 6N + 6 · l · h · q · t
为什么是 6N? 1 次乘加(MAC)= 2 FLOPs。训练一次前向+反向中,每个参数大约参与 3 次 MAC:前向 1 次、激活梯度 1 次、权重梯度 1 次,共 6 FLOPs。

3. CNN 与 ViT 的 FLOPs 计算

操作FLOPs说明
Conv2d 2 · H · W · C_in · K² · C_out H×W 为输出特征图大小,K 为 kernel 边长;乘加拆成 2 FLOPs
Depthwise Separable Conv 2 · H · W · C · (K² + C_out) 先做 depthwise 再做 pointwise,显著降低计算量
ViT Self-Attention 8 · l · n · d² + 4 · l · n² · d n 为 patch/token 数,d 为模型维度;常见 MACs 写法为 4lnd² + 2ln²d
ViT FFN 16 · l · n · d² 典型 expansion ratio 为 4(d→4d→d),两层合计 2nd·4d + 2n·4d·d = 16nd²;常见 MACs 写法为 8lnd²

4. 三种架构的核心差异

维度传统 CNNViTLLM
主要决定因素 输入分辨率、通道数、kernel 大小 patch 数量(n² 关系) 序列长度 t、参数量 N
全局依赖 局部感受野,参数共享 全局 self-attention 全局 causal self-attention
FLOPs 增长瓶颈 分辨率 × 通道平方 token 数二次方 attention 随 t²;总 FLOPs 随 N 线性
参数复用 每个卷积核在 spatial 上复用 同一 patch embedding 复用 同一权重矩阵在每个 token 上复用
典型 FLOPs 公式 ∑ 2HWC_inK²C_out MSA + FFN:~24lnd² + 4ln²d 6N + 12lhqt(causal 时减半)

5. 开源工具选型

工具适用场景优点注意点
calflops 快速估算 LLM / CNN / ViT 的 FLOPs 严格区分 FLOPs 与 MACs;支持 HuggingFace 与 tokenizer 较新,社区生态不如 fvcore
DeepSpeed Flops Profiler 训练全流程 profiling 模块级、前向/反向/耗时全覆盖 需要接入 DeepSpeed 训练循环
fvcore Detectron2 / CV 模型 Meta 官方,生态成熟 返回的是 MACs 但标签为 FLOPs
ptflops 标准 CNN 快速估算 API 简单 只返回 MACs;LLM 支持弱

6. 业界工具到底是怎么数的?

工具计数方式覆盖范围准确性边界
DeepSpeed Flops Profiler PyTorch module hook + 算子级公式表 forward / backward / latency / 模块级 依赖预设算子表,自定义 op、融合算子、共享权重可能统计不准
HuggingFace transformers 早期无统一 API;新版通过集成 PyTorch FlopCounterMode 或 estimate_tokens 估算 mainly forward,部分 helper 支持训练估算 口径不统一,复杂模型常需借助 calflops / DeepSpeed
calflops / ptflops / fvcore hook 式前向遍历,按模块类型查表 主要 forward,部分支持 backward 估算 忽略框架 overhead;MACs/FLOPs 标签混乱
有没有“完全准确”的方案? 没有。解析公式快但近似;运行时 profiler 真实但依赖实现。工程上通常“公式算 MFU 做横向对比,profiler 抓热点做优化”。

7. 常见坑

8. 快速公式卡

MFU = (6N + 12lhqt) × tokens/s ÷ (单卡峰值 × 卡数)
CNN Conv FLOPs = 2 × H × W × C_in × K² × C_out
ViT Attention FLOPs ≈ 8lnd² + 4ln²d
ViT FFN FLOPs ≈ 16lnd²