跳到内容

9.3 初始化、优化与训练循环:Loss 下降只是训练系统的第一条生命体征

计算图已经能给出梯度,模型工坊却没有因此自动稳定:学习率稍大就出现 NaN,换个 seed 又得到不同结果,validation 还会在 train loss 继续下降时恶化。

训练是一个系统:初始化决定初始信号与梯度尺度,optimizer 把梯度变成更新,data loader 决定噪声与吞吐,mode 控制 dropout/normalization,validation 决定何时保存和停止。

本课目标

  • 根据 activation/fan-in/fan-out 选择初始化;
  • 理解 SGD、Momentum、Adam 与 AdamW 的更新差异;
  • 正确处理 learning rate、batch、normalization 与 clipping;
  • 写出包含 train/eval mode 的最小可靠 PyTorch loop;
  • 记录 checkpoint、随机性和诊断信号。

1. 全零权重为什么失败

同一层 hidden units 若以完全相同权重开始,并接收相同 gradient,它们会一直保持相同,无法分化成不同 features。这是 symmetry 问题。

Bias 初始化为 0 通常没这个问题,因为随机 weights 已打破 hidden-unit symmetry;“任何参数都不能为零”也不对。

随机值还要控制 variance。层层传播若 activation/gradient variance 持续放大或缩小,会造成饱和、爆炸或信号消失。

2. Xavier 与 Kaiming 的假设

Xavier/Glorot initialization 根据 fan-in/fan-out 调整 variance,常与近似对称 activation(如 tanh)搭配。Kaiming/He initialization 考虑 ReLU 类 activation 只保留部分输入,常按 fan-in 保持 forward variance。

这些是基于独立、零均值、特定 activation 分布的近似推导,不是架构无关定理。Residual branch、attention、gating、normalization 和极深网络可能使用专门 scale。

PyTorch module 有各自默认初始化;不要一边假设“已经 He 初始化”,一边从未检查当前版本和 module 实现。若自定义:

python
from torch import nn

def initialize(module):
    if isinstance(module, nn.Linear):
        nn.init.kaiming_normal_(module.weight, nonlinearity="relu")
        if module.bias is not None:
            nn.init.zeros_(module.bias)

model.apply(initialize)

如果最后一层/activation 不同,不应无差别套同一初始化。

3. SGD 与 Momentum

Mini-batch gradient $g_t$ 对 full gradient 是带噪估计(在相应抽样条件下)。基本 SGD:

$$ \theta_{t+1}=\theta_t-\eta g_t. $$

Momentum 累积指数加权方向(具体符号/实现约定不同):

$$ v_t=\mu v_{t-1}+g_t,qquad \theta_{t+1}=\theta_t-\eta v_t. $$

它可在持续方向加速、在来回震荡方向平滑,但不保证“避开鞍点”或得到更好泛化。Learning rate、momentum、batch 和 schedule 联合决定轨迹。

4. Adam 与 AdamW

Adam 维护 gradient 一阶矩和平方二阶原始矩的指数估计,做 bias correction 后按坐标缩放 update。它常能快速得到可用训练结果,尤其面对稀疏/不同尺度梯度,但不是所有任务的通用最佳 optimizer。

Adam 中把 L2 penalty 混入 gradient 与真正 decoupled weight decay 不完全等价。AdamW 将 weight decay 与 adaptive gradient update 解耦。还要决定哪些参数 decay:bias、normalization scale 等常被分到 no-decay group,但这也是架构/实验选择。

比较 optimizer 时应给相近调参预算、schedule 与 compute,而不是用 SGD 默认学习率对比精调 AdamW。

5. Learning Rate 是首要稳定参数之一

过大:loss 振荡/发散、activation/gradient 出现 Inf/NaN;过小:有限预算内几乎不动,也可能停在差区域。

常见 schedule:warmup、step/exponential decay、cosine decay、plateau-driven。Warmup 可在大 batch、adaptive optimizer 或不稳定初期减小 update,但并非所有小网络都需要。

“Adam 用 0.001”只是某些配置的起点。有效 learning rate 还受 loss reduction、batch、gradient accumulation、parameterization 和 precision 影响。

6. Batch Size 改变统计与系统行为

小 batch:gradient noise 较大、更新频繁、吞吐可能低。大 batch:硬件利用率可能更好、每 epoch updates 少、内存高,并可能需要调整 learning rate/schedule。

“大 batch 一定泛化差”不是定律。比较时要控制 total examples、updates、schedule 和 compute。Gradient accumulation 模拟较大 effective batch 的梯度平均,但不能完全复制 BatchNorm statistics、optimizer step 频率和随机增强序列。

7. Normalization 不只是防梯度爆炸

BatchNorm

训练时使用 mini-batch statistics 并更新 running estimates;eval 时使用保存的 running statistics。小/非 iid batch、分布式 shard 和 train/eval 切换都会影响结果。

LayerNorm / RMSNorm

沿单样本的特征维归一化,不依赖 batch statistics,常用于序列模型。具体 normalized shape 与 dtype 精度仍要检查。

Norm 放在 activation 前后、residual branch 前后没有一个适用于所有架构的答案。Pre-norm/Post-norm 会改变优化路径;遵循目标架构并用消融验证,不要背“Linear–BN–ReLU”万能顺序。

8. Regularization 作用不同

  • weight decay:收缩 parameters,效果与 parameterization/optimizer 相关;
  • dropout:训练时随机置零并缩放,eval 时关闭;
  • data augmentation:注入任务不变性假设;
  • label smoothing:改变 target distribution 与 calibration;
  • early stopping:用 validation 选择训练轮数;
  • stochastic depth/mixup 等:各自改变目标/架构。

