第 03 章 · 基础
神经网络怎样学会:损失、梯度与优化
用最少数学建立前向计算、交叉熵、反向传播和参数更新的可靠直觉。
学习准备与本章目标
- 开始前
- 向量与矩阵的基本概念
- 学完后
- 说清完整训练循环 · 理解梯度从何而来 · 知道 loss 异常时先检查什么
本章路线图
这一章回答“模型参数为什么会变好”。完整训练链路是:前向计算得到预测,损失函数量化错误,反向传播把错误分配给参数,优化器根据梯度更新参数。
学习时始终跟踪四类量:张量形状、数值范围、是否需要梯度、在一个 optimizer step 内何时发生变化。
一次训练步发生了什么
模型先根据输入得到每个位置对全词表的分数 logits。把它们与真实的下一个 Token 比较,得到交叉熵损失;随后反向传播计算“每个参数改变一点会让损失怎样变化”,优化器再按这些梯度更新参数。
流程可以记成:取 batch → 前向 → 计算 loss → 梯度清零 → 反向传播 → 必要时裁剪梯度 → 优化器更新 → 学习率调度。训练就是把这个过程重复很多次。
交叉熵为什么适合语言模型
Softmax 把 logits 变成概率分布,交叉熵关注真实 Token 被分到的概率。真实 Token 概率越低,惩罚越大。训练时通常一次并行预测序列中的所有下一个 Token,而生成时才逐个产生。
Padding、提示词和不希望参与监督的部分会通过 mask 排除。若标签错位一个位置,模型就会学习错误目标;这是语言模型代码中非常常见、也很隐蔽的问题。
梯度、学习率与 AdamW
梯度指出局部最陡的变化方向,学习率控制一步走多远。太大可能震荡或出现 NaN,太小则学习缓慢。AdamW 会保存梯度的一阶和二阶统计量,并把权重衰减与梯度更新解耦,是大模型训练的常见选择。
训练初期常使用 warmup,先把学习率从很小值升高,避免随机初始化或新任务分布下的剧烈更新;之后再按余弦或线性策略衰减。
训练不稳定时怎样排查
不要只盯总 loss。还要观察学习率、梯度范数、有效 Token 数、不同数据源的 loss、吞吐和显存峰值。出现 NaN 时先寻找“第一个坏点”:脏数据、非法标签、除零、低精度溢出或通信异常。
梯度裁剪只能限制更新幅度,不能修复错误数据;降低学习率也不能代替定位。可靠训练依赖可复现配置、数据版本、随机种子、checkpoint 和最小复现实验。
计算图与链式法则
神经网络的前向过程是一张计算图。若 y = f(x, w),损失为 L(y),参数梯度来自链式法则:
深层网络只是把这条链延长。自动微分会在前向时记录需要的操作和中间量,反向时按相反顺序执行局部导数。detach() 会切断计算图,原地修改可能破坏反向所需的保存值,no_grad() 则用于明确不构建梯度图的推理区域。
梯度回答的是“当前位置附近,参数变化一点会怎样影响 loss”,不是从起点直达最优解的方向。因此训练需要许多小步,还会受 batch 抽样噪声与损失曲面形状影响。
交叉熵、Logits 与困惑度
模型输出 logits,而不是先手工计算概率。对真实类别 ,单位置交叉熵为:
框架通常把 log_softmax 与负对数似然合并计算,以避免指数溢出。语言模型将标签向右错一位,并只对有效 Token 求平均。困惑度是平均交叉熵的指数,但只有在相同 Tokenizer、数据切分和归一化方式下才可比较。
shift_logits = logits[:, :-1].contiguous()
shift_labels = input_ids[:, 1:].contiguous()
loss = torch.nn.functional.cross_entropy(
shift_logits.view(-1, shift_logits.size(-1)),
shift_labels.view(-1),
ignore_index=pad_token_id,
)
梯度累积、裁剪与有效 Batch
显存放不下大 batch 时,可把多个 micro-batch 的梯度累积后再更新。若每个 micro loss 除以累积步数,并保持 token 归一化一致,结果近似一个大 batch;但 dropout、动态长度、混合精度溢出和 batch 统计会带来差异。
梯度裁剪通常在反向累积完成、优化器更新之前执行。按范数裁剪会整体缩放梯度,尽量保留方向;逐值裁剪则改变各元素相对关系。分布式环境还要确认范数是在聚合前还是聚合后计算。
有效训练规模最好用 每次参数更新处理的有效 Token 数 描述:
tokens/update = GPU 数 × micro batch × 序列有效 Token × 累积步数
AdamW 与学习率调度
SGD 直接沿梯度方向更新;Adam 为每个参数维护梯度的一阶矩和二阶矩,自适应调整步长。AdamW 把权重衰减从梯度更新中解耦,避免 L2 项被自适应缩放。
常见参数组会对矩阵权重使用衰减,对 bias 与归一化缩放参数关闭衰减。是否对 Embedding 衰减取决于训练配方,不应机械套用。
学习率决定每一步尺度。warmup 保护训练初期尚不稳定的激活与优化器统计;余弦或线性衰减让后期更新更细。使用梯度累积时,scheduler 一般跟随 optimizer step,而不是每个 micro-step。
低精度训练与数值稳定
FP16 指数范围较小,梯度可能上溢或下溢,常配合动态 loss scaling;BF16 保留与 FP32 相同的指数位数,训练通常更稳定,但有效尾数更少。混合精度不是把所有张量都强制转低精度:矩阵乘可以低精度,归一化、归约、loss 和优化器状态常保留更高精度。
看到 NaN 时应定位第一个非有限值,而不是只看最终 loss。常见来源包括全遮罩后的 softmax、除零、非法标签、异常大 logits、错误恢复的优化器状态和某个 rank 的坏样本。
一个完整而正确的训练循环
model.train()
optimizer.zero_grad(set_to_none=True)
for micro_step, batch in enumerate(loader):
with torch.autocast("cuda", dtype=torch.bfloat16):
output = model(**batch)
loss = output.loss / accumulation_steps
loss.backward()
if (micro_step + 1) % accumulation_steps == 0:
grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
if not torch.isfinite(grad_norm):
raise RuntimeError("non-finite gradient")
optimizer.step()
scheduler.step()
optimizer.zero_grad(set_to_none=True)
真实项目还要处理验证、日志、断点、随机状态、分布式同步和异常恢复。checkpoint 至少包含模型、优化器、scheduler、step、随机数状态,以及混合精度 scaler(若使用)。
训练曲线怎样阅读
训练 loss 与验证 loss 都高,可能是容量不足、学习率错误或数据表达不足;训练继续下降而验证恶化,可能过拟合;验证异常好则要警惕重复数据和评测污染。
同时记录 tokens/s、梯度范数、学习率、显存峰值、数据来源占比和各任务指标。吞吐下降可能来自数据加载、序列变长、通信等待或频繁保存,不一定是模型算子变慢。
系统排错清单
- 固定 seed 和样本 ID,找到第一个异常 step。
- 检查输入 ID、标签范围、mask 与有效 Token 数。
- 在关键层记录激活和梯度的最小值、最大值、均值与有限性。
- 暂时使用 FP32 和单卡最小复现,缩小混合精度或通信变量。
- 核对学习率、累积步数、裁剪时机和恢复状态。
- 修复后保留触发样本与监控,形成回归测试。
章末检查
- 画出
logits → loss → backward → optimizer.step的依赖关系。 - 解释为什么交叉熵应直接接收 logits。
- 说明梯度累积在哪些条件下近似大 batch。
- 列出一次精确断点续训必须保存的状态。
- 给出 Loss NaN 的排查顺序,而不是只说“调小学习率”。
延伸资料
本章进阶内容
原理专题与代码实践
完成主教材后,按顺序阅读原理专题、完成代码实践,并用掌握标准复查本章内容。
- 01
建立损失、梯度与优化的因果链
- 02
推导交叉熵、归一化和梯度稳定性
- 03
亲手写自动求导与完整训练循环
- 04
用过拟合一个小批次验证训练管线
本章术语
六个关键词
- 交叉熵
- 衡量真实 Token 与模型预测分布差异的常用损失。
- 梯度
- 损失对参数的局部变化率。
- 反向传播
- 沿计算图反向应用链式法则求梯度。
- AdamW
- 将权重衰减与自适应梯度更新解耦的优化器。
- Warmup
- 训练初期逐步升高学习率的稳定化阶段。
- 梯度累积
- 多次反向后再更新参数,以模拟更大有效 batch。
章节记录
