云镜收藏

稍后阅读

清单保存在当前浏览器,方便下次回来继续阅读。

清单还是空的

在笔记卡片或正文页点击“加入稍后阅读”即可收藏。

展开课程与本章目录
12 章学习路线01零基础准备:Python、NumPy 与正确安装02起步:环境、设备与第一枚 Tensor03张量基本功:索引、广播与线性代数04自动求导:从计算图到反向传播05神经网络工程:Module、损失与训练循环06数据管线:Dataset、DataLoader 与批处理07亲手搭 Transformer:从 Embedding 到注意力08稳定训练:初始化、AMP、裁剪与排错09性能与显存:Profiler、compile 与检查点10分布式训练:从 DDP 到 FSDP 与 TP11大模型生态:Transformers、PEFT、torchao 与 TorchTitan12方向实战:训练并优化一个迷你语言模型本章课程01nn.Module 不只是一个 Python 类02forward、call 与训练/评估模式03选择损失函数并检查输入约定04优化器、参数组与权重衰减05写对一个最小训练循环06学习率调度与 warmup07checkpoint:保存的不只是模型权重

第 05 章 · 基础

神经网络工程:Module、损失与训练循环

从 `nn.Module` 的参数注册机制出发,亲手搭建可训练、可验证、可保存和可恢复的完整训练器。

7 节课程校订于 2026年8月31日
nn.Modulelossoptimizercheckpoint

预计用时140 分钟

前置知识Autograd · 基础面向对象语法

完成标准写完整训练/验证循环 · 正确保存与恢复状态 · 配置 AdamW 参数组

本章进度0 / 7

nn.Module 不只是一个 Python 类

本节目标

本节围绕“nn.Module 不只是一个 Python 类”展开。运行示例后检查参数注册、训练模式、loss、梯度和优化器更新,确认训练循环各步顺序。

清晰讲解

nn.Module 会递归注册子模块、参数和 buffer,因此 .to(device)state_dict() 与优化器才能自动发现它们。可训练权重用 nn.Parameter,非训练状态如位置索引可用 register_buffer

核对“nn.Module 不只是一个 Python 类”中的参数集合、train/eval 模式、loss 输入约定和更新前后数值。

代码示例

import torch
from torch import nn

class TinyMLP(nn.Module):
    def __init__(self, dim=16):
        super().__init__()
        self.net = nn.Sequential(nn.Linear(dim, 4*dim), nn.GELU(), nn.Linear(4*dim, dim))
        self.register_buffer('scale', torch.tensor(dim**-0.5))
    def forward(self, x):
        return self.net(x) * self.scale

运行结果与观察

打印参数时能看到 MLP 权重,打印 buffer 时能看到 scale;两者会随 .to(device) 移动,但只有参数交给优化器。

常见错误

  • 把层放进普通列表导致无法注册;在 forward 临时创建带参数的层。
  • 排查“nn.Module 不只是一个 Python 类”时只看 loss 是否下降,没有检查模式、参数更新、梯度清零和优化器状态。

与大模型方向的连接

Transformer block 本质也是 Module 树;参数命名与分组决定权重衰减、LoRA 注入和分布式切分。

动手练习

打印模型的 named_parameters()named_buffers(),解释两者区别。

查看参考答案与验收点

验收:参数确实更新,训练与评估模式行为正确,并能说明“nn.Module 不只是一个 Python 类”在训练循环中的位置。

官方资料


forward、call 与训练/评估模式

本节目标

本节围绕“forward、call 与训练/评估模式”展开。运行示例后检查参数注册、训练模式、loss、梯度和优化器更新,确认训练循环各步顺序。

清晰讲解

平时调用 model(x) 而不是直接调用 forward,因为 Module 的调用流程还处理 hooks 和包装逻辑。train()eval() 控制 Dropout、BatchNorm 等层的行为,但不会自动开启或关闭梯度。

核对“forward、call 与训练/评估模式”中的参数集合、train/eval 模式、loss 输入约定和更新前后数值。

代码示例

import torch
from torch import nn
m = nn.Dropout(p=0.5)
x = torch.ones(8)
m.train(); print(m(x))
m.eval();  print(m(x))
with torch.no_grad(): print(m(x))

运行结果与观察

训练模式的 dropout 输出含随机 0,评估模式保持全 1;no_grad 不会替代 eval()

常见错误

  • 认为 eval() 等于禁止梯度;直接调用 forward 绕过框架机制。
  • 排查“forward、call 与训练/评估模式”时只看 loss 是否下降,没有检查模式、参数更新、梯度清零和优化器状态。

与大模型方向的连接

语言模型中的 dropout 在验证与生成时必须关闭;eval 和 no_grad 分别解决“层行为”和“求导开销”。

