训练一个语言模型,真正落到 PyTorch 里以后,很多问题都会变成很具体的计算问题:一个 Tensor 怎么存,矩阵乘法到底有多少 FLOP,反向传播为什么比前向更贵,训练时显存里除了参数还存了什么。
这篇就从这些底层问题开始,把一次完整训练过程拆开来看。
1. 数据类型
训练里最常见的几种浮点格式:
FP32:1 位符号位、8 位指数、23 位尾数;
FP16:1 + 5 + 10;
BF16:1 + 8 + 7;
FP8:常见为 1 + 4 + 3 或 1 + 5 + 2。
BF16 保留了和 FP32 一样的指数位数,因此表示范围接近 FP32,同时内存占用只有一半。
实际训练通常不会全程使用同一种精度,而是使用混合精度:
需要长期保存、对数值稳定性要求高的部分保留高精度;
前向传播等临时计算使用 BF16;
推理阶段通常还可以进一步降低精度。
2. Tensor 的存储
PyTorch Tensor 本质上是对一块底层连续存储的描述。
除了数据地址之外,还需要记录每个维度的 stride。例如一个二维 Tensor:
PYTHON
x.stride(0) == 4
x.stride(1) == 1意味着:
沿第 0 维移动一次,需要跨过 4 个元素;
沿第 1 维移动一次,只需要跨过 1 个元素。
这也是为什么部分切片和转置可以只创建 View,而不复制数据。但 View 依赖底层存储布局。需要连续存储时,可以调用:x.contiguous()。这一步会产生实际的数据复制。
3. 矩阵乘法与 FLOP
LLM 中绝大部分计算成本来自矩阵乘法。假设:
对于输出中的每一个元素,都需要完成 (B) 次乘法和累加,因此 FLOP 数量近似为:。对于线性层来说,可以进一步得到一个常用估算:
这也是训练时经常用参数量和 token 数快速估算计算成本的原因。
4. MFU
MFU(Model FLOPs Utilization)用于衡量实际模型计算对硬件理论算力的利用程度:
实际训练中不可能把所有时间都花在矩阵计算上,通信、访存以及其他操作都会占据时间,因此 MFU 更适合作为整体训练效率的指标。实际上,MFU 能达到 50% 以上已经是比较好的利用率。
5. 反向传播的计算量
假设有两层线性计算:
flowchart LR
X["x"] -->|W1| H1["h1"]
H1 -->|W2| H2["h2"]
H2 --> L["loss"]
第二层:
反向传播时:
同时还需要计算:
这两次矩阵乘法的 FLOP 数量和前向传播是同一级别。因此可以粗略理解为:
前向传播:约为参数量的 2 倍计算;
反向传播:约为参数量的 4 倍计算。
总的来说,训练一次的总计算成本通常约为前向传播的 3 倍。
6. 参数初始化
深层网络里,如果每一层都直接使用方差较大的随机参数,激活值的尺度会随着层数不断变化,最终导致训练不稳定。
Xavier 初始化的思路是根据输入维度缩放初始化权重。对于输入维度 (B),先从正态分布采样,然后按 进行缩放,使每层输出的尺度尽量保持稳定。还可以进一步做区间截断,避免采样到过大的初始值。
7. 优化器
最基本的梯度下降为:
SGD
SGD 每次只使用随机的小批量数据计算梯度。
Momentum
Momentum 为梯度增加历史惯性:
常用的 为 0.9。
Adagrad
Adagrad 会根据历史梯度平方和缩放当前梯度。参数历史梯度越大,后续有效学习率越小。问题是平方和会不断累积,学习率可能最终降得很低。
RMSProp
RMSProp 将 Adagrad 的历史平方和改成指数移动平均,避免分母无限增长。
Adam
Adam 可以理解为 Momentum 和 RMSProp 的组合:
一阶动量控制梯度方向的惯性;
二阶统计量控制不同参数之间的更新尺度。
Muon
Muon保留动量,并对动量矩阵进行正交化处理,用来减小矩阵不同方向之间的尺度差异。
8. 训练时到底要存什么
模型训练时,显存里不只是模型参数。通常还需要保存:
模型权重;
梯度;
前向传播中间激活;
优化器状态。
其中优化器状态可能非常大。以 Adam 为例,需要额外保存一阶和二阶统计量。这也是后面做分布式训练时,ZeRO / FSDP 首先会从 Optimizer State、Gradient 和 Parameter 的切分开始。
9. Checkpoint
一个完整的训练 Checkpoint 至少需要包含:
模型参数;
优化器参数。
否则即使恢复了模型权重,优化器内部的动量和历史统计也会丢失,训练状态并没有真正恢复。
