AI 工程基础体系 · 第 9/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。

深度学习训练:优化器、学习率、批次、混合精度和稳定性

深度学习训练可以看作一个受数据、模型、数值表示和资源约束的迭代系统:

  1. 从数据集中取出一个批次;
  2. 模型执行前向计算,得到预测结果;
  3. 损失函数比较预测值与目标值;
  4. 反向传播计算参数梯度;
  5. 优化器根据梯度更新参数;
  6. 记录指标、保存状态,并在验证集上检查是否真正变好。

优化器、学习率和批次决定“参数如何移动”;混合精度决定“这些计算以什么数值格式进行”;稳定性则要求上述过程在有限精度、有限显存和有限数据质量下仍然可控。


训练问题的数学形式

设模型参数为

θRd\theta \in \mathbb{R}^d

模型对输入 xx 的输出为

y^=fθ(x)\hat{y}=f_\theta(x)

损失函数为

(fθ(x),y)\ell(f_\theta(x),y)

其中 yy 是标签。训练集包含 NN 个样本,经验风险通常写为:

L(θ)=1Ni=1N(fθ(xi),yi)L(\theta)=\frac{1}{N}\sum_{i=1}^{N}\ell(f_\theta(x_i),y_i)

理想目标是寻找使 L(θ)L(\theta) 较小的参数:

θ\*=argminθL(θ)\theta^\*=\arg\min_\theta L(\theta)

但实际训练不会每一步都计算全部 NN 个样本的梯度,而是抽取批次 BtB_t,计算小批次损失:

LBt(θ)=1BtiBt(fθ(xi),yi)L_{B_t}(\theta)=\frac{1}{|B_t|}\sum_{i\in B_t}\ell(f_\theta(x_i),y_i)

其梯度为:

gt=θLBt(θt)g_t=\nabla_\theta L_{B_t}(\theta_t)

如果批次是从训练集均匀抽样得到的,则常见情况下:

E[gt]θL(θt)\mathbb{E}[g_t]\approx \nabla_\theta L(\theta_t)

gtg_t 不是完全准确的全量梯度:

gt=L(θt)+ϵtg_t=\nabla L(\theta_t)+\epsilon_t

其中 ϵt\epsilon_t 是由批次抽样造成的梯度噪声。批次大小、学习率和优化器,本质上都在控制如何利用这份带噪声的方向。


一个训练步骤到底改变了什么

对一个批次执行训练时,参数状态通常经历以下变化:

flowchart LR
    A[DataLoader 取出 batch] --> B[前向计算]
    B --> C[计算 loss]
    C --> D[反向传播得到梯度]
    D --> E{梯度是否有限}
    E -- 否 --> F[跳过更新并记录故障]
    E -- 是 --> G[梯度裁剪 可选]
    G --> H[优化器更新参数]
    H --> I[更新学习率调度器]
    I --> J[记录指标或保存检查点]

一个关键细节是:损失计算、反向传播、优化器更新和学习率调度并不是同一个动作

  • loss.backward() 只计算并累积梯度,不修改参数。
  • optimizer.step() 根据梯度修改参数,同时更新优化器内部状态。
  • optimizer.zero_grad() 清理上一轮残留梯度。
  • scheduler.step() 修改优化器中的学习率,具体应该在每个 batch 还是每个 epoch 调用,取决于调度器设计。
  • model.train()model.eval() 改变 Dropout、BatchNorm 等模块的行为,但不会更新参数。

如果忘记清理梯度,PyTorch 默认会将新梯度加到旧梯度上:

param.gradparam.grad+gt\texttt{param.grad}\leftarrow \texttt{param.grad}+g_t

这在梯度累积时是有意行为,在普通训练中则会导致错误。


优化器:从梯度到参数更新

随机梯度下降

最基本的随机梯度下降(SGD)更新为:

θt+1=θtηtgt\theta_{t+1}=\theta_t-\eta_t g_t

其中:

  • θt\theta_t 是第 tt 步的参数;
  • gtg_t 是当前批次的梯度;
  • ηt>0\eta_t>0 是学习率。

学习率越大,每一步移动越远;学习率越小,每一步移动越近。

一维完整算例

假设当前只有一个参数:

θ0=2\theta_0=2

损失函数为:

L(θ)=(θ5)2L(\theta)=(\theta-5)^2

其梯度是:

L(θ)=2(θ5)\nabla L(\theta)=2(\theta-5)

取学习率 η=0.1\eta=0.1,在 θ0=2\theta_0=2 处:

g0=2(25)=6g_0=2(2-5)=-6

因此:

θ1=20.1×(6)=2.6\theta_1=2-0.1\times(-6)=2.6

