第 2 课:FLOPs 计数内幕与 CNN/ViT 拆解

从 6N 的推导,到 DeepSpeed / transformers 的实现,再到 CNN 与 ViT 的逐项拆解

1. 先纠正一个常见误解:不是“两个乘加”

在上节课里我说“每个参数参与 6 次浮点运算”,有同学会理解为“两个乘加”。这里必须先把单位掰清楚:

1 MAC(multiply-accumulate)= 1 次乘法 + 1 次加法 = 2 FLOPs。

神经网络里的矩阵乘、卷积,本质上都是由 MAC 组成的。我们说的 6 FLOPs 其实是 3 MACs

为什么训练 Transformer 是 3 MACs(6 FLOPs)每参数每 token?

以线性层 y = xW 中的一个权重 W[i,j] 为例:

一个权重在训练的一步里被用了 3 次,每次 1 MAC,所以:

3 MACs = 6 FLOPs per parameter per token

注意:attention 里的 QK^T、softmax、加权求和本身没有可训练权重(除了 Q/K/V 的投影矩阵已计入 N),所以它们的 FLOPs 要单独列出来,就是 12lhqt 那一项。

2. DeepSpeed Flops Profiler 怎么数?

DeepSpeed 的 Flops Profiler 本质上是一个 hook-based profiler

  1. 在 PyTorch 的每个模块上注册 forward / backward hook。
  2. 根据模块类型(Linear、Conv、MatMul、BatchNorm 等)查一张预设计算表。
  3. 把输入输出形状代入公式,累计 MACs,再乘以 2 得到 FLOPs。
  4. 同时记录每层的 forward / backward latency,算出 FLOPS(单位时间运算量)。

它输出的典型报告包含:

DeepSpeed 的边界: 它依赖预设的算子表。遇到自定义算子、融合算子(fused kernel)、共享权重、或者 activation checkpointing 的二次前向时,统计可能偏离理论值。

3. HuggingFace transformers 怎么算?

transformers 库没有一个官方统一的 FLOPs API。不同版本、不同示例代码里口径不一:

FlopCounterMode 的原理:在 forward 执行时拦截 ATen 算子调用(如 aten::mmaten::addmmaten::convolution),按算子形状估算 MACs。它只覆盖被注册的常见算子,对自定义 CUDA kernel 无能为力。

4. 有没有“完全准确”的方案?

没有。业界做法是“两套口径并用”:

方法优点缺点用途
解析公式(6N + 12lhqt 等) 快、可重复、与实现无关 忽略框架 overhead、融合算子、内存搬运、通信 MFU 横向对比、论文复现
运行时 profiler(DeepSpeed / FlopCounterMode) 反映真实执行、能定位热点 依赖具体实现、可能漏算自定义 op、backward 重计算口径混乱 优化热点、验证公式

所以更准确的说法是:

通用方案 = 解析公式做估算 + 运行时 profiler 做验证。 两者差距过大时,就是你的实现里有 overhead 或 profiler 口径没对齐。

5. CNN 的 FLOPs 逐项拆解

以输入 (B, C_in, H_in, W_in)、输出 (B, C_out, H, W) 的 Conv2d 为例:

Conv2d FLOPs = 2 · H · W · C_in · K² · C_out

推导:输出特征图每个位置有 C_out 个通道;每个通道的值是一次 C_in × K × K 的卷积;每次卷积是 C_in · K² 次 MAC;输出图大小是 H · W;1 MAC = 2 FLOPs。

BatchNorm

BN FLOPs ≈ 2 · H · W · C

主要是减 mean、除 std、乘 gamma、加 beta,约等于每个元素 ~1 MAC(或 2 FLOPs)的读写。很多 FLOPs 工具会忽略 BN,因为它占比通常很小。

ReLU / Activation

ReLU FLOPs ≈ H · W · C

ReLU 是逐元素的比较/置零,严格说不是浮点乘加,但部分工具按 1 FLOP 估算。

Pooling

Max/Avg Pooling FLOPs ≈ H · W · C · K²

Max pooling 主要是比较,Avg pooling 是加法平均。不同工具统计方式差异较大。

全连接层(FC)

FC FLOPs = 2 · I · O

输入维度 I,输出维度 O,标准矩阵乘。

6. ViT 的 FLOPs 逐项拆解

ViT 把图像切成 n 个 patch,每个 patch 变成维度为 d 的 token,然后接 l 层 Transformer。

Patch Embedding

Patch Embedding FLOPs = 2 · n · P² · C · d

其中 P² · C 是每个 patch 的像素数,d 是 embedding 维度。也可以理解为一个卷积核为 P、步幅为 P 的卷积。

Multi-Head Self-Attention(MSA)

MSA FLOPs = 8 · l · n · d² + 4 · l · n² · d

拆解:

FFN

FFN FLOPs = 16 · l · n · d²

典型 expansion ratio 为 4:先 d → 4d4d → d,两层线性变换合计 2 · n · d · 4d + 2 · n · 4d · d = 16nd²。某些资料按 MACs 口径简写为 8lnd²,注意上下文单位。

7. 一个对比视角

组件CNN(以 ResNet 为例)ViTLLM
主要算子Conv、BN、ReLU、PoolingLinear、Softmax、LayerNormLinear、Softmax、LayerNorm、RoPE/ALiBi
FLOPs 主导项卷积层:2HWC_inK²C_outMSA + FFN:~24lnd² + 4ln²d(FLOPs);常见 MACs 写法为 ~12lnd² + 2ln²d6N + 12lhqt(causal 减半)
容易忽略的部分BN、ReLU、skip connectionPatch embed、CLS token、position embedCausal mask、activation checkpointing、通信

小测验

8. 推荐延伸阅读

9. 接下来可以问的问题