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 实现。若自定义:
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. 最小可靠训练循环
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 是两套开关。
练习
- 观察全零 hidden weights 导致的 gradient 对称性。
- 比较 Xavier 与 Kaiming 下每层 activation/gradient variance。
- 用相同预算公平比较 SGD+Momentum 与 AdamW。
- 故意忘记
eval(),观察含 Dropout/BatchNorm 的 validation 波动。 - 保存并恢复 optimizer/scheduler/RNG,检查下一步 update 是否可复现。
小结
稳定训练来自一组相互匹配的选择:初始化维持信号尺度,optimizer 与 learning-rate schedule 决定更新,batch/normalization 改变统计,regularization 约束泛化。训练循环必须显式管理 mode、gradient、validation 和 checkpoint。
下一章把这些基础放进结构化深度模型:CNN 如何编码局部性,RNN 如何维护状态,attention/Transformer 如何让位置之间按内容交换信息。