再次计算:

g1=2(2.65)=4.8g_1=2(2.6-5)=-4.8

于是:

θ2=2.60.1×(4.8)=3.08\theta_2=2.6-0.1\times(-4.8)=3.08

参数逐渐接近最优点 55

如果学习率改为 η=1.1\eta=1.1,第一次更新为:

θ1=21.1×(6)=8.6\theta_1=2-1.1\times(-6)=8.6

这一步越过了最优点,而且距离从 33 变成了 3.63.6。对于这个简单二次函数,梯度下降在学习率过大时会来回震荡甚至发散。一般地,对

L(θ)=12a(θθ\*)2L(\theta)=\frac{1}{2}a(\theta-\theta^\*)^2

有:

θt+1θ\*=(1ηa)(θtθ\*)\theta_{t+1}-\theta^\*=(1-\eta a)(\theta_t-\theta^\*)

要收敛,需要:

1ηa<1|1-\eta a|<1

也就是:

0<η<2a0<\eta<\frac{2}{a}

真实神经网络通常是非凸的,且不同参数方向的曲率不同,因此不存在一个简单的全局“正确学习率”。这也是学习率搜索、预热和衰减有价值的原因。


Momentum:让更新保留运动方向

SGD 的每一步只看当前梯度。Momentum 会维护一个速度状态:

vt=βvt1+gtv_t=\beta v_{t-1}+g_t

θt+1=θtηvt\theta_{t+1}=\theta_t-\eta v_t

其中 β\beta 通常接近 1,例如 0.9。

如果多个批次的梯度方向大致一致,历史梯度会累积,参数沿稳定方向加速;如果梯度在某个方向上来回变化,动量会抵消部分震荡。

仍以一维问题为例,令 β=0.9\beta=0.9,初始 v0=0v_0=0,第一次梯度 g0=6g_0=-6

v1=0.9×06=6v_1=0.9\times0-6=-6

η=0.1\eta=0.1

θ1=20.1×(6)=2.6\theta_1=2-0.1\times(-6)=2.6

第二次梯度 g1=4.8g_1=-4.8

v2=0.9×(6)4.8=10.2v_2=0.9\times(-6)-4.8=-10.2

θ2=2.60.1×(10.2)=3.62\theta_2=2.6-0.1\times(-10.2)=3.62

相比不带 Momentum 的第二步结果 3.083.08,它更快向目标移动,但也更依赖合适的学习率。


Adam:按参数维护一阶和二阶统计量

Adam 同时维护梯度的一阶矩和二阶矩:

mt=β1mt1+(1β1)gtm_t=\beta_1m_{t-1}+(1-\beta_1)g_t

vt=β2vt1+(1β2)gt2v_t=\beta_2v_{t-1}+(1-\beta_2)g_t^2

由于初始 m0=v0=0m_0=v_0=0,早期统计量会偏向 0,因此进行偏差修正:

m^t=mt1β1t\hat{m}_t=\frac{m_t}{1-\beta_1^t}

v^t=vt1β2t\hat{v}_t=\frac{v_t}{1-\beta_2^t}

参数更新为:

θt+1=θtηm^tv^t+ϵ\theta_{t+1} = \theta_t-\eta \frac{\hat{m}_t}{\sqrt{\hat{v}_t}+\epsilon}

其中:

  • mtm_t 估计梯度的平均方向;
  • vtv_t 估计梯度平方的平均大小;
  • ϵ\epsilon 防止分母为零;
  • β1,β2\beta_1,\beta_2 控制历史信息的平滑程度。

以单个参数、g1=6g_1=-6β1=0.9\beta_1=0.9β2=0.999\beta_2=0.999 为例:

m1=0.1×(6)=0.6m_1=0.1\times(-6)=-0.6

v1=0.001×36=0.036v_1=0.001\times36=0.036

偏差修正后:

m^1=6,v^1=36\hat m_1=-6,\qquad \hat v_1=36

所以归一化方向约为:

636=1\frac{-6}{\sqrt{36}}=-1

第一次更新的幅度接近 η\eta,而不是直接等于 ηg1\eta |g_1|。这使 Adam 对不同参数的梯度尺度不那么敏感,但不意味着可以忽略学习率。

Adam 与 AdamW 的权重衰减区别

权重衰减(weight decay)通常用于限制参数过大,从而改善泛化。对 SGD,直接把 L2 正则项加入损失,和对参数做衰减在形式上比较接近:

Lreg(θ)=L(θ)+λ2θ2L_{\text{reg}}(\theta)=L(\theta)+\frac{\lambda}{2}\|\theta\|^2

