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
- N:非 embedding 参数量
- l:层数(layers)
- h:注意力头数(heads)
- q:每个头的维度(head dim)
- t:序列长度(sequence length)
自回归语言模型使用 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. 三种架构的核心差异
| 维度 | 传统 CNN | ViT | LLM |
| 主要决定因素 |
输入分辨率、通道数、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. 常见坑
- MACs ≠ FLOPs:1 MAC = 2 FLOPs。看到工具输出先确认是 MACs 还是 FLOPs。
- GPU Utilization 会骗人:100% 利用率可能只是在搬运内存或空转。
- MFU 不算重计算:activation checkpointing 的二次前向不计入 MFU,但计入 HFU。
- Causal mask 要砍半:自回归 attention 计算量约为全 attention 的一半,否则 MFU 可能超过 100%。
- 峰值精度要对应:H100 SXM BF16 峰值 989 TFLOPS,A100 SXM 仅 312 TFLOPS。FP32 峰值约为同代 BF16 的 1/2 ~ 1/3。
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²