工作台课程

CS336 · 从零构建语言模型

Lecture 2 · PyTorch 与资源核算

Lecture 2 · PyTorch 与资源核算

CS336: Language Modeling from Scratch · Stanford · Spring 2025
📅 Apr 3 · 💻 可执行讲义 · 讲者 Percy Liang


承上启下

上一讲(Lecture 1 · 概览与分词

  • 语言模型的流水线从 token 序列开始——BPE 把文本压缩成整数序列,压缩率直接决定训练与推理的成本。
  • 课程主线是 from scratch 与效率:$\text{accuracy} = \text{efficiency} \times \text{resources}$,在资源给定的前提下,一切设计决策都围绕效率展开。

本讲(PyTorch 与资源核算)

  • 从 tensor、operation、gradient、model、optimizer 到 training loop,搭起训练语言模型的最小积木。
  • 重点不在 Transformer 细节(那是 Lecture 3 · 架构与超参数 的事),而在 memory accounting(内存核算)compute accounting(算力核算)
  • 训练前先做 napkin math:模型放不放得下?训练要多久?瓶颈在哪里?
  • 如果需要复习张量与自动微分的入门概念,可参看 CS221 的 Lecture 1 · 张量、梯度与监督学习(未找到对应页面)——那边讲「是什么」,这边讲「花多少钱」。

提示

PyTorch 是机制,资源核算是心态。你要知道每个 tensor 占多少内存、每个矩阵乘法要多少 FLOPs、反向传播为什么约是前向的 2 倍,最后才能判断一个训练方案是不是现实。


1. 两个开场估算

1.1 70B 模型训练要多久?

假设:

  • 参数量 $N = 70\times10^9$;
  • 训练 token 数 $D = 15\times10^{12}$;
  • 设备:1024 张 H100;
  • 训练 FLOPs 近似(推导见第 4 节):$C \approx 6ND$。

$$C \approx 6 \times 70\times10^9 \times 15\times10^{12} = 6.3\times10^{24}\ \text{FLOPs}$$

total_flops = 6 * 70e9 * 15e12          # 6.3e24 FLOPs
h100_flop_per_sec = 1979e12 / 2         # 989.5 TFLOP/s(BF16 稠密峰值)
mfu = 0.5
flops_per_day = h100_flop_per_sec * mfu * 1024 * 60 * 60 * 24
days = total_flops / flops_per_day      # ≈ 144 天

两个容易踩的坑藏在硬件数字里:

  • H100 规格表上的 1979 TFLOP/s 是带 2:1 结构化稀疏的营销数字,稠密矩阵乘法要除以 2,得 989.5 TFLOP/s;
  • 峰值算力永远达不到——通信、内存带宽、kernel launch、数据加载都会拖后腿,实际能达到的比例就是 MFU(Model FLOPs Utilization),能到 50% 已经算好。

于是:理想满血(MFU = 1)约 $6.3\times10^{24} / (989.5\times10^{12} \times 1024) \approx 6.2\times10^6$ 秒 ≈ 72 天;按现实的 MFU = 0.5 算则是 约 144 天

注意

大模型训练的残酷现实:纸面算力先砍一半(稀疏营销数字),再乘 MFU 打对折——端到端训练时间轻松是「理想估算」的 2–4 倍。做估算时一定要写明用的是稠密峰值还是稀疏峰值、假设了多少 MFU。

1.2 8 张 H100 能朴素训练多大模型?

用 AdamW 做最朴素的 fp32 训练时,每个参数占的不只是自己那 4 字节:

内容 dtype 每参数字节
模型参数 FP32 4
梯度 FP32 4
Adam 一阶矩 m FP32 4
Adam 二阶矩 v FP32 4
合计 16 bytes / param
  • 朴素数据并行(每张卡完整复制以上全部状态):单卡 80GB 只能放 $80\times10^9 / 16 \approx 5\times10^9$,即约 5B 参数——这还没算激活值、临时 buffer 和 CUDA workspace,实际连 5B 都紧张。
  • 理想上限(如果能把这些状态完美切分到 8 张卡上——这正是 ZeRO / FSDP 做的事):$8 \times 80\times10^9 / 16 \approx 40\times10^9$,即约 40B 参数的上限(同样未计激活)。

说明

换成混合精度会怎样?BF16 参数 (2) + BF16 梯度 (2) + FP32 master weights (4) + m (4) + v (4) = 仍然 16 字节/参数。混合精度省的是计算时间和激活内存,并不省优化器状态——这个反直觉的结论解释了为什么显存优化的主战场是切分优化器状态(Lecture 7 · 并行 I:基础Lecture 8 · 并行 II:分布式训练实战 的 ZeRO/FSDP)。

提示

这两个开场估算演示了全课最重要的工作习惯:写代码之前先算数量级。两次乘除法就能判断「这个配置根本不可行」,比跑起来 OOM 再调试便宜一万倍。


2. Tensor:所有东西的容器

PyTorch tensor 是训练中一切状态的载体:参数(parameters)、梯度(gradients)、优化器状态(optimizer states)、数据 batch、激活值(activations)。内存核算的第一步就是弄清每个 tensor 占多少字节:

$$\text{memory} = \text{numel} \times \text{bytes per element}$$

其中 numel 是元素个数(各维度之积),每元素字节数由 dtype 决定。

2.1 dtype 决定每个数的成本

浮点数由三部分组成:符号位、指数位(决定动态范围,即能表示多大/多小的数)、尾数位(决定精度,即有效数字位数):

dtype 位分配(符号/指数/尾数) 字节 动态范围 训练中的角色
float32 1/8/23 4 ~10³⁸ 稳定但贵,用于优化器状态与累加
float16 1/5/10 2 最大 65504 快,但小于 6×10⁻⁵ 就下溢
bfloat16 1/8/7 2 与 FP32 相同 现代训练默认选择
float8 (E4M3) 1/4/3 1 最大 448 H100 支持,精度换吞吐
float8 (E5M2) 1/5/2 1 最大 57344 范围换精度的另一种取舍
torch.tensor([1e-8], dtype=torch.float16)   # → 0.(下溢!)
torch.tensor([1e-8], dtype=torch.bfloat16)  # → ≈1e-8(指数位够用)

例如 GPT-3 FFN 中一块大矩阵,形状约 [12288, 49152](即 $d_\text{model} \times 4d_\text{model}$):

numel = 12288 * 49152
fp32_bytes = numel * 4   # 约 2.4 GB
bf16_bytes = numel * 2   # 约 1.2 GB

提示

为什么 bf16 成了训练默认?fp16 把省下来的位数投给尾数(精度),bf16 投给指数(范围)。训练中真正致命的是梯度下溢变零(信号彻底消失)而不是尾数末位的噪声——深度学习本身就对噪声鲁棒。bf16 与 fp32 指数位相同、动态范围一致,从根上消除了下溢问题,所以能省掉 fp16 所需的 loss scaling 全套补丁。

更一般的原则:低精度不是放弃精度,而是把高精度留给敏感位置(累加、归一化、优化器状态),把不敏感的大头(矩阵乘法)交给低精度——这就是第 8.3 节混合精度的全部逻辑。

2.2 Tensor 不是数组本身,而是「指针 + 元数据」

PyTorch tensor 由两层构成:真正连续存放数字的 storage(一维数组),和描述「如何解读这段内存」的元数据——shape、stride(步幅)、dtype、device、offset。

x = torch.arange(12).view(3, 4)
x.shape     # (3, 4)
x.stride()  # (4, 1)

stride=(4, 1) 表示:行方向前进一步要在 storage 里跳 4 个元素,列方向跳 1 个。访问元素 x[r, c] 的物理位置由下式给出:

$$\text{index} = \text{offset} + r \cdot s_0 + c \cdot s_1$$

其中 $s_0, s_1$ 是各维度的 stride。这个「地址计算公式」是理解一切 view 操作的钥匙。

2.3 View 很便宜,copy 很贵

很多操作只改元数据(shape/stride/offset),完全不碰数据本身,因此是 $O(1)$ 的「免费」操作:

x = torch.arange(12).view(3, 4)
y = x[:, 1]            # view:改 offset 和 stride
z = x.transpose(0, 1)  # view:交换 stride,变成 non-contiguous

view 与原 tensor 共享 storage——改一个会影响另一个。而转置后的 tensor 是 non-contiguous(stride 不再是「行优先递减」布局),再做 .view() 会失败:

z = x.transpose(0, 1)
z.view(12)               # 报错:view 只能作用于 contiguous 布局
z.contiguous().view(12)  # 先物理复制成连续存储,再 view

注意

contiguous() 不是免费的:它按新布局完整复制一遍数据,既占内存又消耗显存带宽。在热点路径里反复 transpose + contiguous 是常见的隐形性能杀手——profiler 里看到大量 copy_ kernel 时先查这里。

2.4 einops:给维度起名字

形如 x.transpose(1, 2).reshape(B, T, -1) 的代码充满魔法数字,维度一多就容易错。einops 风格的操作用命名维度表达意图:

from einops import rearrange, reduce, einsum

# 注意力打分:对 hidden 维做内积,保留两个序列维
scores = einsum(q, k, "batch seq1 hidden, batch seq2 hidden -> batch seq1 seq2")

# 沿 hidden 维求和/平均(... 表示任意批维原样保留)
s = reduce(x, "... hidden -> ...", "sum")

# 多头拆分:把 (heads*d) 拆成 (heads, d) 两个维度
x = rearrange(x, "... (heads d) -> ... heads d", heads=8)

维度名写进字符串后,「哪个维和哪个维收缩」一目了然,转置错误在写下式子的当刻就会暴露。配合 jaxtyping 之类的形状注解(Float[Tensor, "batch seq hidden"]),tensor 代码可以做到自我文档化——作业中强烈推荐这样写。


3. Tensor 操作与 FLOPs

3.1 Elementwise 操作:便宜的计算,昂贵的搬运

逐元素操作对每个位置独立施加运算:

y = torch.sin(x) + x * x

FLOPs 是 $O(\text{numel})$,看起来便宜,但它往往是 memory-bound(受内存带宽限制)的:每个元素只做一两次运算,却要把整个 tensor 从 HBM 读进来、再写回去——算术强度(每字节内存流量对应的运算量)太低,GPU 的算力根本用不上。大量细碎的 elementwise op 串联会造成频繁的 HBM 往返,这正是 kernel 融合(如 Flash Attention)要解决的问题——细节见 Lecture 5 · GPULecture 6 · Kernel 与 Triton

3.2 Matrix multiplication 是深度学习的主菜

矩阵乘法 $Y = XW$,其中 $X \in \mathbb{R}^{B \times D}$,$W \in \mathbb{R}^{D \times K}$:每个输出元素 $Y_{bk} = \sum_{d} X_{bd} W_{dk}$ 需要 $D$ 次乘法和 $D$ 次加法,输出共 $BK$ 个元素,所以

$$\text{FLOPs}(XW) = 2BDK$$

系数 2 来自「乘 + 加」各算一次浮点运算。对大矩阵而言 matmul 的 FLOPs 占绝对大头;但占大头的 FLOPs 不等于占大头的时间——只有 shape 足够大、能喂饱 tensor core 时,matmul 才是 compute-bound 的。

3.3 Transformer 训练 FLOPs 的一阶近似:6ND

对参数量 $N$、训练 token 数 $D$ 的 Transformer:

$$C_{\text{forward}} \approx 2ND, \qquad C_{\text{backward}} \approx 4ND, \qquad C_{\text{total}} \approx 6ND$$

前向系数 2 的来源:每个 token 流经每个矩阵参数时,恰好贡献一次乘法和一次加法(矩阵乘法里每个权重 $W_{dk}$ 对每个输入行都参与一次乘加)。反向为什么是前向的 2 倍,见第 4.2 节。

说明

$6ND$ 忽略了注意力里 $QK^\top$ 与加权求和这两个与参数无关的矩阵乘法(它们的代价随序列长度平方增长)。当上下文长度远小于模型宽度的量级时该项占比很小,一阶估算尽可放心用;长上下文场景(32K+)则需要把注意力项加回来。

提示

先用一阶公式判断数量级,再用 profiler 查具体瓶颈——别一上来就陷入每个小算子的精确账本。$6ND$ 之所以够用,是因为它抓住了「每个参数对每个 token 干了多少活」这个主导项;也正因为它只依赖 $N$ 和 $D$,它成了 Lecture 9 · 缩放定律 I:基础 里一切「算力预算」讨论的通用货币。

3.4 FLOPs、FLOP/s 与 MFU

两个拼写相近但含义不同的量:

名称 含义
FLOPs 浮点运算的总量(工作量)
FLOP/s 每秒能做多少浮点运算(硬件速度),也写作 FLOPS

主流训练卡的稠密峰值(tensor core):

硬件 FP32 BF16
A100 19.5 TFLOP/s 312 TFLOP/s
H100 67 TFLOP/s 989.5 TFLOP/s

BF16 峰值是 FP32 的 15–16 倍——这就是「dtype 直接决定吞吐」的含义。

MFU(Model FLOPs Utilization,模型 FLOPs 利用率)

$$\text{MFU} = \frac{\text{实际达成的 FLOP/s}}{\text{硬件峰值 FLOP/s}}$$

分子按「有用功」计:用 $6ND$ 估算模型本身需要的 FLOPs,除以实测耗时——重计算等额外开销不算有用功。经验值:

  • MFU ≥ 0.5 通常已经很好;
  • 小模型、小 batch、不友好的 shape(无法整除 tensor core tile)都会拉低 MFU;
  • 换用 BF16 会提高分母(峰值),所以同样的代码 MFU 数字反而可能变难看——比较 MFU 时必须固定 dtype 口径。

4. Autograd:反向传播的资源账

4.1 Forward 与 Backward

以单参数的最小例子看机制:

y = 0.5 * (x * w - 5) ** 2
y.backward()

前向时 PyTorch 把每一步运算记入计算图并保存反向所需的中间值;调用 backward() 后按链式法则回传。本例中

$$\frac{\partial y}{\partial w} = (xw - 5) \cdot x$$

——注意梯度公式里出现了前向的中间量 $xw$ 与输入 $x$:这就是为什么前向的中间结果(激活)必须保存到反向用完为止,也是激活内存存在的根源。算出的梯度被累积(而非覆盖)到 w.grad 中。

4.2 为什么 backward 约是 forward 的 2 倍?

对一层线性变换 $H = XW$($X$ 是输入,$W$ 是参数,$H$ 是输出),设损失对输出的梯度为 $\nabla_H$,反向需要算两个矩阵乘法:

$$\nabla_W = X^\top \nabla_H \qquad (\text{参数的梯度,用于更新 } W)$$
$$\nabla_X = \nabla_H W^\top \qquad (\text{输入的梯度,继续传给上一层})$$

每个都与前向那一次 $XW$ 同规模($2BDK$ FLOPs),所以反向 ≈ $4BDK$,是前向的 2 倍。汇总成资源账:

阶段 FLOPs 由来
Forward $2ND$ 1 个 matmul:$XW$
Backward $4ND$ 2 个 matmul:$\nabla_W$ 与 $\nabla_X$
Total $6ND$ 上两行之和

提示

反向传播要回答两个独立的问题:「参数该怎么改」($\nabla_W$)和「责任怎么继续往前追」($\nabla_X$)。每个问题都需要一次与前向同规模的矩阵乘法,所以 2 倍不是巧合而是结构性的。首层例外——输入不需要梯度,可以省掉 $\nabla_X$,但在几十层的网络里这点节省可以忽略。

4.3 梯度会占内存

每个可训练参数在 backward() 后都会挂一个同形状的 .grad

loss.backward()
for p in model.parameters():
    print(p.shape, p.grad.shape)  # 一一对应,同形状同 dtype

所以训练内存不能只算参数,完整清单是:参数 + 梯度 + 优化器状态 + 激活 + 临时 buffer + CUDA allocator 保留的缓存。

例子

一个 $L$ 层、宽度 $D$、batch 大小 $B$ 的 MLP:

  • 参数量 ≈ $L \cdot D^2$(每层一个 $D \times D$ 矩阵);
  • 激活 ≈ $B \cdot L \cdot D$(每层输出都要留到反向);
  • 梯度 = 参数量;优化器状态(AdaGrad 一份、Adam 两份)≈ 参数量的 1–2 倍。

总字节数 ≈ $4 \times (\text{参数} + \text{激活} + \text{梯度} + \text{优化器状态})$。注意激活是唯一随 $B$ 增长的项——这就是「OOM 时先调小 batch size」有效的原因;Transformer 的精确激活账本在作业 1 里算。


5. Model:nn.Parameter 与初始化

5.1 参数是什么?

PyTorch 中,可训练参数是 nn.Parameter——一种被 nn.Module 自动登记的特殊 tensor:

class Linear(nn.Module):
    def __init__(self, in_dim, out_dim):
        super().__init__()
        self.weight = nn.Parameter(torch.randn(in_dim, out_dim))

    def forward(self, x):
        return x @ self.weight

登记的意义在于三个下游通道自动打通:model.parameters()(交给 optimizer 更新)、state_dict()(checkpoint 保存/恢复)、.to(device)(整体搬运设备)。

5.2 初始化为什么重要?

如果直接用标准正态初始化:

W = torch.randn(input_dim, output_dim)
y = x @ W

输出的每个分量是 $y_k = \sum_{i=1}^{d_{\text{in}}} x_i W_{ik}$——$d_{\text{in}}$ 个独立随机项之和,方差随之线性增长:

$$\mathrm{Var}(y_k) = d_{\text{in}} \cdot \mathrm{Var}(x_i)\mathrm{Var}(W_{ik})$$

即输出的标准差按 $\sqrt{d_{\text{in}}}$ 放大。层层叠加后激活会指数爆炸(或反过来消失),训练直接发散。解决办法是把权重标准差缩成 $1/\sqrt{d_{\text{in}}}$,抵消求和带来的方差增长(Xavier 初始化的思想):

W = nn.init.trunc_normal_(
    torch.empty(input_dim, output_dim),
    std=1 / math.sqrt(input_dim),
    a=-3, b=3,   # 按 ±3σ 截断,防止极端离群值
)

提示

初始化的目标是让信号在前向流过每一层时方差保持不变(梯度在反向时同理)。$1/\sqrt{d_{\text{in}}}$ 恰好抵消「$d_{\text{in}}$ 项求和把方差放大 $d_{\text{in}}$ 倍」;截断正态则杜绝了小概率的巨大初值把训练在第 0 步就带偏。Lecture 3 · 架构与超参数 会继续讨论残差分支缩放等更精细的初始化技巧。


6. Data Loading:别把数据一次性读进内存

语言模型数据就是 tokenizer 输出的整数序列(上一讲 Lecture 1 · 概览与分词 的产物)。大模型语料可达 TB 级,不可能全部读进内存。标准做法是内存映射:

data = np.memmap("tokens.npy", dtype=np.uint16, mode="r")

memmap 把文件映射进虚拟地址空间,只有实际访问到的片段才由操作系统按页调入物理内存——程序视角是「一个巨大的数组」,物理内存占用却只有热点部分。

说明

dtype 用 uint16 是刻意的:GPT-2 词表 50,257 < 65,536,两字节就能存下一个 token ID,比 int64 省 4 倍磁盘与内存。喂给 PyTorch 的 embedding 前再转成 int64(PyTorch 索引要求)。

一个简单的 batch 采样器——随机取 batch_size 个起点,切出长度 seq_len 的窗口,标签就是右移一位的同一窗口:

def get_batch(data, batch_size, seq_len, device):
    starts = torch.randint(len(data) - seq_len - 1, (batch_size,))
    x = torch.stack([
        torch.from_numpy(data[i:i+seq_len].astype(np.int64))
        for i in starts
    ])
    y = torch.stack([
        torch.from_numpy(data[i+1:i+seq_len+1].astype(np.int64))
        for i in starts
    ])
    return x.to(device), y.to(device)

6.1 Pinned Memory

CPU 内存默认是可分页的(pageable)——操作系统随时可能把它换出,所以 GPU 的 DMA 引擎无法直接异步读取,普通的 CPU→GPU 拷贝必须同步执行。Pinned(锁页)memory 把页面钉在物理内存里,从而允许真正的异步拷贝:

x = x.pin_memory()
x = x.to("cuda", non_blocking=True)  # 拷贝与计算可以重叠

理想流水线由此形成:GPU 计算当前 batch 的同时,CPU 已在准备下一批数据并异步传输——数据加载被完全藏进计算时间里,GPU 永不挨饿。


7. Optimizer:不仅是更新公式,也是内存大户

7.1 从 SGD 到 AdamW

优化器的演化是一条「逐步给梯度加记忆」的线索(记 $g_t$ 为第 $t$ 步梯度,$\eta$ 为学习率,$\theta$ 为参数):

优化器 核心思想 额外状态
SGD 沿负梯度更新
Momentum 梯度的指数滑动平均,平滑震荡 一阶动量
AdaGrad 按历史梯度平方缩放各坐标学习率 梯度平方累计
RMSProp 梯度平方改用指数滑动平均 二阶动量
Adam / AdamW Momentum + RMSProp 一阶 + 二阶动量
  • SGD:$\theta \leftarrow \theta - \eta\, g_t$。
  • AdaGrad:$s_t = s_{t-1} + g_t^2$,$\theta \leftarrow \theta - \eta\, g_t / (\sqrt{s_t} + \epsilon)$——每个坐标按自己的历史梯度大小自适应缩放;缺点是 $s_t$ 单调增长,学习率不可逆地衰减到零。
  • RMSProp:把累计换成指数滑动平均 $s_t = \beta s_{t-1} + (1-\beta) g_t^2$,学习率不再单调死亡。
  • Adam:同时维护一阶矩 $m_t = \beta_1 m_{t-1} + (1-\beta_1) g_t$ 和二阶矩 $v_t = \beta_2 v_{t-1} + (1-\beta_2) g_t^2$,做偏差修正 $\hat{m}_t = m_t/(1-\beta_1^t)$、$\hat{v}_t = v_t/(1-\beta_2^t)$(抵消零初始化在早期造成的低估),然后

$$\theta \leftarrow \theta - \eta \cdot \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon}$$