其梯度为:

Lreg=L+λθ\nabla L_{\text{reg}}=\nabla L+\lambda\theta

但对 Adam,λθ\lambda\theta 会进入自适应的一阶、二阶统计,实际效果不再等价于简单地缩小参数。

AdamW 将权重衰减与梯度更新解耦:

θt+1=(1ηλ)θtηm^tv^t+ϵ\theta_{t+1} = (1-\eta\lambda)\theta_t - \eta \frac{\hat m_t}{\sqrt{\hat v_t}+\epsilon}

PyTorch 中常见写法是:

optimizer = torch.optim.AdamW(
    model.parameters(),
    lr=3e-4,
    betas=(0.9, 0.999),
    weight_decay=0.01,
)

这不是说 AdamW 在所有任务上必然优于 SGD,而是它的衰减机制更符合“独立缩小参数”的设计意图。偏置项、归一化层参数通常不一定适合与普通权重使用相同的 weight decay,生产训练中常通过参数分组分别设置。


学习率:决定更新尺度和训练时间表

固定学习率的局限

固定学习率可能在训练初期过大、后期过大或始终过小:

  • 初期参数尚未进入合适区域,大步更新可能导致 loss 爆炸;
  • 中期需要较大步长快速下降;
  • 后期接近较优区域时,需要更小步长稳定收敛;
  • 学习率过小会使训练看起来“稳定”,但实际上几乎没有学习。

学习率必须结合优化器解释。Adam 中的学习率控制的是归一化梯度的基本步长;SGD 中则直接与梯度绝对尺度相乘。

Warmup

Warmup 是在训练开始阶段逐渐增大学习率。例如线性 warmup:

ηt=ηmaxtTwarmup(1tTwarmup)\eta_t=\eta_{\max}\frac{t}{T_{\text{warmup}}} \quad (1\le t\le T_{\text{warmup}})

原因通常包括:

  • Adam 的矩估计在初期还未稳定;
  • 大批次训练初始梯度可能具有较大方差;
  • Transformer 中残差、注意力和归一化的组合对初始更新较敏感;
  • 从预训练模型微调时,过大的初始更新可能破坏已有表示。

Warmup 不是“越长越好”。过长会浪费训练预算,使有效学习率长期偏低。

衰减策略

常见衰减方法包括:

Step 衰减

每隔若干 epoch 将学习率乘以 γ\gamma

ηk+1=γηk\eta_{k+1}=\gamma\eta_k

指数衰减

ηt=η0γt\eta_t=\eta_0\gamma^t

Cosine 衰减

ηt=ηmin+12(ηmaxηmin)(1+cosπtT)\eta_t = \eta_{\min} + \frac{1}{2}(\eta_{\max}-\eta_{\min}) \left(1+\cos\frac{\pi t}{T}\right)

其中 TT 是计划中的总更新步数。

One-cycle

先增大学习率,再衰减到较低值,常用于有限训练预算下的实验。它改变的不只是最终学习率,也改变了训练过程中的噪声和探索程度。

调度器的步进单位必须明确:

  • 每个 epoch 调一次:tt 表示 epoch;
  • 每个 optimizer update 调一次:tt 表示参数更新次数;
  • 采用梯度累积时,通常应按真实的参数更新次数调度,而不是按每个 micro-batch 调度。

如果每 8 个 micro-batch 才更新一次参数,却每个 micro-batch 调用一次 scheduler,学习率计划会比预期快 8 倍。


批次:统计估计、显存和并行度的共同约束

三种“批次大小”

工程中至少要区分:

  1. micro-batch size:一次前向和反向实际放入单张设备显存的样本数;
  2. per-device batch size:每张设备每次处理的样本数;
  3. global batch size:分布式训练中所有设备一次更新共同使用的样本数;
  4. effective batch size:考虑梯度累积后的等效批次大小。

若有 DD 张设备,每张设备的 micro-batch 为 bb,梯度累积步数为 AA,则:

Beffective=D×b×AB_{\text{effective}}=D\times b\times A

这个公式假设每个 micro-batch 等权,且每一步都包含相同数量的有效样本。

批次大小对梯度噪声的影响

设单样本梯度为 gig_i,批次梯度为:

gB=1Bi=1Bgig_B=\frac{1}{B}\sum_{i=1}^{B}g_i

若样本近似独立,梯度估计的方差通常随 BB 增大而下降,近似为:

Var(gB)1B\operatorname{Var}(g_B)\propto\frac{1}{B}

小批次的特点是:

  • 梯度噪声大;
  • 每次参数更新成本较低;
  • 更新次数多;
  • 噪声有时有助于跳出尖锐区域或改善泛化;
  • loss 曲线更抖动。

