数据截至 (上游 commit c187ef3271d5)
04 · nn.Module 与 Optimizer:每天写的训练循环在干什么
这一章讲什么:
nn.Linear和optim.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 内部状态
Module(torch/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——赋值本身就是注册。