📑 本页目录(点开跳转)
05 · nn.Module 内部:它是怎么找到你的参数的
⏱ 120 分钟 | ⭐ 实测:装在 Python list 里的三个 nn.Linear,梯度照算、权重永远不动、state_dict 里也没有 —— 全程不报一个错
🎯 一句话
nn.Module 认参数只靠一件事:你写 self.x = ... 的那一刻,__setattr__ 做了一次类型检查,把东西分进三个字典(_parameters / _buffers / _modules)。
没进这三个字典的,parameters() 找不到(优化器看不见)、state_dict() 存不了(存盘丢失)、.to() 也搬不动(设备/精度不跟着变)。
⭐ 「为什么我的参数没被优化器更新」,答案几乎全在这一句里。
🚪 一、self.x = ... 的那一刻,海关做了一次类型检查
先看三种写法各去了哪:
import torch
import torch.nn as nn
class M(nn.Module):
def __init__(self):
super().__init__()
self.w = nn.Parameter(torch.ones(2))
self.register_buffer("running", torch.zeros(2))
self.plain = torch.zeros(2) # ⚠️ 普通属性
m = M()
print("m._parameters :", list(m._parameters.keys()))
print("m._buffers :", list(m._buffers.keys()))
print("m.__dict__ 里的:", [k for k in m.__dict__ if not k.startswith("_")])
print()
print("m.named_parameters() → ", [n for n, _ in m.named_parameters()])
print("m.state_dict() → ", list(m.state_dict().keys()))
print("⚠️ 'plain' 在 state_dict 里吗 →", "plain" in m.state_dict())
实测输出:
信息关系
⭐ plain 没有消失,它好好地待在 m.__dict__ 里 —— 它只是没进任何一个注册表。
m.plain 照样读得到、照样能参与 forward 的计算,所以你不会有任何异样感。
(顺带记住 __dict__ 里那个 training —— 它就是 train() / eval() 唯一改的那个布尔,是第 07 章的主角。)
分拣规则只有两行
__setattr__ 认的是类型,不是名字:
import torch
import torch.nn as nn
m = nn.Module()
for name, val in [("t", torch.zeros(2)),
("p", nn.Parameter(torch.zeros(2))),
("sub", nn.Linear(2, 2))]:
setattr(m, name, val)
print("_parameters:", list(m._parameters.keys()))
print("_buffers :", list(m._buffers.keys()))
print("_modules :", list(m._modules.keys()))
print("__dict__ :", [k for k in m.__dict__ if not k.startswith("_")])
try:
m.p = torch.zeros(2) # ⚠️ 往参数槽里塞普通张量
except TypeError as e:
print("往参数槽塞普通张量 ->", e)
实测输出:
关键信息
⭐ 注意 _buffers 是空的:普通张量 t 掉进了 __dict__。
buffer 没有专属类型 —— nn.Parameter 是个类,「buffer」不是,所以它只能靠 register_buffer 显式登记,没有第二条路。这是三个表里唯一一个「你不主动说就不会有」的。
⚠️ 槽位一旦定了就换不掉:w 已经在 _parameters 里,再往它上面赋一个普通张量会直接 TypeError。
(想换成别的值,写 m.w.data.copy_(新值),或者 m.w = nn.Parameter(新值)。)
💀 忘了 super().__init__()
import torch.nn as nn
class Forgot(nn.Module):
def __init__(self):
self.fc = nn.Linear(2, 2) # ⚠️ 上面没写 super().__init__()
try:
Forgot()
except AttributeError as e:
print("AttributeError:", e)
实测输出:
AttributeError: cannot assign module before Module.__init__() call
⭐ 这条报错是这一章最好的证据:三个注册表就是 Module.__init__() 建出来的空字典。
没建表,海关无处安放东西,所以它当场报错 —— 这是本章唯一一个「一定会报错」的坑,其余全是静默的。
一张表记住谁进哪
| 你写的 | 进哪 | parameters() |
state_dict() |
.to() 跟着搬 |
|---|---|---|---|---|
self.w = nn.Parameter(t) |
_parameters |
✅ | ✅ | ✅ |
self.register_buffer("r", t) |
_buffers |
❌ | ✅ | ✅ |
register_buffer("r", t, persistent=False) |
_buffers |
❌ | ⚠️ 不进 | ✅ |
self.fc = nn.Linear(...) |
_modules |
✅(递归进去找) | ✅ | ✅ |
self.plain = torch.zeros(2) |
__dict__ |
❌ | ❌ | ❌ |
最后那个 persistent=False 实测确认过:
import torch
import torch.nn as nn
class M(nn.Module):
def __init__(self):
super().__init__()
self.register_buffer("keep", torch.zeros(2))
self.register_buffer("tmp", torch.zeros(2), persistent=False) # ⭐
m = M()
print("named_buffers:", [n for n, _ in m.named_buffers()])
print("state_dict :", list(m.state_dict().keys()))
m.double()
print("tmp 跟着 .double() 了吗:", m.tmp.dtype)
实测输出:
要点
named_buffers: ['keep', 'tmp']
state_dict : ['keep']
tmp 跟着 .double() 了吗: torch.float64
⭐ persistent=False 是「跟着模型搬设备,但不占 checkpoint 体积」,
适合位置编码表、因果掩码这类能从超参重新算出来的常量。
🛑 读到这里可以停 —— 已经读了约 19 分钟。 后面还有(约 35 分钟):装在 Python list 里的子模块,一个都不算数 · 名字是从根往下递归拼出来的 回来的时候不用重读,直接从下一节接着看就行。
💀 二、装在 Python list 里的子模块,一个都不算数
⚠️ 这是本章最贵的一条,也是「参数不更新」最常见的成因。
import torch
import torch.nn as nn
class BadNet(nn.Module):
def __init__(self):
super().__init__()
self.layers = [nn.Linear(2, 2) for _ in range(3)] # 💀 Python list
def forward(self, x):
for lay in self.layers:
x = lay(x)
return x
class GoodNet(nn.Module):
def __init__(self):
super().__init__()
self.layers = nn.ModuleList([nn.Linear(2, 2) for _ in range(3)]) # ⭐
def forward(self, x):
for lay in self.layers:
x = lay(x)
return x
for cls in (BadNet, GoodNet):
m = cls()
print(f"--- {cls.__name__} ---")
print(" _modules :", list(m._modules.keys()))
print(" 参数个数 :", len(list(m.parameters())))
print(" state_dict key :", list(m.state_dict().keys()))
print(" forward 能跑吗 :", m(torch.randn(1, 2)).shape)
实测输出:
| 检查项 | BadNet | GoodNet |
|---|---|---|
| _modules | [] | ['layers'] |
| 参数个数 | 0 | 6 |
| state_dict key | [] | ['layers.0.weight', 'layers.0.bias', 'layers.1.weight', 'layers.1.bias', 'layers.2.weight', 'layers.2.bias'] |
| forward 能跑吗 | torch.Size([1, 2]) | torch.Size([1, 2]) |
⭐ 最要命的是最后一行:forward 照跑,形状一模一样。
一个 list 不是 nn.Module 也不是 nn.Parameter,海关把它整个丢进 __dict__,
于是 _modules 是空的、参数 0 个、state_dict 是空字典 —— 而 Python 的 for 循环当然还是能遍历它。
⚠️ 全错反而是好事,半错才致命
如果整个模型都在 list 里,你会撞上唯一一条报错:
import torch
import torch.nn as nn
class AllBad(nn.Module):
def __init__(self):
super().__init__()
self.layers = [nn.Linear(2, 2) for _ in range(3)] # 💀 全在 list 里
try:
torch.optim.SGD(AllBad().parameters(), lr=0.1)
except ValueError as e:
print("ValueError:", e)
实测输出:
ValueError: optimizer got an empty parameter list
💀 但真实的模型几乎不会全错。 典型形态是:stem 和 head 好好写着,中间那一堆重复的 block 顺手用了 list:
import torch
import torch.nn as nn
class HalfBad(nn.Module):
def __init__(self):
super().__init__()
self.stem = nn.Linear(4, 4) # ⭐ 注册了
self.blocks = [nn.Linear(4, 4) for _ in range(3)] # 💀 没注册
self.head = nn.Linear(4, 1) # ⭐ 注册了
def forward(self, x):
x = self.stem(x)
for b in self.blocks:
x = torch.relu(b(x))
return self.head(x)
torch.manual_seed(0)
m = HalfBad()
snap = {"stem": m.stem.weight.clone(), "block0": m.blocks[0].weight.clone()}
opt = torch.optim.SGD(m.parameters(), lr=0.1) # ⚠️ 不报错,因为不是空的
print("优化器管着几个张量:", sum(len(g["params"]) for g in opt.param_groups), "个")
print("模型里其实有几个 :", 2 + 2 * 3 + 2, "个")
for _ in range(20):
opt.zero_grad()
m(torch.randn(16, 4)).pow(2).mean().backward()
opt.step()
print("stem 变了吗:", not torch.equal(snap["stem"], m.stem.weight))
print("blocks 变了吗:", not torch.equal(snap["block0"], m.blocks[0].weight))
print("blocks[0] 有梯度吗:", m.blocks[0].weight.grad is not None)
print("state_dict:", list(m.state_dict().keys()))
实测输出:
对照
优化器管着几个张量: 4 个
模型里其实有几个 : 10 个
stem 变了吗: True
blocks 变了吗: False
blocks[0] 有梯度吗: True
state_dict: ['stem.weight', 'stem.bias', 'head.weight', 'head.bias']
⭐⭐ 把这四行连起来读,就是这个 bug 的完整画像:
| 现象 | 为什么不会引起怀疑 |
|---|---|
| 优化器建得起来(4 个,不是 0 个) | ⚠️ ValueError 不会出现 |
blocks[0].weight.grad is not None |
⭐ 梯度是真的算出来了 —— autograd 只看图,不看注册表 |
| 训练照跑,loss 也在降 | 因为 stem 和 head 确实在学 |
blocks 的权重 20 步后一位没动 |
优化器手里根本没有它们,step() 无从下手 |
| checkpoint 少 6 个张量 | state_dict 里只有 4 个 key —— 模型的一多半从来没被存过 |
💀 代价:模型容量凭空少了一大截(这里 10 个张量只有 4 个在学),指标就是「差一点」; 你会去调学习率、加正则、换优化器 —— 而没有任何一条日志、报错或警告指向真正的原因。
修法:五个容器,各修各的
import torch
import torch.nn as nn
class Mixed(nn.Module):
def __init__(self):
super().__init__()
self.a = [nn.Parameter(torch.ones(2))] # 💀 list of Parameter
self.b = nn.ParameterList([nn.Parameter(torch.ones(2))]) # ⭐
self.c = {"head": nn.Linear(2, 2)} # 💀 dict of Module
self.d = nn.ModuleDict({"head": nn.Linear(2, 2)}) # ⭐
self.e = nn.Sequential(nn.Linear(2, 2)) # ⭐ 也是注册的
m = Mixed()
print("参数名:", [n for n, _ in m.named_parameters()])
print()
print("子模块(named_modules):")
for n, mod in m.named_modules():
print(" ", repr(n), "->", type(mod).__name__)
实测输出:
关键信息
⚠️ a 和 c 连名字都没出现 —— 它们不在任何遍历结果里。
| 你想装的东西 | 用这个 |
|---|---|
| 一串子模块 | nn.ModuleList(纯容器)或 nn.Sequential(还负责依次调用) |
| 名字 → 子模块 | nn.ModuleDict |
一串 Parameter |
nn.ParameterList |
名字 → Parameter |
nn.ParameterDict |
| 一个不参与梯度的张量 | ⭐ register_buffer |
⚠️ nn.ModuleList 不是 nn.Sequential:前者没有 forward,直接 m.layers(x) 会报错,你必须自己写 for 循环;后者会按顺序调用。想要 for 循环里做点别的(残差、条件分支)就用 ModuleList。
🏷️ 三、名字是从根往下递归拼出来的
import torch
import torch.nn as nn
class Block(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(2, 2)
self.register_buffer("cnt", torch.zeros(1))
class Net(nn.Module):
def __init__(self):
super().__init__()
self.enc = Block()
self.dec = Block()
m = Net()
print("named_parameters:", [n for n, _ in m.named_parameters()])
print("named_buffers :", [n for n, _ in m.named_buffers()])
print()
print("children (只往下一层):", [n for n, _ in m.named_children()])
print("modules (全展开) :", [n for n, _ in m.named_modules()])
实测输出:
对照
named_parameters: ['enc.fc.weight', 'enc.fc.bias', 'dec.fc.weight', 'dec.fc.bias']
named_buffers : ['enc.cnt', 'dec.cnt']
children (只往下一层): ['enc', 'dec']
modules (全展开) : ['', 'enc', 'enc.fc', 'dec', 'dec.fc']
⭐ 遍历 API 就四对,区别只有两条轴:
| 只往下一层 | 全展开(递归) | |
|---|---|---|
| 子模块 | children() / named_children() |
modules() / named_modules() |
| 参数 | _parameters 直接读 |
⭐ parameters() / named_parameters() |
| buffer | _buffers 直接读 |
buffers() / named_buffers() |
⚠️ named_modules() 的第一条是模型自己,名字是空字符串 ''。
写「给每个子模块挂 hook」之类的循环时,别忘了这一条也在里面(否则你会给根模块也挂一份)。
⭐ apply(fn) 是同一套递归,最常见的用法是自定义初始化:
import torch.nn as nn
net = nn.Sequential(nn.Linear(3, 4), nn.ReLU(), nn.Linear(4, 2))
hit = []
net.apply(lambda mod: hit.append(type(mod).__name__))
print("apply 走过:", hit)
实测输出:
要点
apply 走过: ['Linear', 'ReLU', 'Linear', 'Sequential']
⚠️ 注意顺序是从叶子往上,根模块 Sequential 排在最后。
⭐ 权重共享:两个遍历给的答案不一样
import torch
import torch.nn as nn
class Tied(nn.Module):
def __init__(self):
super().__init__()
shared = nn.Linear(2, 2)
self.a = shared
self.b = shared # ⭐ 同一个对象挂在两个属性上(权重共享)
m = Tied()
print("named_parameters:", [n for n, _ in m.named_parameters()])
print("参数张量个数 :", len(list(m.parameters())))
print("state_dict keys :", list(m.state_dict().keys()))
print("a.weight is b.weight:", m.a.weight is m.b.weight)
print()
print("关掉去重后:", [n for n, _ in m.named_parameters(remove_duplicate=False)])
实测输出:
要点
named_parameters: ['a.weight', 'a.bias']
参数张量个数 : 2
state_dict keys : ['a.weight', 'a.bias', 'b.weight', 'b.bias']
a.weight is b.weight: True
关掉去重后: ['a.weight', 'a.bias', 'b.weight', 'b.bias']
⭐ parameters() 去重(2 个),state_dict() 不去重(4 个 key)。 两个后果:
- ✅ 优化器不会把共享权重更新两次 —— 这正是你要的(梯度已经在反向里加过一遍了,见第 02 章「一个参数被多条路径用到」)。
- ⚠️ checkpoint 里同一块张量存了两份:key 是双份、体积也是双份。加载回去结果仍然正确(两次
copy_写的是同一块存储),只是白占地方;⭐ 但如果你手工改 key、只改了a.*忘了b.*,加载完两边就不再是同一块内存,权重共享当场断掉而不报错。
🛑 读到这里可以停 —— 前半章讲完了(约 30 分钟):参数是怎么被找到的(
__setattr__的类型检查 + 三个注册表)、装进 Python list 会静默丢掉一多半模型、以及名字是怎么递归拼出来的。 后半章还有:.to()/.float()为什么能一次改掉全部(以及谁不会跟着搬) ·forward和__call__的区别(hook 会被跳过) · 一个 20 行的自检器 回来的时候不用重读,直接从下一节接着看就行。
🚚 四、.to() / .float() 为什么能一次改掉全部
答案就是第一节那三个表 —— 因为有表可遍历。
nn.Module._apply(fn) 做三件事:递归所有子模块 → 对 _parameters 里每一个套 fn → 对 _buffers 里每一个套 fn。
⚠️ __dict__ 它根本不看。
import torch
import torch.nn as nn
class M(nn.Module):
def __init__(self):
super().__init__()
self.w = nn.Parameter(torch.ones(2))
self.register_buffer("running", torch.zeros(2))
self.plain = torch.zeros(2) # ⚠️ 普通属性
m = M()
print("改之前:", m.w.dtype, m.running.dtype, m.plain.dtype)
m.double() # 等价于 m.to(torch.float64)
print("double() 之后:")
print(" w (Parameter):", m.w.dtype)
print(" running (buffer) :", m.running.dtype)
print(" plain (普通属性) :", m.plain.dtype, " ⚠️ 没跟着变")
实测输出:
对照
改之前: torch.float32 torch.float32 torch.float32
double() 之后:
w (Parameter): torch.float64
running (buffer) : torch.float64
plain (普通属性) : torch.float32 ⚠️ 没跟着变
⭐ 同一个机制,换成 device 就是那条最著名的报错。
.to("cuda") 只搬 _parameters 和 _buffers;self.plain 留在 CPU 上,前向撞上 Expected all tensors to be on the same device —— 而模型本身明明已经 .to("cuda") 过了,所以你会先怀疑数据、怀疑 DataLoader,很难想到怀疑一个属性。
(🗓️ 未实跑 —— 需要 GPU。本机 CUDA 不可用,无法给出这条报错的原文。下面用 dtype 复现同一个机制。)
import torch
import torch.nn as nn
class M(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(2, 2)
self.proj = torch.eye(2) # 💀 普通属性,参与了计算
def forward(self, x):
return self.fc(x) @ self.proj
m = M().double()
print("fc.weight:", m.fc.weight.dtype, " proj:", m.proj.dtype)
try:
m(torch.ones(1, 2, dtype=torch.float64))
except RuntimeError as e:
print("RuntimeError:", e)
实测输出:
fc.weight: torch.float64 proj: torch.float32
RuntimeError: expected m1 and m2 to have the same dtype, but got: double != float
💀 这还算走运的 —— 它报错了。 我在实测里顺手试过把 @ self.proj 换成 + self.bias2(一个普通属性张量),
同一段代码一声不吭地跑完 —— 因为加法有类型提升,float32 会被自动提上 float64。
⭐ 判据:矩阵乘法卡 dtype,逐元素运算不卡。 所以「没报错」完全不能证明属性搬对了。
Module.to 是原地的,Tensor.to 不是
import torch
import torch.nn as nn
m = nn.Linear(2, 2)
print("Module.to 返回的是自己吗:", m.to(torch.float64) is m)
x = torch.ones(2)
print("Tensor.to 返回的是自己吗:", x.to(torch.float64) is x)
x.to(torch.float64) # 💀 没接返回值
print("忘了接返回值,x 还是:", x.dtype)
实测输出:
要点
Module.to 返回的是自己吗: True
Tensor.to 返回的是自己吗: False
忘了接返回值,x 还是: torch.float32
⭐ 这解释了 ML 基础 15 章 那份模板里两种写法为什么不一样:
model = make_model().to(device) 接不接返回值都行(Module.to 改的是注册表里的东西,返回 self 只是为了链式书写),
而 xb, yb = xb.to(device), yb.to(device) 必须接 —— 张量是不可变对象,.to() 给你一个新的。
⭐ 转换之后,Parameter 对象没换
import torch
import torch.nn as nn
m = nn.Linear(3, 2)
before = m.weight
m.float() # 已经是 float32,什么都不用做
print("dtype 没变时是同一个对象吗:", m.weight is before)
before_ptr = m.weight.data_ptr()
m.double()
print("dtype 变了之后还是同一块存储吗:", m.weight.data_ptr() == before_ptr)
print("变完还是 Parameter 吗:", type(m.weight).__name__,
" requires_grad =", m.weight.requires_grad)
opt = torch.optim.SGD(m.parameters(), lr=0.1)
print("⭐ 优化器拿到的是变换【之后】的对象吗:",
opt.param_groups[0]["params"][0] is m.weight)
实测输出:
要点
dtype 没变时是同一个对象吗: True
dtype 变了之后还是同一块存储吗: False
变完还是 Parameter 吗: Parameter requires_grad = True
⭐ 优化器拿到的是变换【之后】的对象吗: True
⭐ 换了存储,没换对象。 这一条决定了「先建优化器还是先 .to()」有没有关系 —— 实测:
import torch
import torch.nn as nn
torch.manual_seed(0)
m = nn.Linear(2, 2)
opt = torch.optim.SGD(m.parameters(), lr=0.1) # ⭐ 先建优化器
held = opt.param_groups[0]["params"][0]
m.double() # 再转换
print("优化器手里的还是同一个对象吗:", held is m.weight)
print("dtype:", held.dtype)
before = m.weight.clone()
m(torch.ones(4, 2, dtype=torch.float64)).pow(2).mean().backward()
opt.step()
print("step() 之后权重变了吗:", not torch.equal(before, m.weight))
实测输出:
要点
优化器手里的还是同一个对象吗: True
dtype: torch.float64
step() 之后权重变了吗: True
⭐ 顺序不影响正确性 —— 优化器持有的是 Parameter 对象的引用,.to() 只换它内部的存储。
⚠️ 但仍然建议先 .to(device) 再建优化器:优化器状态(Adam 的 exp_avg 等)是在第一次 step() 时照着参数当时的设备创建的,先搬完再建,能避免状态和参数分处两地这类边角问题。
🛑 第二个休息点 —— 中段讲完了(约 20 分钟)。 最后一段还有:
forward和__call__不是一回事 · 一个 20 行的自检器 这一章确实长,分三次读完全没问题 —— 回来直接从下一节接着看。
📞 五、forward 和 __call__ 不是一回事
m(x) 走的是 nn.Module.__call__,它按顺序做四件事:
跑 forward pre-hook → 跑 forward → 跑 forward hook → 给输出挂上 backward hook。
m.forward(x) 只做中间那一件。
import torch
import torch.nn as nn
fired = []
class Net(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(2, 2)
def forward(self, x):
return self.fc(x)
m = Net()
m.register_forward_pre_hook(lambda mod, inp: fired.append("pre(Net)"))
m.register_forward_hook(lambda mod, inp, out: fired.append("post(Net)"))
m.fc.register_forward_hook(lambda mod, inp, out: fired.append("post(fc)"))
x = torch.ones(1, 2)
fired.clear()
m(x)
print("m(x) 触发了:", fired)
fired.clear()
m.forward(x)
print("m.forward(x) 触发了:", fired, " ⚠️ 少了两条")
实测输出:
对照
m(x) 触发了: ['pre(Net)', 'post(fc)', 'post(Net)']
m.forward(x) 触发了: ['post(fc)'] ⚠️ 少了两条
⭐ 注意 post(fc) 还在。 被跳过的只有你直接点名的那一层的 hook —— 它内部的子模块照常走各自的 __call__。
⚠️ 这让 bug 更难发现:hook 不是全没了,是少了一部分,日志里还有东西在打。
后果不止是「少打了日志」
import torch
import torch.nn as nn
class Net(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(2, 1, bias=False)
with torch.no_grad():
self.fc.weight.fill_(1.0)
def forward(self, x):
return self.fc(x)
m = Net()
# ⭐ 一个真会改数的 pre-hook:把输入减掉均值(很多归一化/量化封装就是这么挂上去的)
m.register_forward_pre_hook(lambda mod, inp: (inp[0] - inp[0].mean(),))
x = torch.tensor([[1.0, 3.0]])
print("m(x) =", m(x).item(), " ← hook 生效")
print("m.forward(x) =", m.forward(x).item(), " ⚠️ 数值不同,而且不报错")
实测输出:
对照
m(x) = 0.0 ← hook 生效
m.forward(x) = 4.0 ⚠️ 数值不同,而且不报错
💀 0.0 和 4.0。 同一个模型、同一个输入,两种调用方式给出两个答案,没有任何警告。
谁在偷偷用 hook(也就是「你会跳过谁」):
| 谁 | 挂的是哪种 |
|---|---|
| 特征提取、中间层可视化、激活值统计 | forward hook |
| 梯度监控、梯度裁剪的诊断 | register_full_backward_hook |
量化、剪枝、nn.utils.parametrize 这类「改写权重再用」的封装 |
⭐ forward pre-hook |
| 凡是「在模型外面包一层」的工具(编译、图捕获、分布式包装) | 入口都是 __call__ |
⚠️ 规则一句话:永远写 model(x),不要写 model.forward(x)。
唯一合理的例外,是你在自己的 forward 里显式调用父类实现(super().forward(x))。
🛑 读到这里可以停 —— 已经读了约 83 分钟。 最后一段还有(约 34 分钟):一个 20 行的自检器 · 检查点与走神救援 回来的时候不用重读,直接从下一节接着看就行。
🔍 六、一个 20 行的自检器
前面五节的坑全是静默的,所以值得在模型构造完之后加一道体检。
⭐ 判据要认结构:不看你给属性起的名字,只看 __dict__ 里躺着的东西是不是张量或模块。
import torch
import torch.nn as nn
def find_unregistered(model):
"""扫出所有【躺在 __dict__ 里】的张量和子模块 —— 它们不在任何注册表里。"""
bad = []
for name, mod in model.named_modules():
prefix = name + "." if name else ""
for k, v in vars(mod).items(): # ⭐ vars() 就是 __dict__
if k.startswith("_"):
continue
if isinstance(v, (torch.Tensor, nn.Module)):
bad.append((prefix + k, type(v).__name__))
elif isinstance(v, (list, tuple)) and any(
isinstance(i, (torch.Tensor, nn.Module)) for i in v):
bad.append((prefix + k, f"{type(v).__name__}[{len(v)}]"))
elif isinstance(v, dict) and any(
isinstance(i, (torch.Tensor, nn.Module)) for i in v.values()):
bad.append((prefix + k, f"dict[{len(v)}]"))
return bad
class HalfBad(nn.Module):
def __init__(self):
super().__init__()
self.stem = nn.Linear(4, 4)
self.blocks = [nn.Linear(4, 4) for _ in range(3)] # 💀
self.mask = torch.ones(4) # 💀
self.head = nn.Linear(4, 1)
for name, ty in find_unregistered(HalfBad()):
print(f"⚠️ 没注册: {name} ({ty})")
print("对照:干净的模型扫出", len(find_unregistered(nn.Sequential(nn.Linear(2, 2)))), "条")
实测输出:
对照
⚠️ 没注册: blocks (list[3])
⚠️ 没注册: mask (Tensor)
对照:干净的模型扫出 0 条
⭐ 注意它跳过了 _ 开头的键(三个注册表自己就叫 _parameters / _buffers / _modules),
也跳过了 training(那是个 bool 不是张量)—— 干净模型上的误报是 0 条。
还有两个更粗、但一秒钟就能做的自检:
import torch.nn as nn
model = nn.Sequential(nn.Linear(4, 4), nn.ReLU(), nn.Linear(4, 1))
print("参数量:", sum(p.numel() for p in model.parameters()))
print("state_dict 条数:", len(model.state_dict()))
实测输出:
要点
参数量: 25
state_dict 条数: 4
⭐ 和你心算的对不上,就去查注册表。 参数量少一大截、state_dict 条数少一半,这两个数字比 loss 曲线诚实得多。
🔗 这一章连到哪里
| 相关的地方 | 为什么 |
|---|---|
06 · state_dict 到底装了什么 |
⭐ 下一章。本章讲谁会被找到(进不进那三个表),那一章讲找到之后长什么样、怎么存怎么取,以及加载报错怎么读 |
| 02 · autograd 怎么建图 | ⭐ 第二节那三个没注册的层梯度是真的算出来了的 —— 因为 autograd 只看图不看注册表。⚠️ 另外 02 章说过 optimizer.zero_grad() 只清「这个优化器管着的」参数,管着谁正是本章决定的 |
| 04 · detach 与 no_grad | 冻结 backbone 的写法是遍历 parameters() 把 requires_grad_ 关掉 —— ⭐ 能这么遍历,靠的就是本章的注册表;反过来,没注册的层你连冻都冻不了 |
07 · train() / eval() 改了什么 |
第一节 m.__dict__ 里那个 training 就是下下章的全部主角;⭐ 而 BN 的 running_mean 是 buffer,正落在本章那张表「进 _buffers、不进 _parameters」那一格 |
| 09 · 自己写一个 autograd 算子 | ⚠️ 两个 forward 不是一回事:那一章写的是 autograd.Function 的 forward,要配一个 backward;本章说的是 nn.Module 的 forward,反向由 autograd 自动来 |
| 01 · 张量到底是什么 | .to() 改的 dtype / device 就是说明书上那两个字段;⭐ 本章讲的是「怎么一次改掉几百个张量的这两个字段,以及谁会被漏下」 |
| 《机器学习与深度学习基础》15 · PyTorch 实战手册 | 那份模板里 model = make_model().to(device) 和 xb, yb = xb.to(device), yb.to(device) 是两种写法 —— ⭐ 第四节实测了为什么模型那行可以不接返回值、数据那行必须接 |
| 《机器学习与深度学习基础》附录C 第 5 题 | 手写 BatchNorm 时 running_mean 用 register_buffer 而不是 Parameter(那里的理由是「不参与梯度更新,但必须跟着 state_dict 存盘」)—— ⭐ 本章第一节那张表就是这句话的完整版,还多出一格 persistent=False |
✅ 检查点
self.w = nn.Parameter(...)、self.register_buffer("r", ...)、self.plain = torch.zeros(2)三者分别进哪里?哪些进parameters()、哪些进state_dict()?- 为什么普通张量赋值不会被自动认成 buffer?三个注册表里哪一个「你不主动说就不会有」?
- 忘写
super().__init__()的报错原文是什么?为什么它是本章唯一一个一定会报错的坑? - 把三个
nn.Linear装进 Python list,实测_modules、参数个数、state_dict各是什么?forward还跑得动吗? - 「半对半错」(stem/head 注册了、中间 blocks 在 list 里)为什么比全错更危险?实测优化器管了几个张量、模型里其实有几个?那些没注册的层有没有梯度?
children()和modules()的区别?named_modules()的第一条是什么?- 同一个
nn.Linear挂在两个属性上(权重共享),named_parameters()和state_dict()的条数一样吗?分别是几条? .double()为什么能一次改掉所有参数和 buffer,却改不掉普通属性张量?底层遍历的是什么?Module.to()和Tensor.to()的返回值差在哪?.to()之后Parameter还是同一个对象吗?m(x)和m.forward(x)差在哪?实测 hook 触发条数和输出数值分别是多少?为什么post(fc)那条还在?
👀 答案
- 分别进
_parameters、_buffers、__dict__。parameters()只看_parameters(加上递归进子模块);state_dict()看_parameters+_buffers。实测:named_parameters()是['w'],state_dict()是['w', 'running'],⚠️'plain' in m.state_dict()是False。 - 因为
__setattr__认的是类型:nn.Parameter是一个类、nn.Module是一个类,而「buffer」不是任何类型,普通张量和普通张量长得一模一样。所以_buffers是唯一一个只能靠register_buffer显式登记的表(实测直接setattr一个张量,_buffers是空的[],东西掉进了__dict__)。 AttributeError: cannot assign module before Module.__init__() call。因为那三个注册表就是Module.__init__()建出来的空字典 —— 没建表,海关无处安放,只能当场报错。其余的坑(list 装子模块、普通属性不搬设备、.forward跳 hook)全是静默的。_modules是[]、参数个数 0、state_dict是 空 list;而forward照跑,输出形状torch.Size([1, 2])和正确版本完全一样。对照nn.ModuleList版本:_modules是['layers']、6 个参数、6 个 key。- 因为全错会撞上
ValueError: optimizer got an empty parameter list(报错反而是好事);半错则一条错都不报。实测HalfBad:优化器管着 4 个张量、模型里其实有 10 个;20 步之后stem变了、blocks一位没动;⚠️ 但blocks[0].weight.grad is not None是True—— 梯度是真的算出来了的,autograd 只看图不看注册表,只是优化器手里没有这些张量。state_dict也只有 4 个 key,模型的一多半从来没被存过。 children()只往下一层,modules()递归全展开。实测Net(enc=Block, dec=Block):children 是['enc', 'dec'],modules 是['', 'enc', 'enc.fc', 'dec', 'dec.fc']。⚠️ 第一条''是模型自己,写「给每个子模块挂 hook」时别忘了它也在里面。- 不一样。实测:
named_parameters()是 2 条(['a.weight', 'a.bias'],默认remove_duplicate=True去重了),state_dict()是 4 个 key(a.*和b.*都有,不去重)。去重让优化器不会把共享权重更新两次;⚠️ 不去重让 checkpoint 把同一块张量存两份,而且手工改 key 时只改一半会让权重共享静默断掉。 - 因为
_apply(fn)递归子模块之后,只遍历_parameters和_buffers两个字典,__dict__它根本不看。实测m.double()之后w和running都是torch.float64,而plain还是torch.float32。换成device就是那条「Expected all tensors to be on the same device」。⚠️ 而且不一定会报错:矩阵乘法卡 dtype(实测expected m1 and m2 to have the same dtype, but got: double != float),而加法有类型提升,一声不吭地跑完。 Module.to()是原地改注册表里的东西、返回self(实测m.to(torch.float64) is m为True);Tensor.to()返回一个新张量(x.to(...) is x为False,不接返回值x还是float32)。.to()之后Parameter对象没换(data_ptr变了,但opt.param_groups[0]["params"][0] is m.weight仍是True),所以先建优化器再.to()也不会坏;不过仍建议先搬设备再建优化器,因为优化器状态是照参数当时的设备创建的。m(x)走__call__:pre-hook →forward→ forward hook → 给输出挂 backward hook;m.forward(x)只做中间一件。实测触发条数 3 条 vs 1 条(['pre(Net)', 'post(fc)', 'post(Net)']vs['post(fc)'])。post(fc)还在,是因为只有你直接点名的那一层被跳过,它内部的子模块照常走自己的__call__—— ⚠️ 所以 hook 不是全没了、是少了一部分,更难发现。后果不止少日志:挂一个会改数的 pre-hook 之后,实测m(x)是 0.0、m.forward(x)是 4.0,没有任何警告。
🛑 可以停在这里
⚡ 走神救援
⭐
nn.Module认参数只靠__setattr__的一次类型检查:nn.Parameter进_parameters、nn.Module进_modules、其余一律掉进普通的__dict__。⭐ buffer 是唯一没有专属类型的,所以只能靠register_buffer显式登记。这三个表决定了三件事:优化器看不看得见、
state_dict()存不存、.to()搬不搬得动。这就是「为什么我的参数没被优化器更新」的全部答案。💀 最贵的形态是把子模块装进 Python list:参数 0 个、
state_dict是空的,⭐ 而forward照跑、输出形状和正确版本一模一样。
全错会报「optimizer got an empty parameter list」——⭐ 报错反而是好事;
半错才致命:优化器只管住一部分,训练几十步后一部分变了另一部分一位没动,⚠️ 而且那些没被注册的参数梯度并不是 None(autograd 只看图不看注册表)——模型的一多半既没在学也没被存过,没有一条日志会告诉你。
修法是
nn.ModuleList/ModuleDict/ParameterList/register_buffer。⚠️ModuleList没有forward,要自己写循环。⭐ 权重共享时两个遍历给的答案不同:
named_parameters()去重(让优化器不重复更新),state_dict()不去重(于是 checkpoint 存了两份,而手工改 key 只改一半会静默断掉共享)。⚠️
.to()只搬那两个表里的东西,普通属性它根本不看——换 dtype 时矩阵乘法会报错,但加法有类型提升,会一声不吭跑完。
下一节 👉 06-state_dict与存取.md