符号含义:$\beta_1 \approx 0.9$、$\beta_2 \approx 0.95{\sim}0.999$ 控制两个滑动平均的记忆长度;分母 $\sqrt{\hat{v}_t}$ 让每个参数按自身梯度的典型幅度归一化——所有参数的有效步长尺度一致,这是 Adam 对超参数不敏感、成为 LLM 默认优化器的关键。

  • AdamW 在此之上把权重衰减从梯度中解耦出来,单独执行 $\theta \leftarrow \theta - \eta \lambda \theta$:若像原始 Adam 那样把 L2 项混进梯度,它会被 $\sqrt{\hat{v}_t}$ 除掉,导致梯度大的参数几乎不被正则化。解耦后衰减强度对所有参数一致。

7.2 AdamW 为什么吃内存?

每个参数要陪绑一阶矩 $m$、二阶矩 $v$(通常均为 fp32),混合精度下还要一份 fp32 master weights——优化器状态比参数本身还大(见 1.2 节的 16 字节账本)。这直接催生了两类工程方案:


8. Training Loop 与工程习惯

把前面所有积木拼起来,最小训练循环只有六行:

for step in range(num_steps):
    x, y = get_batch(data, batch_size, seq_len, device)

    logits = model(x)
    loss = F.cross_entropy(
        logits.view(-1, vocab_size),
        y.view(-1),
    )

    optimizer.zero_grad(set_to_none=True)
    loss.backward()
    optimizer.step()