它们不是可随意叠加的“防过拟合套餐”。先确认 failure mode,并把所有强度在 validation 内选择。

9. Gradient Clipping 是护栏,不是修理工

Global norm clipping:

$$ g\leftarrow g\cdot\min\left(1,\frac{c}{\lVert g\rVert}\right). $$

它限制单步 gradient norm,常用于序列/不稳定训练。阈值太小会持续改变优化方向/尺度;持续触发提示应检查 loss scaling、data outlier、初始化、learning rate 和数值溢出。

记录 clip 前 norm 与触发比例。只在 optimizer.step() 前 clip;使用 mixed precision 时需按当前 AMP API 先 unscale 再 clip。

10. 最小可靠训练循环

python
import torch
from torch import nn

def train_one_epoch(model, loader, optimizer, loss_fn, device):
    model.train()
    total_loss = 0.0
    total_examples = 0

    for features, target in loader:
        features = features.to(device)
        target = target.to(device)

        optimizer.zero_grad(set_to_none=True)
        logits = model(features)
        loss = loss_fn(logits, target)

        if not torch.isfinite(loss):
            raise FloatingPointError(f"non-finite loss: {loss.item()}")

        loss.backward()
        nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0)
        optimizer.step()

        batch = target.shape[0]
        total_loss += loss.item() * batch
        total_examples += batch

    return total_loss / total_examples

def evaluate(model, loader, loss_fn, device):
    model.eval()
    total_loss = 0.0
    total_examples = 0
    all_logits, all_targets = [], []

    with torch.inference_mode():
        for features, target in loader:
            features = features.to(device)
            target = target.to(device)
            logits = model(features)
            loss = loss_fn(logits, target)

            batch = target.shape[0]
            total_loss += loss.item() * batch
            total_examples += batch
            all_logits.append(logits.cpu())
            all_targets.append(target.cpu())

    return {
        "loss": total_loss / total_examples,
        "logits": torch.cat(all_logits),
        "targets": torch.cat(all_targets),
    }

这里假设 loss 默认是每 batch mean;若使用 reduction="sum"、sample weights、token-level padding 或最后不完整 batch,累计方式要相应改变。Metric 应在聚合 predictions 后按正确总体计算,不能盲目平均 batch AUC/F1。

11. Train/Eval Mode 与 Validation

model.eval() 不等于关闭 gradient;torch.inference_mode()/no_grad() 也不等于切换 dropout/BatchNorm。验证通常两者都要。

Validation 用于:

  • early stopping/checkpoint 选择;
  • schedule/超参数比较;
  • calibration/threshold(需合适数据边界);
  • failure slice 与 learning curve。

Test 不进入每 epoch loop。若每轮都查看 test accuracy,它已经成为 validation。

12. Checkpoint 必须能恢复训练状态

只保存 weights 适合 inference,不足以无缝 resume。训练 checkpoint 常包括:

  • model/optimizer/scheduler state;
  • AMP scaler state(若使用);
  • epoch/global step/best metric;
  • sampler/data position(需要精确恢复时);
  • RNG states 与 seed policy;
  • config、代码/依赖/数据版本;
  • normalization/tokenizer/label mapping。

保存“last”与“best validation”两个 checkpoint,明确 best 的 metric、方向和 tie rule。原子写入并验证能 load/推理,避免训练数小时后才发现 checkpoint 损坏。

13. Reproducibility 有边界

设置 Python/NumPy/PyTorch seeds 只能控制部分随机源。Hardware、并行执行、nondeterministic kernels 和库版本都可能改变结果;CPU/GPU 或版本之间不保证 bitwise 相同。

Deterministic algorithms 可能降低性能或对无确定实现的 op 报错。开发/回归阶段可提高确定性,最终 benchmark 要记录模式。报告多个 seeds/置信范围通常比只展示最好 seed 更诚实。

14. 训练前后的三组测试

训练前

  • shape/dtype/device;
  • label range 与 mask;
  • 参数数、initial logits/loss;
  • 数据重复/泄漏;
  • 单 batch forward/backward finite。

极小数据 Overfit

尝试把 1–2 个 batch 的 train loss 降到很低。失败通常提示实现、容量、loss/label 或 optimizer 问题;成功不证明泛化。

正式训练

记录 train/validation loss、任务指标、gradient/update norm、throughput、memory、learning rate 和 checkpoint。看到异常先缩小复现,不要直接堆 optimizer 技巧。

常见误区

  • 所有参数都不能初始化为 0:关键是打破同层 weights 的对称性。
  • Adam 是任何任务的默认最优:它是候选,需要公平调参与验证。
  • BatchNorm 固定放在 activation 前:顺序属于架构设计。
  • Batch 越大收敛越快/泛化越差:要控制 update、schedule 与 compute。
  • model.eval() 已关闭 autograd:mode 与 gradient recording 是两套开关。

练习

  1. 观察全零 hidden weights 导致的 gradient 对称性。
  2. 比较 Xavier 与 Kaiming 下每层 activation/gradient variance。
  3. 用相同预算公平比较 SGD+Momentum 与 AdamW。
  4. 故意忘记 eval(),观察含 Dropout/BatchNorm 的 validation 波动。
  5. 保存并恢复 optimizer/scheduler/RNG,检查下一步 update 是否可复现。

小结

稳定训练来自一组相互匹配的选择:初始化维持信号尺度,optimizer 与 learning-rate schedule 决定更新,batch/normalization 改变统计,regularization 约束泛化。训练循环必须显式管理 mode、gradient、validation 和 checkpoint。

下一章把这些基础放进结构化深度模型:CNN 如何编码局部性,RNN 如何维护状态,attention/Transformer 如何让位置之间按内容交换信息。

Built with VitePress | Software Systems Atlas