动手练习

写出一份正确的验证阶段模板,并解释每行状态切换。

查看参考答案与验收点

用一个极小 batch 过拟合,并记录 loss、梯度范数和参数更新;断点恢复后比较下一步结果,而不只检查文件存在。

官方资料


选择损失函数并检查输入约定

本节目标

本节围绕“选择损失函数并检查输入约定”展开。运行示例后检查参数注册、训练模式、loss、梯度和优化器更新,确认训练循环各步顺序。

清晰讲解

损失函数不仅是公式,还有严格的形状与数值约定。CrossEntropyLoss 接受未经 softmax 的 logits 与整数类别;语言模型通常把 [B,T,V] 展平为 [B*T,V],标签展平为 [B*T]

核对“选择损失函数并检查输入约定”中的参数集合、train/eval 模式、loss 输入约定和更新前后数值。

代码示例

import torch
from torch import nn
B, T, V = 2, 4, 10
logits = torch.randn(B, T, V)
targets = torch.randint(0, V, (B, T))
loss = nn.functional.cross_entropy(logits.view(-1, V), targets.view(-1))
print(loss.item())

运行结果与观察

交叉熵应是有限标量;把 logits 或 targets 的 shape/dtype 改错后应得到明确错误,而不是静默训练。

常见错误

  • 先 softmax 再交叉熵;标签用了浮点型;把词表维放错。
  • 排查“选择损失函数并检查输入约定”时只看 loss 是否下降,没有检查模式、参数更新、梯度清零和优化器状态。

与大模型方向的连接

next-token prediction 就是在每个有效 token 位置做多分类;padding 可通过 ignore_index 排除。

动手练习

加入 padding 标签 -100,验证这些位置不影响 loss。

查看参考答案与验收点

验收:参数确实更新,训练与评估模式行为正确,并能说明“选择损失函数并检查输入约定”在训练循环中的位置。

官方资料


优化器、参数组与权重衰减

本节目标

本节围绕“优化器、参数组与权重衰减”展开。运行示例后检查参数注册、训练模式、loss、梯度和优化器更新,确认训练循环各步顺序。

清晰讲解

优化器读取参数的 .grad 更新数值。AdamW 将权重衰减与梯度更新解耦,是 Transformer 常用选择;bias 与归一化参数通常不做衰减,可按名字或维度分组。

核对“优化器、参数组与权重衰减”中的参数集合、train/eval 模式、loss 输入约定和更新前后数值。

代码示例

import torch
from torch import nn
model = nn.Sequential(nn.Linear(8, 16), nn.LayerNorm(16), nn.Linear(16, 2))
decay, no_decay = [], []
for name, p in model.named_parameters():
    (decay if p.ndim >= 2 else no_decay).append(p)
opt = torch.optim.AdamW([
    {'params': decay, 'weight_decay': 0.1},
    {'params': no_decay, 'weight_decay': 0.0},
], lr=3e-4)

运行结果与观察

两个参数组不重不漏:二维权重进入 decay,bias 与归一化参数进入 no_decay;元素总数应等于全部可训练参数。

常见错误

  • 把被冻结参数也交给优化器;参数在多个 group 重复;将 L2 正则与 AdamW 衰减混为一谈。
  • 排查“优化器、参数组与权重衰减”时只看 loss 是否下降,没有检查模式、参数更新、梯度清零和优化器状态。

与大模型方向的连接

大模型训练脚本常按参数性质分组;先验证参数不重不漏,再讨论复杂策略。

动手练习

统计两个参数组的元素总数,并断言等于所有可训练参数总数。

查看参考答案与验收点

用一个极小 batch 过拟合,并记录 loss、梯度范数和参数更新;断点恢复后比较下一步结果,而不只检查文件存在。

官方资料


写对一个最小训练循环

本节目标

本节围绕“写对一个最小训练循环”展开。运行示例后检查参数注册、训练模式、loss、梯度和优化器更新,确认训练循环各步顺序。

清晰讲解

标准顺序是取批次、前向、算损失、清梯度、反向、更新。顺序本身容易背,难点是每一步的设备、模式、形状和梯度状态都一致。先用小数据过拟合,是验证管线正确性的最快方法。

核对“写对一个最小训练循环”中的参数集合、train/eval 模式、loss 输入约定和更新前后数值。

代码示例

for inputs, targets in loader:
    inputs, targets = inputs.to(device), targets.to(device)
    optimizer.zero_grad(set_to_none=True)
    logits = model(inputs)
    loss = loss_fn(logits, targets)
    loss.backward()
    optimizer.step()

运行结果与观察