zero_grad(set_to_none=True).grad 置为 None 而不是填零——省一次显存写入,还能让下一次 backward 直接赋值而非累加。

生产级训练在此之上还必须加:seed 管理(可复现)、checkpoint(容错)、学习率调度(warmup + 衰减)、梯度裁剪(防 loss spike)、混合精度(吞吐)、logging 与 validation(可观测性)、OOM 恢复。每一件都不难,但缺一件就会在某个深夜付出代价。

8.1 随机性

随机性藏在五个地方:参数初始化、dropout、batch 采样、CUDA kernel 的非确定性归约、数据顺序。想要可复现,前面几个用种子控制:

random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)

注意

设了所有 seed 仍可能得到不同结果——部分 CUDA kernel(如某些 atomicAdd 归约)本身是非确定性的,浮点加法不满足结合律,线程完成顺序会改变末位数值。完全确定性需要 torch.use_deterministic_algorithms(True),但会牺牲性能;实践上通常接受「统计意义上的可复现」。

8.2 Checkpointing

长时间训练一定会崩(硬件故障、抢占、OOM),checkpoint 不是可选项。保存的不只是模型——优化器状态(Adam 的 m/v)和 RNG 状态同样必须保存,否则恢复后动量清零、数据顺序错位,loss 曲线会出现肉眼可见的断层:

torch.save({
    "model": model.state_dict(),
    "optimizer": optimizer.state_dict(),
    "step": step,
    "rng_state": torch.get_rng_state(),
}, "checkpoint.pt")

恢复:

ckpt = torch.load("checkpoint.pt")
model.load_state_dict(ckpt["model"])
optimizer.load_state_dict(ckpt["optimizer"])
step = ckpt["step"]

8.3 Mixed Precision

低精度快且省内存但不稳定,fp32 稳定但贵——混合精度的答案是按算子的敏感度分配精度

with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
    logits = model(x)
    loss = loss_fn(logits, y)

loss.backward()
optimizer.step()

autocast 的策略:

  • 矩阵乘法等 compute-heavy 且对精度不敏感的算子 → BF16(吃满 tensor core 的 15 倍吞吐);
  • 归约类算子(softmax、norm、loss)在 fp32 中累加——长序列求和最容易累积舍入误差;
  • 优化器状态与 master weights 保持 FP32;
  • 用 FP16 时梯度可能下溢,需要 loss scaling(把 loss 放大 $2^k$ 倍让梯度进入可表示范围,更新前再缩回);BF16 因动态范围与 fp32 相同而免掉这套机制,故成为现代默认;
  • 更激进的 FP8 训练由 NVIDIA Transformer Engine 支持(H100 起),需要逐层的缩放因子管理。

