第 2 讲:PyTorch 与资源核算(FLOPs、显存、算术强度)
第 2 讲:PyTorch 与资源核算(FLOPs、显存、算术强度)
日期:4 月 1 日(周三,Spring 2026) | 讲师:Percy Liang | 材料:lecture_02.py
概览
这是一讲”系统思维”课:在你谈”如何在固定资源下训练最好的模型”之前,必须先学会核算(accounting)一次计算消耗的显存与算力。本讲覆盖 PyTorch 张量基础(数据类型、显存占用)、用 einops 写出可读的张量运算、FLOPs 计数,以及最关键的算术强度 / roofline 分析(你是 compute-bound 还是 memory-bound?),最后把核算应用到训练循环、梯度累积与激活重计算上。
核心概念与定义
- 张量(Tensor):深度学习里存储一切的基本单元——数据、参数、梯度、优化器状态、激活值。它有秩(rank,即维度数);Transformer 里常见秩为 4 的张量,形状为 (B=32, S=16, H=16, D=64)(batch、序列位置、head、每头维度)。
- 类比:张量就像一张固定轴数的电子表格;秩 4 的张量像一座按”批次、位置、注意力头、特征”四个方向组织的仓库,每层货架上都是一张表格。
- 精度类型(dtype)与权衡:fp32(4 字节,默认)、fp16(2 字节,但动态范围小——
torch.tensor([1e-8], dtype=torch.float16)会下溢为 0!)、bf16(2 字节,动态范围与 fp32 相同但分辨率更差——1e-8 不会下溢)、fp8(E4M3/E5M2,H100 支持)、fp4(NVFP4,仅 4 比特,Blackwell)。- 类比:fp16 像一把只标厘米的尺子——便宜,但量不出一根头发的粗细;bf16 像一把量程极大、但刻度略粗的尺子;fp32 是精密千分尺。用 fp16/bf16 训练时,小数值塌缩成 0 会造成”训练不稳定”。
- 混合精度训练(mixed precision):参数/激活/梯度用 bf16,优化器状态用 fp32(因为要在很多步上累积,需要精度)。PyTorch 的
torch.amp.autocast会自动处理。 - FLOPs 与 FLOP/s(发音相同、极易混淆):FLOPs = 浮点运算次数(衡量做了多少”工作量”,例如 GPT-3 约 3.14e23);FLOP/s = 每秒浮点运算次数(衡量硬件速度)。
- MFU(Model FLOPs Utilization,模型浮点利用率):实际 FLOP/s ÷ 峰值(承诺)FLOP/s。≥0.5 就算相当不错;它永远达不到 1,因为受限于显存带宽、kernel 效率等。
- 算术强度(arithmetic intensity):某次计算的 FLOPs ÷ 搬运的字节数。加速器强度(accelerator intensity):硬件峰值 FLOP/s ÷ 显存带宽(H100 约 295 FLOP/byte)。若负载的算术强度 < 加速器强度 → memory-bound(访存受限);若大于 → compute-bound(计算受限)。
- 类比:工厂(计算单元)由一条传送带(显存带宽)供料。带子送料不够快,工厂就闲置——你是”带子受限”(memory-bound);带子送料过剩而工厂本身慢,则是”工厂受限”(compute-bound)。
- 6ND 法则:训练一个 N 参数模型、喂 D 个 token,总计算量约 6·N·D FLOPs(前向 2ND,反向 4ND)。反向传播是前向的 2 倍。
- Roofline 模型:把算术强度(横轴)与达到的 FLOP/s(纵轴)画在一起;拐点就是加速器强度;MFU = min(1, 算术强度 / 加速器强度)。
- 梯度累积(gradient accumulation):为了用”大逻辑 batch”而不承担大 batch 的显存,在多个 micro-batch 上分别算梯度并累加(不清零),最后统一更新一次。
- 激活重计算(activation checkpointing / gradient checkpointing / rematerialization):只保存部分层的激活,其余在反向时重新计算。显存-算力权衡:每层都存是 O(L) 显存、零重算;完全不存是 O(1) 显存但 O(L²) 重算;每隔 √L 层存一次则 O(√L) 显存、O(L) 重算。
代码示例:用 einops 写出可读的张量数学
代码(Python):
from einops import rearrange, einsum, reduce
import torch
x = torch.ones(2, 3, 4) # batch seq hidden
y = torch.ones(2, 3, 4) # batch seq hidden
# 传统写法(很容易把 -2、-1 搞混):
z = x @ y.transpose(-2, -1) # batch seq seq
# einops 写法:
z = einsum(x, y, "batch seq1 hidden, batch seq2 hidden -> batch seq1 seq2")
# 用 '...' 表示对任意个前导维度做广播:
z = einsum(x, y, "... seq1 hidden, ... seq2 hidden -> ... seq1 seq2")
# reduce:对最后一个维度求和
y_sum = reduce(x, "... hidden -> ...", "sum")
# rearrange:把被压平的维度拆成 (heads, hidden1)
w = torch.ones(4, 4)
x = rearrange(x, "... (heads hidden1) -> ... heads hidden1", heads=2)
x = einsum(x, w, "... hidden1, hidden1 hidden2 -> ... hidden2")
x = rearrange(x, "... heads hidden2 -> ... (heads hidden2)")
代码做了什么:
einsum是”带维度命名”的广义矩阵乘法:同时在两个操作数中出现、但不出现在输出里的维度会被求和消去(contracted)。第一个例子按 batch 元素计算x @ yᵀ(正是 attention 打分矩阵的模式)。reduce对命名维度做聚合(sum/mean/max/min)。rearrange只改变形状、不改变数据,包括拆分/合并带括号的维度((heads hidden1))。
实现深挖:
- 为什么要给维度起名:
x @ y.transpose(-2, -1)语义不透明——到底哪个维度被消去?einops 把”收缩模式”明确写进字符串,并在运行时检查形状错误,而讲义认为手写 transpose 极易出错。生产代码中,普通torch.matmul/torch.bmm通常比 einops 的 einsum 更快;所以实践建议是”原型期用 einops 保证清晰”,而作业 1 的 handout 正是用这种 einops 风格来表述 Transformer 前向计算的。 - 为什么要有
...(省略号):它让同一个表达式既能处理带 batch 的情况、也能处理不带 batch 的情况,对任意个前导维度广播——写”维度无关”的层时非常方便。
与作业的联系:作业 1 的 Transformer 前向计算正是用这套记号表述的(handout 里 attention 与 RMSNorm 都给出了 einops 写法)。作业 2 的 FlashAttention kernel 也需要以命名维度(B、H、S、D)的思维来做 tiling。
代码示例:线性层与训练步的 FLOPs 计数
代码(Python):
B, D, K = 1024, 256, 64 # batch、输入维度、输出维度
x = torch.ones(B, D)
w = torch.randn(D, K)
y = x @ w
# 每个 (i, j, k) 三元组对应一次乘法 + 一次加法:
actual_num_flops = 2 * B * D * K
# 两层 MLP:
# 前向: h1 = x @ w1 -> 2*B*D*D FLOPs
# h2 = h1 @ w2 -> 2*B*D*D FLOPs
# 反向: h1.grad = h2.grad @ w2^T -> 2*B*D*D
# w2.grad = h2.grad^T @ h1 -> 2*B*D*D
# 每层合计:2(前向)+ 4(反向)= 6 * B * D * D
代码做了什么:
- 统计矩阵乘法(
2 * M * N * K)的 FLOPs,并说明一层的反向传播恰好是前向的 2 倍(两次矩阵乘法:一次算输入梯度、一次算权重梯度),从而导出著名的 6ND 法则。
实现深挖:
- 为什么是 2·B·D·K:每个输出元素是一个长度为 K 的点积,约含 K 次乘法 + (K−1) 次加法 ≈ 2K FLOPs,乘以输出元素个数(讲义按”每个 (i,j,k) 三元组一次乘加”计为 2·B·D·K)。真正重要的是比例:反向 = 2 × 前向。
- 为什么它对 Transformer 只是近似:6ND 对 MLP 精确,对短上下文的 Transformer 是很好的近似(长上下文下 attention 会额外贡献不可忽略的一项)。
- 为什么要实测而非只看规格:峰值 FLOP/s 依赖精度与稀疏性(H100:含稀疏 1979 TFLOP/s,不含则减半);实际吞吐要用”FLOPs ÷ 时间”测出来,再算出 MFU。
与作业的联系:作业 1 的资源核算部分要求你针对给定配置,算出 Transformer 每个组件的 FLOPs(embedding、attention 的 QKᵀ、softmax、attention·V、MLP 的 up/gate/down 投影、LM head)——正是这种 2·M·N·K 计数。显存核算(bf16 下 2 字节参数 + 2 字节梯度 + 8 字节 AdamW 优化器状态 = 每参数 12 字节)同样属于作业 1,并且是作业 2 分布式显存规划的基础。
代码示例:从零实现 AdaGrad 与训练循环
代码(Python):
class AdaGrad(torch.optim.Optimizer):
def __init__(self, params, lr=0.01):
super().__init__(params, dict(lr=lr))
def step(self):
for group in self.param_groups:
lr = group["lr"]
for p in group["params"]:
state = self.state[p]
grad = p.grad.data
g2 = state.get("g2", torch.zeros_like(grad)) # 梯度平方的累积和
g2 += torch.square(grad)
state["g2"] = g2
p.data -= lr * grad / torch.sqrt(g2 + 1e-5)
# 标准训练循环:
for t in range(num_train_steps):
x, y = get_batch()
pred_y = model(x).mean()
loss = F.mse_loss(pred_y, y)
loss.backward()
optimizer.step()
optimizer.zero_grad(set_to_none=True)
代码做了什么:
AdaGrad为每个参数维护”梯度平方的累积和”g2,并用sqrt(g2 + eps)去除梯度——过去的梯度越大,当前有效步长越小(逐坐标自适应学习率)。- 训练循环完成标准动作:采样 batch → 前向 → 算损失 → 反向 → 优化器更新 → 清零梯度。
实现深挖:
- 为什么要自己写优化器:讲义用 AdaGrad 作为铺垫,串起优化器家族谱系:momentum = SGD + 梯度的指数平均;AdaGrad = SGD + 梯度平方的(累积)平均;RMSProp = 把 AdaGrad 的梯度平方改成指数平均;Adam = RMSProp + momentum。作业 1 要求从零实现 AdamW——你写的类基本上就是上面这个加上一阶/二阶矩估计与权重衰减。
- 为什么优化器状态用 fp32:讲义原话是”习惯上使用 fp32 以保证稳定性(要在很多步上累积幂次平均)”——这正是混合精度的由来:参数/梯度 bf16,优化器状态 fp32(Adam 两个矩共 8 字节/参数;AdaGrad 一个 4 字节/参数)。
- 为什么
zero_grad(set_to_none=True):置为 None 比清零更快释放显存。
与作业的联系:作业 1:实现交叉熵损失、AdamW,以及带 checkpoint 保存/加载的训练循环——正是上面这个模式。显存表(参数 2B + 梯度 2B + 优化器状态 4–8B + 激活)就是你在资源核算里要报告的;作业 2 在这个循环之上构建分布式版本(DDP/FSDP)。
代码示例:梯度累积与激活重计算
代码(Python):
# 梯度累积:用 256 的 micro-batch 模拟 4096 的大 batch
for micro_step in range(accumulation_steps):
x, y = get_micro_batch()
loss = loss_fn(model(x), y) / accumulation_steps # 损失要缩放
loss.backward() # 累积进 .grad(不清零)
optimizer.step() # 整个逻辑 batch 只更新一次
optimizer.zero_grad(set_to_none=True)
# 激活重计算:
for layer in self.layers:
x = torch.utils.checkpoint.checkpoint(layer, x) # 反向时重算
代码做了什么: 第一段把一次优化器更新摊到多个 micro-batch 上,用”一个 micro-batch 的激活显存”换到与大批量相同的梯度统计量。第二段用 torch.utils.checkpoint.checkpoint 包裹每层:前向丢弃中间激活,反向时重新计算。
实现深挖:
- 为什么要梯度累积:激活显存随 batch 大小线性增长;逻辑 batch 为 64×1024 维 × 16 层时激活需要 2·64·1024·16 字节,容易爆显存。切成 256 的 micro-batch 可把激活显存降低 4 倍。
- 为什么重计算是”用算力换显存”:全部保存是 O(L) 显存、零重算;完全不存是 O(1) 显存但 O(L²) 算力(每层都从头重算);每隔 √L 层存一次则在两者之间取得平衡。
- 为什么注意损失缩放:损失要除以累积步数,否则等价于把学习率乘以累积步数。
与作业的联系:作业 2 的激活重计算任务(为 TransformerBlock 实现重计算,并用显存 hook 验证)正是这个例子的直系后代——handout 里同样给出了 pack_hook/unpack_hook 插桩来测量被保存张量的显存。
关键要点
- 一切都是张量:参数、梯度、激活、优化器状态——每种的显存 = 元素个数 × 单元素字节数(bf16 为 2B,fp32 为 4B)。
- 6ND 法则:训练成本约等于”每参数每 token 6 个 FLOPs”(前向 2 + 反向 4);反向是前向的 2 倍。
- 用 roofline 思维:矩阵乘法是 compute-bound(算术强度约 n/3),逐元素运算(ReLU/GELU)与点积是 memory-bound——所以”孤立地看,ReLU 并不比 GELU 快”,而推理(矩阵-向量乘)必然是 memory-bound。
- 显存技巧很重要:混合精度(bf16 + fp32 优化器状态)、梯度累积、激活重计算,让你能塞下更大的 batch 或模型。
- einops 让张量数学可读可调;MFU ≥ 0.5 已算不错,而且计时一定要配合
torch.cuda.synchronize()。
常见陷阱
- fp16 下溢:1e-8 这类值会塌成 0 导致训练不稳定;优先用 bf16,或用 fp32 做累加。
- benchmark 忘记
torch.cuda.synchronize():CUDA 是异步的,不同步测到的只是 kernel 启动开销,不是 kernel 时间。应使用 CUDA events。 - 混淆 FLOPs(工作量)与 FLOP/s(速度):两者读音相同但是不同量;另外峰值 FLOP/s 取决于精度与稀疏性。
- 梯度累积时忘记缩放损失:等价于改变了有效学习率。
- 核算显存时忽略激活:”8 张 H100 能训多大模型”的纸面推算只是上界——激活取决于 batch 与序列长度,可能成为主导项。
- 盲目使用
-2, -1转置:attention/MLP 代码里的维度错乱是经典 bug 来源;要么命名维度(einops),要么在注释里写清形状。
复习题
- 问: 某 GPU 峰值 1000 TFLOP/s、带宽 3.35 TB/s。某负载搬运 1 GB、做 1 TFLOP。它是 memory-bound 还是 compute-bound?
- 答: 加速器强度 ≈ 1000e12 / 3.35e12 ≈ 298 FLOP/byte;负载强度 = 1e12 / 1e9 = 1000 FLOP/byte > 298 → compute-bound。
- 问: 为什么线性层的反向传播是前向的 2 倍?
- 答: 需要算两个梯度:传给前面层的输入梯度,以及本层权重梯度——各自都是与前后向同量级的矩阵乘法,因此前向 1 次 + 反向 2 次 ≈ 每层 6·B·D²,即前向 2ND、反向 4ND。
- 问: 为什么推理的算术强度低,而训练不低?
- 答: 训练处理大的批量矩阵乘法(B≫1,compute-bound)。推理一次只解码一个 token(B=1):每步都要读全部参数(矩阵-向量乘),强度约等于 1,远低于加速器强度——因此是 memory-bound。
