跳到主要内容

数据截至 (上游 commit c187ef3271d5)

04 · nn.Module 与 Optimizer:每天写的训练循环在干什么

这一章讲什么: nn.Linearoptim.AdamW 这两个最常见的类内部到底长什么样。读完你会理解:model.parameters() 为什么能找全所有参数、model(x) 为什么等价于 forward 又不只是 forward、以及 AdamW 相对 Adam 到底改了哪一行。


1. 它要解决的小问题

写模型的人只想做两件事:把层当属性赋值self.fc = nn.Linear(...))、把层当函数调y = self.fc(x))。但框架需要随时能枚举全部参数(给优化器、给 state_dict、给 .to(device))。

矛盾在于:普通 Python 对象赋值不会通知任何人。Module 的全部魔法就是劫持 __setattr__,在你赋值的同时完成注册。


2. Module:三本字典 + 属性拦截

2.1 内部状态

Moduletorch/nn/modules/module.py:407)本质上是三本字典的持有者(类型声明见 torch/nn/modules/module.py:456-457:478):

字典装什么
_parameters: dict[str, Parameter | None]可学习参数
_buffers: dict[str, Tensor | None]非参数但属于状态的量(如 BN 的 running_mean)
_modules: dict[str, Module]子模块

2.2 赋值即注册

__setattr__torch/nn/modules/module.py:1980)按值类型分派:

  • 赋的是 Parameter → 从普通属性/其它字典里移除同名项,塞进 _parameters
  • 赋的是 Module → 塞进 _modules
  • 其余 → 正常对象属性。

由此得到几个免费推论:

  • parameters() / named_parameters()torch/nn/modules/module.py:2703)= 递归遍历 _modules,收集每层的 _parameters
  • state_dict()torch/nn/modules/module.py:2203)= 同一套遍历,按 子模块名.参数名 拼 key;
  • .to(device) / .apply(fn)torch/nn/modules/module.py:1040)= 递归遍历后对每个参数/模块动手。

所以「把一个子模块挂到模型上」这件事,在 PyTorch 里没有任何注册 API——赋值本身就是注册。

2.3 model(x) 不是直接调 forward

__call__ 绑到 _wrapped_call_impltorch/nn/modules/module.py:1926),再到 _call_impltorch/nn/modules/module.py:1791)。它的快路径一句话:如果没有任何 hook,直接调 forwardtorch/nn/modules/module.py:1796-1800);有 hook 则按「forward pre-hook → forward → forward hook → backward hook 注册」的顺序绕一圈。

这个设计让「无 hook 时零开销、有 hook 时全功能」两全,也是各种 profiling/剪枝/量化工具能挂在任意层上的原因。


3. Linear:标准层到底算什么

以最常见的 nn.Lineartorch/nn/modules/linear.py:53)为例。

3.1 初始化

reset_parameterstorch/nn/modules/linear.py:117):weight 用 kaiming_uniform_(a=sqrt(5)),注释里点破这等价于 uniform(-1/sqrt(in_features), 1/sqrt(in_features))(并引了 pytorch#57109);bias 再按 fan_in 取 ±1/sqrt(fan_in) 均匀分布。

3.2 前向

forward 只有一行(torch/nn/modules/linear.py:130):F.linear(input, self.weight, self.bias)

F.linear 本身只是 C++ 算子的 Python 包装(torch/nn/functional.py:2382torch._C._nn.linear 挂上文档)。真正的 at::linearaten/src/ATen/native/Linear.cpp:84)做的事:

  1. 把输入的多维 batch 维展平成 2D;
  2. dense 情形直接 at::addmm(bias, input_flattened, weight.t())aten/src/ATen/native/Linear.cpp:67)——即 y = x·Wᵀ + b,一次 GEMM 完成;
  3. 再把输出形状还原。

一个「全连接层」的全部计算,就是一次矩阵乘加偏置。 上面叠的所有 Python 层(Module、functional)都只是到达 addmm 前的登记与转发。