总结

mindmap
  root((PyTorch 与资源核算))
    Memory
      dtype 位分配
      numel x element_size
      params gradients optimizer states
      activations 随 batch 增长
    Tensor
      storage 指针加元数据
      shape stride offset
      view vs copy
      contiguous
      einops 命名维度
    Compute
      FLOPs vs FLOP/s
      matmul 2BDK
      6ND 法则
      MFU
    Autograd
      前向存激活
      backward 两个 matmul
      gradient memory
    Training
      初始化 1/√fan_in
      memmap 与 pinned memory
      SGD 到 AdamW
      checkpoint 含优化器与 RNG
      mixed precision

关键要点

  1. 资源核算是训练前的 sanity check
    - $C \approx 6ND$ 估训练 FLOPs,除以(稠密峰值 × MFU × 卡数)估时间;
    - 每参数字节数(朴素 AdamW 为 16)估训练显存;
    - 先算数量级,再写代码。

  2. Tensor 的 view/copy 差异非常重要
    - view 只改元数据(shape/stride/offset),$O(1)$ 且共享存储;
    - contiguous() 会物理复制,占内存也占带宽——热点路径上要警惕。

  3. 矩阵乘法主导 FLOPs,但不一定主导时间
    - matmul 是 compute-bound(shape 够大时);
    - elementwise、norm、softmax 往往 memory-bound——算术强度太低,瓶颈在 HBM 带宽而非算力。

  4. 反向传播约是前向的 2 倍
    - 每层线性变换的反向要算 $\nabla_W$ 和 $\nabla_X$ 两个与前向同规模的矩阵乘法;
    - 总训练 FLOPs ≈ 前向的 3 倍,即 $6ND$。

  5. Optimizer state 是显存大户
    - AdamW 的 m/v(加 master weights)让「每参数字节数」翻到参数本身的 4 倍;
    - 混合精度不省优化器状态——所以显存优化的主战场是 ZeRO/FSDP 式的切分。