至少记录首尾 epoch 的 loss 与 accuracy:在内置可分数据上 loss 应明显下降、准确率应高于随机基线;若二者不动,先检查 zero_grad → backward → step 顺序。

常见错误

  • no_grad 里训练;忘记搬标签;对验证 loss 反向传播。
  • 排查“写对一个最小训练循环”时只看 loss 是否下降,没有检查模式、参数更新、梯度清零和优化器状态。

与大模型方向的连接

预训练循环会增加梯度累积、AMP、裁剪与调度器,但骨架没有改变。复杂化前先确保一个 batch 能让 loss 下降。

动手练习

用随机生成的线性数据训练一个回归模型,确认几十步内 loss 明显下降。

查看参考答案与验收点

用一个极小 batch 过拟合,并记录 loss、梯度范数和参数更新;断点恢复后比较下一步结果,而不只检查文件存在。

官方资料


学习率调度与 warmup

本节目标

本节围绕“学习率调度与 warmup”展开。运行示例后检查参数注册、训练模式、loss、梯度和优化器更新,确认训练循环各步顺序。

清晰讲解

学习率决定每次更新幅度。Transformer 常先 warmup,降低初始化早期不稳定,再按 cosine 或线性策略衰减。调度器应与“优化器更新次数”对齐,而不是盲目与 dataloader 次数对齐。

核对“学习率调度与 warmup”中的参数集合、train/eval 模式、loss 输入约定和更新前后数值。

代码示例

import math
def lr_scale(step, warmup, total):
    if step < warmup:
        return (step + 1) / warmup
    progress = (step - warmup) / max(1, total - warmup)
    return 0.5 * (1 + math.cos(math.pi * progress))

scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lambda s: lr_scale(s, 100, 1000))

运行结果与观察

逐步打印学习率,warmup 阶段应单调升高,随后按既定调度下降;第一步、峰值步和最后一步都应与手算值一致。

常见错误

  • 把 batch 数当 update 数;恢复 checkpoint 后调度器从零重启。
  • 排查“学习率调度与 warmup”时只看 loss 是否下降,没有检查模式、参数更新、梯度清零和优化器状态。

与大模型方向的连接

梯度累积时 4 个 micro-step 才发生一次参数更新,scheduler 通常也应只走一步。

动手练习

画出 1000 步 warmup+cosine 曲线,检查首尾和峰值。

查看参考答案与验收点

验收:参数确实更新,训练与评估模式行为正确,并能说明“学习率调度与 warmup”在训练循环中的位置。

官方资料


checkpoint:保存的不只是模型权重

本节目标

本节围绕“checkpoint:保存的不只是模型权重”展开。运行示例后检查参数注册、训练模式、loss、梯度和优化器更新,确认训练循环各步顺序。

清晰讲解

只做推理时保存 state_dict 即可;要精确续训,还需优化器、调度器、step、随机状态和 AMP scaler。保存普通数据结构比 pickle 整个模型对象更稳健。

核对“checkpoint:保存的不只是模型权重”中的参数集合、train/eval 模式、loss 输入约定和更新前后数值。

代码示例

checkpoint = {
    'model': model.state_dict(),
    'optimizer': optimizer.state_dict(),
    'scheduler': scheduler.state_dict(),
    'step': step,
    'cpu_rng': torch.get_rng_state(),
}
torch.save(checkpoint, 'checkpoint.pt')

state = torch.load('checkpoint.pt', map_location='cpu', weights_only=True)
model.load_state_dict(state['model'])

运行结果与观察

保存后重建模型与优化器,恢复出的下一次 loss 应与不中断训练的对照分支在浮点误差内一致;只恢复权重而不恢复优化器时通常不会一致。

常见错误

  • 保存整个 Module 后依赖代码路径;加载时不检查 missing/unexpected keys;只恢复模型却声称无缝续训。
  • 排查“checkpoint:保存的不只是模型权重”时只看 loss 是否下降,没有检查模式、参数更新、梯度清零和优化器状态。

与大模型方向的连接

大模型 checkpoint 可能分片保存;无论规模多大,核心目标仍是明确恢复哪些状态、由谁重组。

动手练习

保存训练 5 步的状态,重启后继续 5 步,并与不中断的 10 步结果比较。

查看参考答案与验收点

用一个极小 batch 过拟合,并记录 loss、梯度范数和参数更新;断点恢复后比较下一步结果,而不只检查文件存在。

官方资料

章末资料

小测、项目与相关面试题

1. eval() 会自动关闭梯度吗?

不会。eval 控制 Dropout/BatchNorm 行为,关闭梯度要用 no_grad 或 inference_mode。

2. 精确续训只保存模型权重够吗?

不够,还需优化器、调度器、step、随机状态和 AMP scaler 等状态。