大批次的特点是:

  • 梯度估计更接近全量梯度;
  • 单次更新更稳定;
  • 需要更多显存或通信;
  • 在固定样本数下,参数更新次数更少;
  • 不一定带来更好的验证集效果。

因此“大批次更稳定”只描述训练曲线的一个方面,不等同于“泛化更好”。

梯度累积的正确形式

若一个有效批次由 AA 个 micro-batch 构成,应该将每个 micro-batch 的损失除以 AA

Lmicro,j=1biBjiL_{\text{micro},j} = \frac{1}{b}\sum_{i\in B_j}\ell_i

反向传播:

1ALmicro,j\frac{1}{A}\nabla L_{\text{micro},j}

累积后:

j=1A1ALmicro,j=1Aj=1ALmicro,j\sum_{j=1}^{A}\frac{1}{A}\nabla L_{\text{micro},j} = \frac{1}{A}\sum_{j=1}^{A}\nabla L_{\text{micro},j}

这才是各 micro-batch 平均梯度。

如果忘记除以 AA,累积梯度大约会变成原来的 AA 倍。虽然可以同时把学习率缩小 AA 倍抵消参数更新幅度,但梯度裁剪、混合精度溢出、优化器状态和日志中的梯度范数都会改变,因此不应依赖这种补偿。

还要注意最后一个不完整累积组。如果数据量不能被 AA 整除,最后一次更新中的有效 micro-batch 数可能小于 AA。严格处理时应使用实际累积数量归一化,或者设置 drop_last=True。对于 token 级语言模型,若不同序列的有效 token 数不同,按“序列平均 loss”再平均,和按“所有有效 token 总和”计算 loss,并不等价。


混合精度:降低成本,但不是免费加速

混合精度不是“把所有张量都改成半精度”,而是:

  • 对适合低精度的算子使用 float16bfloat16
  • 对敏感操作保留 float32
  • 通常以 float32 保存模型参数和优化器状态;
  • float16 梯度使用 loss scaling,避免小梯度下溢。

为什么低精度会出问题

IEEE 浮点格式同时受精度和表示范围限制。

  • float16 的表示范围较小,容易出现 overflow;
  • 很小的梯度可能下溢为 0;
  • bfloat16 指数范围接近 float32,不易 overflow,但有效尾数较少,舍入误差更明显;
  • 某些归约、归一化和 softmax 操作对数值误差更敏感。

因此混合精度的目标是让计算吞吐和显存占用受益,同时把关键状态保留在更可靠的格式中。

autocast

PyTorch 的 autocast 会根据算子类型选择适合的计算 dtype。典型 CUDA 用法:

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

autocast 的选择是常见实现策略,不应理解成对所有模型和所有算子的数学精确保证。自定义算子、第三方扩展和不支持 autocast 的算子需要单独验证。

GradScaler

float16 下,梯度可能太小而下溢。GradScaler 先把 loss 乘以一个缩放因子 ss

L=sLL'=sL

反向传播得到:

L=sL\nabla L'=s\nabla L

在更新前再除以 ss,恢复原始梯度。如果检测到 infnan,则跳过这次参数更新并减小缩放因子。

PyTorch 端到端训练片段如下:

use_amp = device.type == "cuda"
amp_dtype = torch.float16

scaler = torch.amp.GradScaler(
    "cuda",
    enabled=use_amp,
)

for x, y in train_loader:
    x = x.to(device, non_blocking=True)
    y = y.to(device, non_blocking=True)

    optimizer.zero_grad(set_to_none=True)

    with torch.autocast(
        device_type=device.type,
        dtype=amp_dtype,
        enabled=use_amp,
    ):
        logits = model(x)
        loss = criterion(logits, y)

    scaler.scale(loss).backward()
    scaler.unscale_(optimizer)

    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

    scaler.step(optimizer)
    scaler.update()

这里有几个不可交换的顺序:

  1. scale(loss).backward() 在缩放后的 loss 上反向传播;
  2. unscale_(optimizer) 必须在梯度裁剪前执行,否则裁剪的是放大后的梯度;
  3. scaler.step(optimizer) 会在梯度有限时执行真实更新;
  4. scaler.update() 根据溢出情况调整缩放因子。

GradScaler 不是数学上保证训练稳定的工具。即使使用 scaler,前向计算也可能产生 nan,或者模型本身已经因为学习率过大而发散。

float16 与 bfloat16

在支持 bfloat16 的硬件上,可以使用:

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

