线性模型的计算量(FLOPs)怎么算?
什么是 FLOP?
一个 FLOP(Floating Point Operation)就是一次基本的浮点运算,比如:
- 一次加法:x + y
- 一次乘法:x * y
矩阵乘法里,每个输出元素都涉及"先乘后加",所以通常把一次乘加算作 2 个 FLOP(1 次乘法 + 1 次加法)
手算矩阵乘法 FLOPs
假设
text
x 形状 (B=2, D=3)
w 形状 (D=3, K=2)
y = x @ w,形状 (B=2, K=2)
具体数值:
text
x = [[1, 2, 3],
[4, 5, 6]]
w = [[1, 0],
[0, 1],
[1, 1]]
计算 y[0, 0]:
text
y[0,0] = x[0,0]*w[0,0] + x[0,1]*w[1,0] + x[0,2]*w[2,0]
= 1*1 + 2*0 + 3*1
= 1 + 0 + 3
= 4
这里有:
- 3 次乘法:11、20、3*1
- 2 次加法:1+0、1+3
- 总共 5 次浮点运算,约等于 2D = 6(D=3,2D=6,实际是 5,因为加法比乘法少一次)
对于每个输出元素 y[i, k],都需要 D 次乘法和 D−1 次加法,约 2D 个 FLOP。
输出 y 有 B × K 个元素,所以总 FLOPs:
FLOPs=B×K×2D=2×B×D×K \text{FLOPs} = B \times K \times 2D = 2 \times B \times D \times K FLOPs=B×K×2D=2×B×D×K
代入真实数字
文档里的例子:
text
B = 16384
D = 32768
K = 8192
计算:
FLOPs=2×16384×32768×8192 \text{FLOPs} = 2 \times 16384 \times 32768 \times 8192 FLOPs=2×16384×32768×8192
先算:
text
16384 × 32768 = 536,870,912
536,870,912 × 8192 = 4,398,046,511,104
× 2 = 8,796,093,022,208
约 8.8 × 10¹² FLOPs,即 8.8 TFLOPs。
一个重要洞察:FLOPs ≈ 2 × 数据量 × 参数量
注意:
D × K正好是权重矩阵 w 的元素个数,也就是模型参数量 P;B是数据点数量(在语言模型里就是 token 数量)。
所以:
FLOPs≈2×B×P \text{FLOPs} \approx 2 \times B \times P FLOPs≈2×B×P
这就是"前向传播 FLOPs ≈ 2 × token 数 × 参数量"的由来。
例如 GPT-3 有 1750 亿参数,训练用了 3000 亿 token:
前向 FLOPs≈2×3×1011×1.75×1011=1.05×1023 \text{前向 FLOPs} \approx 2 \times 3 \times 10^{11} \times 1.75 \times 10^{11} = 1.05 \times 10^{23} 前向 FLOPs≈2×3×1011×1.75×1011=1.05×1023
训练还要反向传播,总 FLOPs 约前向的 3 倍,即:
≈3.15×1023 \approx 3.15 \times 10^{23} ≈3.15×1023
GPT-3 训练约 3.14×102310^{23}1023 FLOPs 吻合。
MFU(Model FLOPs Utilization)是什么?
定义
MFU 衡量的是:你实际跑出来的计算速度,占硬件理论峰值的比例。
MFU=实际 FLOP/s硬件峰值 FLOP/s \text{MFU} = \frac{\text{实际 FLOP/s}}{\text{硬件峰值 FLOP/s}} MFU=硬件峰值 FLOP/s实际 FLOP/s
- 实际 FLOP/s = 总 FLOPs ÷ 实际耗时
- 硬件峰值 FLOP/s = 厂商公布的该精度下的理论最大速度
小例子
假设:
- 矩阵乘法总 FLOPs = 8.8×101210^{12}1012
- 实际跑完用了 0.05 秒
- 硬件是 H100,bf16 峰值约 990×101210^{12}1012 FLOP/s
实际速度:
实际 FLOP/s=8.8×10120.05=1.76×1014=176 TFLOP/s \text{实际 FLOP/s} = \frac{8.8 \times 10^{12}}{0.05} = 1.76 \times 10^{14} = 176 \text{ TFLOP/s} 实际 FLOP/s=0.058.8×1012=1.76×1014=176 TFLOP/s
MFU:
MFU=176990≈0.178=17.8% \text{MFU} = \frac{176}{990} \approx 0.178 = 17.8\% MFU=990176≈0.178=17.8%
怎么理解这个数字?
| MFU | 含义 |
|---|---|
| ≥ 0.5 | 相当不错 |
| 接近 1.0 | 极难,几乎不可能 |
| 0.1 ~ 0.3 | 常见,尤其小矩阵或访存瓶颈 |
| < 0.1 | 可能矩阵太小,或没用好硬件 |
为什么很难接近 1.0?因为 GPU 不可能 100% 满负荷:
- 要读内存、写内存;
- 有 kernel 启动开销;
- 有通信开销(多卡训练);
- 有同步等待。
为什么 bf16 的 MFU 通常比 fp32 高?
因为硬件对低精度做了专门优化:
- H100 的 fp32 峰值约 67 TFLOP/s;
- H100 的 bf16 峰值约 990 TFLOP/s(甚至 1979 带稀疏)。
同样一个矩阵乘法,bf16 的理论上限高得多,实际跑出来也更快,所以 MFU 往往更高。
注意点
MFU 公式忽略了通信和系统开销,只看纯计算效率。如果多卡训练通信很慢,MFU 会低,但这不代表计算本身有问题。
反向传播(梯度)的 FLOPs 怎么算?
例子:两层线性模型
文档里的例子:
text
x --w1--> h1 --w2--> h2 -> loss
形状:
text
x: (B, D)
w1: (D, D)
h1: (B, D)
w2: (D, K)
h2: (B, K)
loss = mean(h2²)
前向 FLOPs:
- 第一层
x @ w1:2×B×D×D2×B×D×D2×B×D×D - 第二层
h1 @ w2:2×B×D×K2×B×D×K2×B×D×K
总计:
前向 FLOPs=2BD2+2BDK \text{前向 FLOPs} = 2BD^2 + 2BDK 前向 FLOPs=2BD2+2BDK
反向传播要算哪些梯度?
调用 loss.backward() 后,PyTorch 需要算:
h2.grad:loss 对 h2 的梯度w2.grad:loss 对 w2 的梯度h1.grad:loss 对 h1 的梯度w1.grad:loss 对 w1 的梯度
逐个分析 FLOPs
(1)算 h2.grad
loss = mean(h2²),对 h2 求导:
text
h2.grad = d(loss)/d(h2) = h2 / (B*K) (大致)
这是逐元素操作,FLOPs 约 O(BK),相比矩阵乘法可忽略。
(2)算 w2.grad
text
w2.grad = h1.T @ h2.grad
形状:
text
h1.T: (D, B)
h2.grad: (B, K)
w2.grad: (D, K)
这是一个矩阵乘法,FLOPs:
2×D×B×K=2BDK 2 \times D \times B \times K = 2BDK 2×D×B×K=2BDK
(3)算 h1.grad
链式法则:
text
h1.grad = h2.grad @ w2.T
形状:
text
h2.grad: (B, K)
w2.T: (K, D)
h1.grad: (B, D)
又是一个矩阵乘法,FLOPs:
2×B×K×D=2BDK 2 \times B \times K \times D = 2BDK 2×B×K×D=2BDK
所以第二层的反向总共:
2BDK+2BDK=4BDK 2BDK + 2BDK = 4BDK 2BDK+2BDK=4BDK
正好是第二层前向 2BDK2BDK 的 2 倍。
(4)算 w1.grad
text
w1.grad = x.T @ h1.grad
形状:
text
x.T: (D, B)
h1.grad: (B, D)
w1.grad: (D, D)
FLOPs:
2×D×B×D=2BD2 2 \times D \times B \times D = 2BD^2 2×D×B×D=2BD2
(5)算 x.grad(本例跳过)
text
x.grad = h1.grad @ w1.T
FLOPs 也是 2BD22BD^{2}2BD2但本例 x 没有 requires_grad=True,PyTorch 跳过这一步。
总反向 FLOPs
本例:
反向 FLOPs=2BD2⏟w1.grad+4BDK⏟w2相关 \text{反向 FLOPs} = \underbrace{2BD^2}{w1.grad} + \underbrace{4BDK}{w2相关} 反向 FLOPs=w1.grad 2BD2+w2相关 4BDK
前向:
前向 FLOPs=2BD2+2BDK \text{前向 FLOPs} = 2BD^2 + 2BDK 前向 FLOPs=2BD2+2BDK
比值:
反向前向=2BD2+4BDK2BD2+2BDK \frac{\text{反向}}{\text{前向}} = \frac{2BD^2 + 4BDK}{2BD^2 + 2BDK} 前向反向=2BD2+2BDK2BD2+4BDK
当 K 和 D 同量级时,这个比值约等于:
6BD24BD2=1.5 \frac{6BD^2}{4BD^2} = 1.5 4BD26BD2=1.5
但这是本例首层输入不求梯度的特殊情况。在真实多层网络中,每一层反向都包含"权重梯度"和"输入梯度"两个矩阵乘法,各等于该层前向,所以每层反向 = 2 × 前向。
因此整体:
反向 FLOPs≈2×前向 FLOPs \text{反向 FLOPs} \approx 2 \times \text{前向 FLOPs} 反向 FLOPs≈2×前向 FLOPs
训练总 FLOPs ≈ 6 × token 数 × 参数量
前向:
前向≈2×B×P \text{前向} \approx 2 \times B \times P 前向≈2×B×P
反向:
反向≈2×前向≈4×B×P \text{反向} \approx 2 \times \text{前向} \approx 4 \times B \times P 反向≈2×前向≈4×B×P
总:
训练总 FLOPs=前向+反向≈2BP+4BP=6BP \text{训练总 FLOPs} = \text{前向} + \text{反向} \approx 2BP + 4BP = 6BP 训练总 FLOPs=前向+反向≈2BP+4BP=6BP
这就是著名的 6× 规则:
训练一次总 FLOPs≈6×token 数×参数量 \boxed{\text{训练一次总 FLOPs} \approx 6 \times \text{token 数} \times \text{参数量}} 训练一次总 FLOPs≈6×token 数×参数量
验证文档中的例子
文档:
text
B = 16384, D = 32768, K = 8192
前向:
text
2BD² = 2 × 16384 × 32768² ≈ 3.52 × 10¹³
2BDK = 2 × 16384 × 32768 × 8192 ≈ 8.80 × 10¹²
前向总计 ≈ 4.40 × 10¹³
反向:
text
2BD² = 3.52 × 10¹³ (w1.grad)
4BDK = 1.76 × 10¹³ (w2相关)
反向总计 ≈ 5.28 × 10¹³
总:
text
4.40 × 10¹³ + 5.28 × 10¹³ ≈ 9.68 × 10¹³
而 6× 规则:
text
6 × B × P,其中 P ≈ D² + DK = 32768² + 32768×8192 ≈ 1.34 × 10⁹
6 × 16384 × 1.34 × 10⁹ ≈ 1.32 × 10¹⁴
量级一致。差异来自首层输入不求梯度、以及 K 和 D 不同导致的比例偏差。
总结
| 概念 | 通俗解释 | 公式 |
|---|---|---|
| FLOPs | 做了多少次浮点运算,衡量"干了多少活" | 矩阵乘法:2×B×D×K2 \times B \times D \times K2×B×D×K |
| FLOP/s | 每秒能做多少次浮点运算,衡量"干得多快" | 实际 FLOP/s = 总 FLOPs ÷ 耗时 |
| MFU | 实际速度占硬件峰值的比例,衡量"硬件用得多满" | 实际 FLOP/s硬件峰值 FLOP/s\frac{\text{实际 FLOP/s}}{\text{硬件峰值 FLOP/s}}硬件峰值 FLOP/s实际 FLOP/s |
| 反向 FLOPs | 反向传播的运算量,约等于前向的 2 倍 | 每层:权重梯度 + 输入梯度 = 2 × 前向 |
| 6× 规则 | 训练一次的总计算量 | 6×token 数×参数量6 \times \text{token 数} \times \text{参数量}6×token 数×参数量 |
一句话记忆
- 前向:2×token×参数,因为每个参数参与一次乘加。
- 反向:约 4×token×参数,因为每层要算权重梯度和输入梯度两个矩阵乘法。
- 训练总共:6×token×参数。
- MFU:实际速度 ÷ 理论峰值,0.5 以上算优秀,接近 1.0 几乎不可能。