第 08 章 · 进阶
稳定训练:初始化、AMP、裁剪与排错
让模型不仅能跑,还能稳定地学:掌握初始化、混合精度、梯度累积、裁剪、监控和 NaN 定位。
初始化如何影响信号与梯度
本节目标
本节围绕“初始化如何影响信号与梯度”展开。运行示例后记录 loss、梯度范数、学习率、缩放因子与异常值,比较配置变化前后。
清晰讲解
初始化要让前向激活和反向梯度在多层传播时不过度放大或缩小。Linear 默认初始化通常够用;自定义结构应根据激活与 fan-in/fan-out 选择 Xavier 或 Kaiming,并观察实际统计。
核对“初始化如何影响信号与梯度”中的 loss、gradient norm、学习率、数值范围和更新是否连续。
代码示例
import torch
from torch import nn
layer = nn.Linear(1024, 4096)
nn.init.xavier_uniform_(layer.weight)
nn.init.zeros_(layer.bias)
x = torch.randn(32, 1024)
y = layer(x)
print(x.std().item(), y.std().item())
运行结果与观察
输入输出标准差应处在同一数量级;若层数增加后快速趋零或爆炸,需重新检查初始化与残差缩放。
常见错误
- 所有权重统一正态且方差过大;忘记 bias;初始化后又被
reset_parameters覆盖。 - 排查“初始化如何影响信号与梯度”时只看最终 loss,没有保存梯度范数、学习率、AMP 状态或 NaN/Inf 出现位置。
与大模型方向的连接
深层 Transformer 还会对残差分支做缩放;不要脱离架构照搬某个初始化常数。
动手练习
记录 12 层 MLP 每层输出标准差,观察普通正态初始化的漂移。
查看参考答案与验收点
用一个极小 batch 过拟合,并记录 loss、梯度范数和参数更新;断点恢复后比较下一步结果,而不只检查文件存在。
官方资料
AMP:autocast 与 GradScaler 分工
本节目标
本节围绕“AMP:autocast 与 GradScaler 分工”展开。运行示例后记录 loss、梯度范数、学习率、缩放因子与异常值,比较配置变化前后。
清晰讲解
autocast 为不同算子选择合适精度,GradScaler 在 float16 训练时放大 loss,降低小梯度下溢风险。bfloat16 通常不需要 scaler。优化器看到梯度前如要裁剪,应先 unscale_。
核对“AMP:autocast 与 GradScaler 分工”中的 loss、gradient norm、学习率、数值范围和更新是否连续。
代码示例
scaler = torch.amp.GradScaler('cuda')
for x, y in loader:
optimizer.zero_grad(set_to_none=True)
with torch.autocast('cuda', dtype=torch.float16):
loss = loss_fn(model(x), y)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
运行结果与观察
CUDA 上 autocast 区域的矩阵输出通常为低精度,而 loss/关键归约保持合适精度;缩放后梯度应为有限值,发生溢出时 scale 会下降并跳过更新。
常见错误
- 把模型和输入全部粗暴 half;裁剪缩放后的梯度;CPU 上写死 CUDA autocast。
- 排查“AMP:autocast 与 GradScaler 分工”时只看最终 loss,没有保存梯度范数、学习率、AMP 状态或 NaN/Inf 出现位置。
与大模型方向的连接
大模型训练普遍依赖低精度;理解哪些状态仍保持高精度,比简单 .half() 更重要。
动手练习
给循环加入设备判断,在 CUDA 不可用时自动回退 float32。
查看参考答案与验收点
验收:监控项足以复现并定位异常,并能说明“AMP:autocast 与 GradScaler 分工”影响稳定性的具体路径。
官方资料
梯度累积如何模拟更大 batch
本节目标
本节围绕“梯度累积如何模拟更大 batch”展开。运行示例后记录 loss、梯度范数、学习率、缩放因子与异常值,比较配置变化前后。
清晰讲解
显存放不下目标 batch 时,把它拆成多个 micro-batch,累积梯度后统一更新。若每个 micro loss 是均值,需要除以累积步数,才能保持与大 batch 相近的梯度尺度。
核对“梯度累积如何模拟更大 batch”中的 loss、gradient norm、学习率、数值范围和更新是否连续。
代码示例
accum_steps = 4
optimizer.zero_grad(set_to_none=True)
for micro_step, (x, y) in enumerate(loader, 1):
loss = loss_fn(model(x), y) / accum_steps
loss.backward()
if micro_step % accum_steps == 0:
optimizer.step()
optimizer.zero_grad(set_to_none=True)
运行结果与观察
在无 dropout、相同样本与等价 loss 归约下,累积多个 micro-batch 得到的参数更新应接近一次大 batch;若差异大,通常漏除了累积步数。
常见错误
- 最后不足累积步数的 batch 未更新;日志显示被除后的 loss;DDP 每个 micro-step 都同步梯度。
- 排查“梯度累积如何模拟更大 batch”时只看最终 loss,没有保存梯度范数、学习率、AMP 状态或 NaN/Inf 出现位置。
与大模型方向的连接
全局 batch 通常等于 micro batch × 累积步数 × 数据并行进程数;学习率配方要基于真正的全局 batch。
动手练习
处理 dataloader 末尾不足 4 步的情况,并记录未缩放 loss。
查看参考答案与验收点
用一个极小 batch 过拟合,并记录 loss、梯度范数和参数更新;断点恢复后比较下一步结果,而不只检查文件存在。
官方资料
梯度裁剪解决什么、不解决什么
本节目标
本节围绕“梯度裁剪解决什么、不解决什么”展开。运行示例后记录 loss、梯度范数、学习率、缩放因子与异常值,比较配置变化前后。
清晰讲解
按全局范数裁剪限制一次更新的最大梯度规模,能缓解偶发尖峰,但不能修复错误数据、过大学习率或 NaN 算子。裁剪前记录原始范数,才能知道它是否频繁触发。
核对“梯度裁剪解决什么、不解决什么”中的 loss、gradient norm、学习率、数值范围和更新是否连续。
代码示例
loss.backward()
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
if not torch.isfinite(total_norm):
raise RuntimeError(f'non-finite grad norm: {total_norm}')
optimizer.step()
运行结果与观察
裁剪前总范数应大于阈值,裁剪后重新计算应不超过阈值加浮点误差;已有 NaN 的梯度不会被“修好”,应触发异常排查。
常见错误
- 先 optimizer.step 再裁剪;只裁某一小部分参数;把裁剪当成所有不稳定问题的万能药。
- 排查“梯度裁剪解决什么、不解决什么”时只看最终 loss,没有保存梯度范数、学习率、AMP 状态或 NaN/Inf 出现位置。
与大模型方向的连接
训练大模型时梯度范数是关键健康指标;突然抬升常比 loss 更早暴露异常 batch。
动手练习
制造一个异常大 loss,比较裁剪前后参数更新范数。
查看参考答案与验收点
用一个极小 batch 过拟合,并记录 loss、梯度范数和参数更新;断点恢复后比较下一步结果,而不只检查文件存在。
官方资料
系统定位 NaN 与 Inf
本节目标
本节围绕“系统定位 NaN 与 Inf”展开。运行示例后记录 loss、梯度范数、学习率、缩放因子与异常值,比较配置变化前后。
清晰讲解
排查顺序应从输入、loss、梯度到参数更新,定位第一个非有限值。先缩小到单 batch、关闭复杂优化,再用 anomaly detection 或 hooks 精确到算子。
核对“系统定位 NaN 与 Inf”中的 loss、gradient norm、学习率、数值范围和更新是否连续。
代码示例
torch.autograd.set_detect_anomaly(True)
logits = model(batch)
if not torch.isfinite(logits).all(): raise ValueError('bad logits')
loss = loss_fn(logits, labels)
if not torch.isfinite(loss): raise ValueError('bad loss')
loss.backward()
for name, p in model.named_parameters():
if p.grad is not None and not torch.isfinite(p.grad).all():
raise ValueError(f'bad grad: {name}')
运行结果与观察
打开 anomaly detection 后,故意构造非法运算应给出产生非有限梯度的算子栈;逐层检查时记录第一个从有限变为非有限的张量,而不是只看最终 loss。
常见错误
- 长期打开 anomaly detection 导致极慢;只检查最终 loss;异常后仍保存覆盖好 checkpoint。
- 排查“系统定位 NaN 与 Inf”时只看最终 loss,没有保存梯度范数、学习率、AMP 状态或 NaN/Inf 出现位置。
与大模型方向的连接
大模型 NaN 可能来自脏 token、极端长度、低精度溢出或通信;先找“第一个坏点”,不要只调低学习率碰运气。
动手练习
写一个函数返回第一个出现非有限梯度的参数名。
查看参考答案与验收点
验收:监控项足以复现并定位异常,并能说明“系统定位 NaN 与 Inf”影响稳定性的具体路径。
官方资料
最小可复现实验与训练监控
本节目标
本节围绕“最小可复现实验与训练监控”展开。运行示例后记录 loss、梯度范数、学习率、缩放因子与异常值,比较配置变化前后。
清晰讲解
有效监控至少包含训练/验证 loss、学习率、tokens/s、显存峰值、梯度范数和数据进度。遇到问题时保存配置与小批样本,构造能独立复现的最小脚本。
核对“最小可复现实验与训练监控”中的 loss、gradient norm、学习率、数值范围和更新是否连续。
代码示例
metrics = {
'step': step,
'loss': float(loss.detach()),
'lr': optimizer.param_groups[0]['lr'],
'max_memory_mib': torch.cuda.max_memory_allocated() / 1024**2 if torch.cuda.is_available() else 0,
}
print(metrics)
运行结果与观察
固定配置与随机种子重复运行时,关键指标曲线应在声明的容差内一致;日志至少能还原代码版本、数据版本、超参数、设备和最佳 checkpoint。
常见错误
- 每步
.item()造成频繁 GPU 同步;只记录均值掩盖尖峰;没有配置快照。 - 排查“最小可复现实验与训练监控”时只看最终 loss,没有保存梯度范数、学习率、AMP 状态或 NaN/Inf 出现位置。
与大模型方向的连接
tokens/s 比 samples/s 更适合变长文本;同时报告有效 token 数,性能比较才公平。
动手练习
每 20 步聚合一次指标,避免每步同步,并计算 tokens/s。
查看参考答案与验收点
用一个极小 batch 过拟合,并记录 loss、梯度范数和参数更新;断点恢复后比较下一步结果,而不只检查文件存在。
官方资料
章末资料
小测、项目与相关面试题
1. AMP 中为何裁剪前要 unscale?
否则裁剪的是被 scaler 放大的梯度,阈值失去原有意义。
2. 梯度裁剪能修复错误数据吗?
不能。它限制更新幅度,但不会修复 NaN 算子、脏数据或错误标签。