通常 bfloat16 的动态范围更接近 float32,因此很多场景下对 loss scaling 的依赖较小。不过是否使用 GradScaler、哪些算子保留 float32,仍取决于 PyTorch 版本、硬件和模型结构。不能把“bfloat16 不容易溢出”误解成“完全不会发生数值问题”。


稳定性:从数据、梯度到状态的一整条链路

训练不稳定不只意味着 loss 变成 NaN。以下现象都可能表示稳定性问题:

  • loss 在初期突然上升多个数量级;
  • 梯度范数持续增大;
  • 验证指标快速恶化;
  • 梯度大量为 0;
  • logits 变成 inf
  • 每次运行结果差异异常大;
  • 恢复检查点后训练轨迹突然改变。

数据稳定性

在调优化器之前,应先确认数据和目标正确:

  • 分类标签是否从 0num_classes - 1
  • CrossEntropyLoss 的输入是否是未归一化 logits,而不是 softmax 后的概率;
  • 回归目标是否存在极端异常值;
  • 输入归一化统计量是否只使用训练集计算;
  • 训练集和验证集是否发生泄漏;
  • padding、mask 和 attention mask 是否语义一致;
  • token、图像或音频预处理是否与推理阶段一致。

例如,CrossEntropyLoss 通常内部执行 log-softmax 和负对数似然。如果先手动执行 softmax 再传入,可能造成数值稳定性和梯度质量变差。应优先传入原始 logits:

logits = model(x)          # shape: [batch, num_classes]
loss = torch.nn.functional.cross_entropy(logits, labels)

梯度裁剪

梯度裁剪用于限制一次更新的最大梯度规模。按全局范数裁剪时,若:

g2>c\|g\|_2>c

则替换为:

g=cg2gg'=\frac{c}{\|g\|_2}g

其中 ccmax_norm

它可以降低偶发异常批次造成的巨大更新,但不能修复:

  • 标签全部错误;
  • 学习率高出合理范围;
  • 模型输出和损失函数不匹配;
  • 输入中存在 nan
  • 数据分布发生根本变化。

在混合精度下必须先 unscale:

scaler.unscale_(optimizer)
grad_norm = torch.nn.utils.clip_grad_norm_(
    model.parameters(),
    max_norm=1.0,
)

grad_norm 可用于日志和报警。若裁剪几乎每一步都触发,通常说明学习率、初始化、归一化或数据尺度需要重新检查,而不是无限增大裁剪阈值。

梯度消失与爆炸

若网络中连续出现导数小于 1 的变换,反向传播中的梯度会连乘并逐层缩小;若连续出现导数大于 1,则可能快速增大。循环结构、深层 Transformer、极端 logits 和不当初始化都可能放大该问题。

常见缓解手段及其因果关系包括:

  • 合理初始化,使初始激活和梯度尺度可控;
  • 使用合适的归一化,降低层间尺度漂移;
  • 对深层残差结构采用成熟的架构设计;
  • 降低初始学习率并使用 warmup;
  • 对异常大梯度进行裁剪;
  • 检查输入、标签和损失的数值范围。

归一化并不自动保证稳定。BatchNorm 依赖批次统计量,小批次或分布式设置下统计量可能不可靠;LayerNorm 通常按单个样本的特征维度归一化,更常见于 Transformer。两者的训练/推理行为和状态管理不同,不能仅凭“有归一化层”判断模型稳定。

NaN/Inf 诊断

可以在关键位置检查:

def assert_finite(name, tensor):
    if not torch.isfinite(tensor).all():
        raise FloatingPointError(f"{name} contains NaN or Inf")

assert_finite("input", x)

with torch.autocast(
    device_type=device.type,
    dtype=torch.float16,
    enabled=use_amp,
):
    logits = model(x)
    assert_finite("logits", logits)
    loss = criterion(logits, y)

assert_finite("loss", loss)

发现异常后,应按数据流逆向定位:

  1. 输入是否已包含 NaN/Inf;
  2. 模型哪个中间层首先出现异常;
  3. loss 是否与输出语义匹配;
  4. unscale 后梯度是否有限;
  5. 学习率是否在异常发生前改变;
  6. 混合精度是否触发溢出;
  7. 当前批次是否包含极端长度、极端值或异常标签。

torch.autograd.set_detect_anomaly(True) 可以帮助定位部分反向传播异常,但会显著降低速度,应仅用于小规模复现,而不是长期生产训练。


PyTorch 训练循环:一个可运行的最小工程骨架

下面示例使用随机生成的二分类数据,展示数据、模型、优化器、梯度累积、混合精度、验证和检查点的完整关系。它不能证明模型在真实任务上的效果,但可以验证训练循环的状态变化。

import random
from pathlib import Path

