最近为了捡起来一些基本忘干净的基础知识,尝试速刷Stanford CS336课程,因此打算开个新坑。

这篇文章是 CS336 课程笔记系列的第一篇。

主题:Resource Accounting(资源核算)
核心问题:给定计算和显存资源,如何估算模型能否训练、需要多长时间,以及性能瓶颈在哪里。

1. 用张量理解训练系统

深度学习训练中的数据、参数、激活值、梯度和优化器状态都以张量保存。理解一个张量时,至少要明确形状、数据类型、所在设备以及是否需要梯度。张量秩表示维度数量,例如 Transformer 中常见的注意力张量可写为 $[B,S,H,D]$,分别表示批量大小、序列长度、注意力头数和每个头的维度。

张量的原始存储量由元素数量和每个元素的字节数决定:

$$\text{memory}=\operatorname{numel}(x)\times\operatorname{element\_size}(x).$$

FP32 每个元素占 $4$ 字节,FP16 和 BF16 占 $2$ 字节。FP16 的尾数精度较高,但指数范围较小;BF16 保留了 FP32 的 $8$ 位指数,动态范围更大,通常更适合大模型训练。混合精度训练常让参数、激活和梯度使用 BF16,并让优化器状态及部分敏感计算保留 FP32。

FP16、BF16 与 FP8 位分配示意

图 1 FP16、BF16 与两种 FP8 格式的位分配。来源:NVIDIA Transformer Engine

若模型有 $P$ 个参数,参数和梯度使用 BF16,Adam 的一阶矩与二阶矩使用 FP32,则仅这三部分约占

$$2P+2P+4P+4P=12P\ \text{bytes}.$$

这还没有计入激活值、临时缓冲区、通信缓冲区及某些实现中的 FP32 主权重。激活显存还会随批量大小、序列长度、隐藏维度和层数增长。

2. 用 einops 表达张量操作

einops​ 用有含义的维度名描述张量变化,可以减少 transpose​、permute​ 和 reshape​ 中由下标造成的错误。表达式左侧描述输入维度,右侧描述输出维度;没有出现在输出中的维度会在 einsum​ 或 reduce 中被聚合,括号用于拆分或合并维度。

from einops import einsum, rearrange, reduce

# x、y 的形状分别为 [batch, seq1, hidden] 和 [batch, seq2, hidden]
# 对 hidden 维做内积,得到 [batch, seq1, seq2]
scores = einsum(
    x,
    y,
    'batch seq1 hidden, batch seq2 hidden -> batch seq1 seq2',
)

# 将 d_model 拆成多个注意力头
q = rearrange(q, 'batch seq (heads dim) -> batch heads seq dim', heads=12)

# 对 hidden 维求和
summed = reduce(x, 'batch seq hidden -> batch seq', 'sum')

阅读表达式时,先写出每个轴的语义和长度,再检查哪些轴被保留、合并、拆分或聚合。详细操作示例可继续参考当天日记中的 einops 整理。

3. FLOPs、FLOP/s 与 MFU

FLOP 表示一次浮点加法或乘法,FLOPs 表示完成任务需要的浮点运算量,FLOP/s 表示硬件每秒的浮点运算吞吐。硬件峰值会随设备型号、数据类型、稀疏模式和所用计算单元变化,因此实际 BF16 计算不能与 FP32 峰值直接比较。

设 $A\in\mathbb{R}^{m\times n}$、$B\in\mathbb{R}^{n\times p}$,矩阵乘法 $C=AB$ 包含 $mp$ 个输出元素,每个元素约需 $n$ 次乘法和 $n$ 次加法,所以

$$\mathrm{FLOPs}_{\mathrm{matmul}}\approx 2mnp.$$

模型 FLOPs 利用率(Model FLOPs Utilization,MFU)衡量有效模型吞吐占同精度硬件理论峰值的比例:

$$\mathrm{MFU}=\frac{\text{模型有效 FLOP/s}}{\text{相同精度下的硬件理论峰值 FLOP/s}}.$$

MFU 低于 $100\%$ 的原因包括显存读写、设备间通信、内核启动、非矩阵运算、填充 token 和等待时间。它与监控工具中的 GPU utilization 不同:GPU 可以持续忙碌,但大量时间用于访存或通信时,MFU 仍可能较低。

4. 训练计算量为什么约为 $6NP$

设稠密模型有 $P$ 个参数,训练共处理 $N$ 个数据位置;对语言模型而言,$N$ 通常是 token 数。线性层前向传播需要一次矩阵乘法,计算量约为 $2NP$。反向传播既要计算参数梯度,也要计算输入梯度,相当于两次同规模矩阵乘法,因此约为 $4NP$。总计算量近似为

$$\mathrm{FLOPs}_{\mathrm{training}}\approx 2NP+4NP=6NP.$$

于是每个 token 的训练计算量可粗略估为 $6P$ FLOPs。若吞吐为 $r$ tokens/s,则有效模型吞吐约为 $6Pr$ FLOP/s,可进一步估算 MFU 或训练时间。