下一讲预告(Lecture 3 · 架构与超参数

有了 PyTorch 训练积木和资源核算的本能后,下一步是决定模型长什么样:Pre-Norm、RMSNorm、SwiGLU、RoPE、GQA/MQA、FFN 维度、head 数、词表大小,以及稳定训练所需的 z-loss、QK norm 等技巧——每一个选择背后依然是本讲的两本账:内存与 FLOPs。


复习自测

题目

对每层 $H = XW$,前向是 1 个 matmul($2BDK$ FLOPs)。反向要回答两个问题:参数怎么改——$\nabla_W = X^\top \nabla_H$;责任怎么往前传——$\nabla_X = \nabla_H W^\top$。两个都是与前向同规模的 matmul,共 $4BDK$。逐层累加即得前向 $2ND$、反向 $4ND$、合计 $6ND$。

题目

fp16 只有 5 位指数,最小正规数约 $6\times10^{-5}$,反向传播中大量小梯度会下溢变零,必须用 loss scaling 先放大再缩回。bf16 有 8 位指数、动态范围与 fp32 相同,梯度不会因量级太小而下溢,代价是尾数只剩 7 位、精度更粗——而训练对精度噪声远比对「信号归零」鲁棒。

题目

每参数 = 4(参数)+ 4(梯度)+ 4(m)+ 4(v)= 16 字节。单卡:$80\times10^9 / 16 = 5\times10^9$,约 5B。8 卡完美切分:$640\times10^9 / 16 = 40\times10^9$,约 40B 上限。真实可训规模还要再扣掉激活、buffer 与碎片。

题目

$C = 6 \times 70\times10^9 \times 15\times10^{12} = 6.3\times10^{24}$ FLOPs。有效吞吐 $= 989.5\times10^{12} \times 0.4 \times 1024 \approx 4.05\times10^{17}$ FLOP/s。时间 $= 6.3\times10^{24} / 4.05\times10^{17} \approx 1.55\times10^7$ 秒 ≈ 180 天。注意峰值用的是稠密 BF16(989.5 TFLOP/s),不是营销的稀疏数字。

题目

transpose 只交换 stride,不搬数据——结果的内存布局不再是行优先连续,而 view 要求新形状能用「同一段连续 storage + 新 stride」表达,对 non-contiguous 布局无法做到,故报错。contiguous() 按新布局完整复制数据后再 view:正确性恢复了,但多花了一份内存和一次全量带宽读写——在热点路径上这可能比计算本身还贵。


参考资料