import numpy as np
import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset


def seed_everything(seed: int = 42):
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    if torch.cuda.is_available():
        torch.cuda.manual_seed_all(seed)


def evaluate(model, loader, device):
    model.eval()
    total_loss = 0.0
    total_correct = 0
    total_count = 0
    criterion = nn.CrossEntropyLoss(reduction="sum")

    with torch.no_grad():
        for x, y in loader:
            x = x.to(device, non_blocking=True)
            y = y.to(device, non_blocking=True)

            logits = model(x)
            total_loss += criterion(logits, y).item()
            total_correct += (logits.argmax(dim=1) == y).sum().item()
            total_count += y.numel()

    return total_loss / total_count, total_correct / total_count


def main():
    seed_everything()

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    use_amp = device.type == "cuda"
    amp_dtype = torch.float16

    # 构造可运行的二分类数据
    n_train, n_val, feature_dim = 4096, 1024, 20
    x = torch.randn(n_train + n_val, feature_dim)
    y = (x[:, :5].sum(dim=1) > 0).long()

    train_ds = TensorDataset(x[:n_train], y[:n_train])
    val_ds = TensorDataset(x[n_train:], y[n_train:])

    train_loader = DataLoader(
        train_ds,
        batch_size=64,
        shuffle=True,
        num_workers=0,
        pin_memory=(device.type == "cuda"),
    )
    val_loader = DataLoader(
        val_ds,
        batch_size=256,
        shuffle=False,
        num_workers=0,
        pin_memory=(device.type == "cuda"),
    )

    model = nn.Sequential(
        nn.Linear(feature_dim, 64),
        nn.GELU(),
        nn.LayerNorm(64),
        nn.Linear(64, 2),
    ).to(device)

    optimizer = torch.optim.AdamW(
        model.parameters(),
        lr=3e-4,
        weight_decay=1e-2,
    )

    epochs = 5
    accumulation_steps = 2
    updates_per_epoch = (
        len(train_loader) + accumulation_steps - 1
    ) // accumulation_steps
    total_updates = epochs * updates_per_epoch

    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
        optimizer,
        T_max=total_updates,
        eta_min=3e-5,
    )

    scaler = torch.amp.GradScaler(
        "cuda",
        enabled=use_amp,
    )

    global_update = 0
    checkpoint_dir = Path("checkpoints")
    checkpoint_dir.mkdir(parents=True, exist_ok=True)

    for epoch in range(epochs):
        model.train()
        optimizer.zero_grad(set_to_none=True)

        running_loss = 0.0
        seen = 0

        for step, (x_batch, y_batch) in enumerate(train_loader):
            x_batch = x_batch.to(device, non_blocking=True)
            y_batch = y_batch.to(device, non_blocking=True)

            with torch.autocast(
                device_type=device.type,
                dtype=amp_dtype,
                enabled=use_amp,
            ):
                logits = model(x_batch)
                raw_loss = nn.functional.cross_entropy(
                    logits,
                    y_batch,
                )
                loss = raw_loss / accumulation_steps

            scaler.scale(loss).backward()

            is_last_batch = step == len(train_loader) - 1
            should_update = (
                (step + 1) % accumulation_steps == 0
                or is_last_batch
            )

            if should_update:
                # 将不同长度的最后一组累积梯度修正为平均梯度
                remainder = (step + 1) % accumulation_steps
                actual_steps = (
                    remainder if is_last_batch and remainder else accumulation_steps
                )

                if actual_steps != accumulation_steps:
                    correction = accumulation_steps / actual_steps
                    for p in model.parameters():
                        if p.grad is not None:
                            p.grad.mul_(correction)

                scaler.unscale_(optimizer)

                grad_norm = torch.nn.utils.clip_grad_norm_(
                    model.parameters(),
                    max_norm=1.0,
                )

                scaler.step(optimizer)
                scaler.update()
                optimizer.zero_grad(set_to_none=True)

                scheduler.step()
                global_update += 1

                if not torch.isfinite(grad_norm):
                    raise FloatingPointError("gradient norm is not finite")

            running_loss += raw_loss.item() * y_batch.size(0)
            seen += y_batch.size(0)

        val_loss, val_acc = evaluate(model, val_loader, device)
        train_loss = running_loss / seen
        lr = optimizer.param_groups[0]["lr"]

        print(
            f"epoch={epoch + 1} "
            f"updates={global_update} "
            f"train_loss={train_loss:.4f} "
            f"val_loss={val_loss:.4f} "
            f"val_acc={val_acc:.4f} "
            f"lr={lr:.6g}"
        )

        torch.save(
            {
                "model": model.state_dict(),
                "optimizer": optimizer.state_dict(),
                "scheduler": scheduler.state_dict(),
                "scaler": scaler.state_dict(),
                "epoch": epoch,
                "global_update": global_update,
            },
            checkpoint_dir / "latest.pt",
        )