该结论主要适用于矩阵乘法占主导的稠密模型,也是短上下文 Transformer 的常用近似。长上下文注意力、激活函数、归一化、优化器更新、通信与激活重计算会带来额外开销;MoE 应使用每个 token 实际激活的参数量估算。

5. 算术强度与 Roofline

一次 GPU 计算通常包含从显存读取输入、执行运算和写回输出三个阶段。算术强度(Arithmetic Intensity)衡量每搬运一个字节能够完成多少浮点运算:

$$I_{\mathrm{op}}=\frac{\mathrm{FLOPs}}{\mathrm{Bytes\ transferred}}.$$

硬件自身也有一个计算与带宽的比值:

$$I_{\mathrm{device}}=\frac{\text{峰值 FLOP/s}}{\text{显存带宽 Byte/s}}.$$

Roofline 模型给出的可达吞吐上界为

$$\text{performance}\leq\min\bigl(\text{peak FLOP/s},\ \text{bandwidth}\times I_{\mathrm{op}}\bigr).$$

当 $I_{\mathrm{op}}<I_{\mathrm{device}}$ 时,计算通常受显存带宽限制,属于 memory-bound;当 $I_{\mathrm{op}}>I_{\mathrm{device}}$ 时,计算更可能受算力限制,属于 compute-bound。逐元素操作需要读写大量数据,却只对每个元素做少量运算,通常是 memory-bound。足够大的矩阵乘法能够重复利用读入的数据,算术强度随矩阵规模提高,通常是 compute-bound。小矩阵、狭长矩阵和矩阵向量乘法仍可能受带宽限制,这也是低批量推理常见的瓶颈。

Roofline 模型中的带宽受限与计算受限区域

图 2 Roofline 模型。横轴为算术强度,纵轴为实际吞吐;斜线区域受带宽限制,水平区域受峰值算力限制。来源:JAX Scaling Book

6. 训练循环与优化器状态

一个标准训练步骤包括取数据、前向计算损失、反向求梯度、更新参数和清空梯度。loss.backward()​ 会把梯度累加到参数的 .grad 中,因此每次参数更新后都要显式清空。

# 清空上一步梯度;set_to_none=True 通常能减少不必要的写入
optimizer.zero_grad(set_to_none=True)

# 前向传播并计算损失
prediction = model(x)
loss = loss_fn(prediction, target)

# 反向传播并更新参数
loss.backward()
optimizer.step()

Adam 类优化器除参数和梯度外,还要保存一阶矩与二阶矩,若两者均为 FP32,则额外占用约 $8P$ 字节。实际训练能否装入显存,必须同时核算参数、梯度、优化器状态、激活值和临时缓冲区。

7. 梯度累积

大批量通常能提高训练稳定性,但激活显存随微批量大小增加。梯度累积把一个大批量拆成多个 micro-batch,依次前向和反向,但暂不清空梯度;累计若干次后再执行一次参数更新。数据并行时,有效全局批量为

$$B_{\mathrm{global}}=B_{\mathrm{micro}}\times K\times N_{\mathrm{device}},$$

其中 $K$ 是累积步数。若希望梯度等价于整个大批量的平均梯度,通常应将每个 micro-batch 的损失除以 $K$,或在更新前对累计梯度取平均。

梯度累积让显存中只保留一个 micro-batch 的激活,能够增大有效批量,但会增加串行执行次数。它不会减少参数、梯度缓冲区和优化器状态本身的占用。

8. 激活检查点

反向传播需要前向过程中的中间激活值。激活检查点(Activation Checkpointing,也称 Gradient Checkpointing 或 Rematerialization)只保存部分层的激活;反向传播需要缺失激活时,从最近的检查点重新执行一段前向计算。

普通训练与激活检查点的保存和重计算对比

图 3 普通训练与激活检查点的对比。检查点区域只保留入口张量,反向传播时重新计算内部激活。来源:PyTorch

这种方法用额外计算换取显存。全部保存时,激活显存随层数 $L$ 近似线性增长且无需重算;若每隔约 $\sqrt{L}$ 层保存一次,可将检查点相关存储量降到约 $O(\sqrt{L})$,同时保持约 $O(L)$ 量级的重计算。检查点间隔越大,显存越省,但反向传播需要重算的内容越多。

9. 本讲总结

模型训练可以统一看成张量上的前向、反向与参数更新。资源核算时应同时关注张量形状与数据类型、参数和激活显存、总 FLOPs、硬件 FLOP/s、显存带宽以及 MFU。矩阵乘法通常提供较高算术强度,逐元素操作和矩阵向量乘法更容易受显存带宽限制。梯度累积通过缩小 micro-batch 获得更大的有效批量,激活检查点通过重算减少中间激活,两者都体现了显存、计算量和训练速度之间的权衡。

参考:Stanford CS336 Spring 2026

最后修改:2026 年 08 月 03 日
如果觉得我的文章对你有用,请随意赞赏