4. Optimizer:param_groups + step

4.1 基类骨架

Optimizertorch/optim/optimizer.py:368)持有的核心结构是 self.param_groupstorch/optim/optimizer.py:425):一个 list,每项是 {"params": [...], "lr": ..., "weight_decay": ...} 这样的 dict——所以给不同参数组配不同学习率是原生能力。state_dicttorch/optim/optimizer.py:710)把分组和每参数状态(动量缓冲等)打包,checkpoint 续训靠它。

zero_gradtorch/optim/optimizer.py:1058)默认 set_to_none=True:直接把 .grad 设为 None 而不是清零——docstring 里写明了理由:省显存、省一次 kernel,且让优化器能区分「梯度为 0」和「这一步没收到梯度」(后者直接跳过)。

4.2 AdamW:相对 Adam 就改了一个位置

AdamWtorch/optim/adamw.py:20直接继承 Adam,构造时只多塞一个开关:decoupled_weight_decay=Truetorch/optim/adamw.py:48)。

真正的更新循环在 _single_tensor_adamtorch/optim/adam.py:348)。关键三行:

步骤代码(torch/optim/adam.py数学
权重衰减(decoupled):417-418 param.mul_(1 - lr * weight_decay)θ ← θ·(1 − η·λ),与梯度无关
一阶动量:455 exp_avg.lerp_(grad, 1 - beta1)m ← β₁m + (1−β₁)g
二阶动量与更新同函数后续v ← β₂v + (1−β₂)g²;θ ← θ − η·m̂/(√v̂+ε)

对比之下,非 decoupled(经典 Adam 的 L2)是把 weight_decay * param 加进梯度再算动量(torch/optim/adam.py:421-429 的 else 分支)。AdamW 的论文主张「衰减不该被动量缩放」,落在这个实现里就是一行 mul_ 的位置差异


5. 串起来:一个训练 step 在 nn/optim 层的路径

model(x) Module.__call__ → Linear.forward → F.linear → addmm
loss.backward() autograd 引擎(第 3 章)把梯度写进每个 Parameter 的 .grad
opt.step() 遍历 param_groups → 逐参数按 AdamW 公式原地更新
opt.zero_grad() 把所有 .grad 置 None,等下一轮

四行代码,每一行都对应前面某一章的机制。


6. 坑与边界

  • self.fc = nn.Linear(...)self.fc = [nn.Linear(...)] 天壤之别:list/dict 不会被 __setattr__ 展开注册,要用 nn.ModuleList/ModuleDict,否则参数进不了 parameters(),优化器根本看不到它们。
  • 直接调 forward 会绕过 hook——永远不要写 model.forward(x)
  • Parameter 就是 Tensor 子类torch/nn/parameter.py:30),默认 requires_grad=True;把普通 Tensor 塞进 _parameters 不会自动变成参数。
  • in-place 更新参数要用 param.data 或在 no_grad,否则会在反向图上留下痕迹;优化器内部的 mul_/lerp_ 都是在引擎的 inference 模式下执行的(_use_grad_for_differentiable 装饰器控制,torch/optim/adam.py:215-216)。
  • 看不出来:foreach/fused 多 tensor 路径的 kernel 细节(_multi_tensor_adamtorch/optim/adam.py:554),它把逐参数循环换成批量 kernel 启动,逻辑等价但性能路径不同。

7. 本章代码地图

主题文件路径符号名
Module 基类torch/nn/modules/module.pyModule__setattr___call_implnamed_parametersstate_dict
参数类型torch/nn/parameter.pyParameter
全连接层torch/nn/modules/linear.pyLinearreset_parameters
linear 的 C++ 实现aten/src/ATen/native/Linear.cppat::linear
Python functional 包装torch/nn/functional.pylinear(:2382)
优化器基类torch/optim/optimizer.pyOptimizerzero_gradparam_groups
AdamWtorch/optim/adamw.pyAdamW
Adam 更新核心torch/optim/adam.py_single_tensor_adam