if __name__ == "__main__":
    main()

代码中的关键状态

model.state_dict() 只保存模型参数和模型缓冲区,例如某些归一化层的运行统计量;它不保存 AdamW 的一阶、二阶矩,也不保存学习率调度器位置。

要真正恢复训练,还需要:

  • optimizer.state_dict():恢复动量和自适应统计量;
  • scheduler.state_dict():恢复学习率时间表;
  • scaler.state_dict():恢复混合精度的 loss scale;
  • 当前 epoch 和更新步数;
  • 随机数状态、数据采样器状态,以及分布式训练时的 rank 相关状态。

如果只加载模型参数,却继续使用旧训练计划,恢复后的过程并不等价于中断前继续训练。

PyTorch API 会随版本演进。上例使用当前 PyTorch 2.x 文档中的 torch.amp.GradScaler("cuda", ...)torch.autocast(...) 风格;旧版本可能使用 torch.cuda.amp.autocasttorch.cuda.amp.GradScaler。实际运行时应以安装版本对应的 PyTorch 文档为准,不要混用不同版本示例。


训练、验证和推理的生命周期

训练阶段

训练阶段通常需要:

model.train()
optimizer.zero_grad(set_to_none=True)
# forward -> loss -> backward -> optimizer.step()

model.train() 会让 Dropout 启用随机丢弃,并让 BatchNorm 使用当前批次统计量、更新其运行统计量。

验证阶段

验证不应更新参数或梯度:

model.eval()

with torch.no_grad():
    logits = model(x)

torch.no_grad() 减少 autograd 记录;model.eval() 改变模块行为。二者职责不同,不能只调用其中一个。

验证集只能用于模型选择、超参数决策和早停等有限用途。如果反复根据验证集调参,验证集也会逐渐参与训练决策,最终评测应使用未参与决策的测试集或线上回放集。

检查点与恢复

检查点不仅是“保存一个 .pt 文件”,还涉及:

  • 文件写入是否原子化,避免进程中断产生损坏文件;
  • 是否包含模型结构版本、数据版本和代码提交版本;
  • 是否记录优化器、调度器和 scaler 状态;
  • 是否限制访问权限,因为检查点可能包含训练数据记忆或敏感信息;
  • 是否定期验证可以成功加载和继续训练。

生产环境应将检查点目录与密钥、原始数据目录分离,并使用最小权限的对象存储凭据。训练进程不应拥有删除全部历史检查点的权限,否则单次错误操作可能破坏恢复能力。


学习率、批次与优化器的联动

这三个变量不能独立调参。

大批次不等于简单放大学习率

一种常见经验是批次扩大 kk 倍时,学习率也扩大 kk 倍,即线性缩放规则。它在某些 SGD 场景有理论和经验支持,但并非普遍定律:

  • Adam、AdamW 的自适应统计改变了缩放关系;
  • 梯度累积得到的是多个梯度的平均,而不是总和;
  • 大批次降低了梯度噪声;
  • warmup 长度可能需要同步变化;
  • BatchNorm 的统计行为可能变化;
  • 训练总更新次数会减少。

因此批次变化后,应至少重新检查初始 loss、梯度范数、训练曲线和验证指标,而不是机械地乘以一个比例。

参数更新次数必须显式计算

若数据集有 NN 个样本,有效批次为 BB,训练 EE 个 epoch,则大致更新次数为:

T=ENBT=E\left\lceil\frac{N}{B}\right\rceil

学习率调度器的 T_max、warmup 步数和日志中的 global_step 应围绕 TT 定义。若把 epoch、batch、micro-batch 和 optimizer update 混为一谈,调度和检查点恢复都容易出错。


常见失败模式与诊断

loss 从第一步开始变成 NaN

优先检查:

  1. 输入或标签是否包含非有限值;
  2. 学习率是否过大;
  3. logits 是否出现 inf
  4. 是否错误地对 logits 进行了额外指数运算;
  5. float16 是否发生溢出;
  6. 自定义 loss 是否含有 log(0)、除零或非法开方;
  7. 是否错误处理了 padding 和 mask。

不要先盲目增加 loss scale。若 NaN 出现在前向阶段,GradScaler 无法修复。

loss 不变,梯度几乎为零

可能原因包括:

  • 参数被冻结或没有加入 optimizer;
  • 计算图被错误地 detach()
  • 学习率过小;
  • 激活进入饱和区;
  • 标签或损失实现错误;
  • 混合精度下梯度下溢;
  • 数据本身没有有效信号。

可以打印少量参数的梯度范数:

for name, p in model.named_parameters():
    if p.grad is not None:
        print(name, p.grad.norm().item())
        break

同时确认参数更新前后确实发生变化:

before = model[0].weight.detach().clone()
optimizer.step()
after = model[0].weight.detach()

print((after - before).norm().item())

训练集准确率很高,验证集很差

这通常是过拟合、数据分布差异或数据泄漏边界错误,而不是优化器必然失效。应区分:

  • 训练 loss 是否继续下降;
  • 验证 loss 何时开始上升;
  • 训练和验证预处理是否一致;
  • 样本是否按用户、时间或实体正确拆分;
  • 是否因为重复样本导致评测失真;
  • 任务指标是否与训练 loss 一致。

权重衰减、数据增强、早停和更强的验证设计可能有帮助,但必须先确认评测数据没有被训练流程间接使用。

混合精度比 float32 更差

混合精度改变了数值路径,可能造成:

  • 少数算子溢出;
  • 自定义算子未适配 autocast;
  • 梯度下溢;
  • 归约误差累积;
  • 不同硬件内核产生不同结果。

诊断方式是先用小数据运行 float32 基线,再逐步启用 autocast、GradScaler 和更大的批次。不要同时改变模型、数据、学习率和精度,否则无法判断故障来源。


随机性、可复现性和分布式训练

设置随机种子只能减少一部分差异,不能保证跨设备、跨版本和跨硬件完全一致。差异来源包括:

  • 数据加载器 worker 的随机状态;
  • CUDA 算子选择;
  • 非确定性并行归约;
  • 分布式 AllReduce 的浮点加法顺序;
  • 混合精度内核;
  • 不同 PyTorch、CUDA、驱动和硬件版本。

分布式数据并行中,每张设备通常计算本地梯度,再通过 AllReduce 聚合。若每张设备的本地 loss 和梯度归一化方式不同,global batch 的梯度就可能不符合预期。尤其在最后不完整批次、变长序列和 token masking 场景,应明确聚合的是“样本平均”还是“有效 token 总和”。

确定性设置可以用于调试和回归测试,但可能降低性能,且无法覆盖所有算子。生产系统更现实的目标通常是记录环境、固定数据版本、保存完整状态,并用指标范围而不是单个浮点结果判断回归。


模型、数据、评测、权限和成本应作为一个系统

训练配置的选择会直接改变成本和风险:

  • 更大的 batch 可能提高硬件利用率,但增加单次失败损失;
  • 混合精度可能降低显存压力,但增加数值诊断复杂度;
  • 更长 warmup 会增加训练步数预算;
  • 保存更多检查点提升恢复能力,但增加存储成本;
  • 更频繁验证提高故障发现速度,但减少训练吞吐;
  • 共享数据集和模型工件需要明确读写权限和审计记录。

一次可靠的训练运行至少应记录:

  • 模型结构和初始化方式;
  • 数据集版本、切分规则和预处理配置;
  • optimizer、学习率、scheduler、batch、累积步数;
  • 精度模式、硬件和软件版本;
  • train/validation/test 指标;
  • 梯度范数、loss scale、跳过更新次数;
  • 检查点位置和恢复记录;
  • 运行身份、权限范围和资源消耗。

当训练失败时,这些记录决定了问题能否复现;当模型效果变好时,它们决定了结果能否被验证和复用。


一套可操作的排查顺序

面对“训练不收敛”时,建议按因果链排查,而不是同时修改多个超参数:

  1. 用极小数据集验证模型能否过拟合少量样本;
  2. 用 float32 跑通单个 batch 的前向、反向和更新;
  3. 检查输入、标签、logits、loss 和梯度是否有限;
  4. 确认 train()eval()no_grad() 的生命周期正确;
  5. 打印真实的 optimizer update 次数和学习率;
  6. 检查梯度累积是否按实际 micro-batch 数归一化;
  7. 再启用混合精度;
  8. 最后扩大 batch、启用分布式和复杂调度器;
  9. 对每次改变保存配置、指标和检查点,确保可以回退。

这个顺序的核心是先验证计算图和数据语义,再验证优化过程,最后才优化吞吐和成本。学习率、批次、优化器、混合精度和稳定性并不是互相独立的配置项,而是同一条训练状态链上的不同环节。只要能明确每一步的梯度、参数、学习率、精度和检查点状态,训练问题就能从“经验调参”转化为可观测、可复现的工程问题。


系列导航与关联阅读

官方资料

本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。