🏠 总目录📚 本教程 05 · nn.Module 内部 ← →
📑 本页目录(点开跳转)

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())

实测输出:

信息关系

m._parameters : ['w']
m._buffers : ['running']
m.__dict__ 里的: ['training', 'plain']
m.named_parameters()→['w']
m.state_dict()→['w', 'running']
⚠️ 'plain' 在 state_dict 里吗→False

⭐ 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)

实测输出:

关键信息

_parameters: ['p']
_buffers : []
_modules : ['sub']
__dict__ : ['training', 't']
往参数槽塞普通张量 -> cannot assign 'torch.FloatTensor' as parameter 'p' (torch.nn.Parameter or None expected)

⭐ 注意 _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)

实测输出:

能前向运行,不等于参数已被注册
检查项BadNetGoodNet
_modules[]['layers']
参数个数06
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__)

实测输出:

关键信息

参数名: ['b.0', 'd.head.weight', 'd.head.bias', 'e.0.weight', 'e.0.bias']
子模块(named_modules):
'' -> Mixed
'b' -> ParameterList
'd' -> ModuleDict
'd.head' -> Linear
'e' -> Sequential
'e.0' -> Linear

⚠️ 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)。 两个后果:


🛑 读到这里可以停 —— 前半章讲完了(约 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

✅ 检查点

  1. self.w = nn.Parameter(...)、self.register_buffer("r", ...)、self.plain = torch.zeros(2) 三者分别进哪里?哪些进 parameters()、哪些进 state_dict()?
  2. 为什么普通张量赋值不会被自动认成 buffer?三个注册表里哪一个「你不主动说就不会有」?
  3. 忘写 super().__init__() 的报错原文是什么?为什么它是本章唯一一个一定会报错的坑?
  4. 把三个 nn.Linear 装进 Python list,实测 _modules、参数个数、state_dict 各是什么?forward 还跑得动吗?
  5. 「半对半错」(stem/head 注册了、中间 blocks 在 list 里)为什么比全错更危险?实测优化器管了几个张量、模型里其实有几个?那些没注册的层有没有梯度?
  6. children() 和 modules() 的区别?named_modules() 的第一条是什么?
  7. 同一个 nn.Linear 挂在两个属性上(权重共享),named_parameters() 和 state_dict() 的条数一样吗?分别是几条?
  8. .double() 为什么能一次改掉所有参数和 buffer,却改不掉普通属性张量?底层遍历的是什么?
  9. Module.to() 和 Tensor.to() 的返回值差在哪?.to() 之后 Parameter 还是同一个对象吗?
  10. m(x) 和 m.forward(x) 差在哪?实测 hook 触发条数和输出数值分别是多少?为什么 post(fc) 那条还在?
👀 答案
  1. 分别进 _parameters、_buffers、__dict__。parameters() 只看 _parameters(加上递归进子模块);state_dict() 看 _parameters + _buffers。实测:named_parameters() 是 ['w'],state_dict() 是 ['w', 'running'],⚠️ 'plain' in m.state_dict() 是 False。
  2. 因为 __setattr__ 认的是类型:nn.Parameter 是一个类、nn.Module 是一个类,而「buffer」不是任何类型,普通张量和普通张量长得一模一样。所以 _buffers 是唯一一个只能靠 register_buffer 显式登记的表(实测直接 setattr 一个张量,_buffers 是空的 [],东西掉进了 __dict__)。
  3. AttributeError: cannot assign module before Module.__init__() call。因为那三个注册表就是 Module.__init__() 建出来的空字典 —— 没建表,海关无处安放,只能当场报错。其余的坑(list 装子模块、普通属性不搬设备、.forward 跳 hook)全是静默的。
  4. _modules 是 []、参数个数 0、state_dict 是 空 list;而 forward 照跑,输出形状 torch.Size([1, 2]) 和正确版本完全一样。对照 nn.ModuleList 版本:_modules 是 ['layers']、6 个参数、6 个 key。
  5. 因为全错会撞上 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,模型的一多半从来没被存过。
  6. children() 只往下一层,modules() 递归全展开。实测 Net(enc=Block, dec=Block):children 是 ['enc', 'dec'],modules 是 ['', 'enc', 'enc.fc', 'dec', 'dec.fc']。⚠️ 第一条 '' 是模型自己,写「给每个子模块挂 hook」时别忘了它也在里面。
  7. 不一样。实测:named_parameters() 是 2 条(['a.weight', 'a.bias'],默认 remove_duplicate=True 去重了),state_dict() 是 4 个 key(a.* 和 b.* 都有,不去重)。去重让优化器不会把共享权重更新两次;⚠️ 不去重让 checkpoint 把同一块张量存两份,而且手工改 key 时只改一半会让权重共享静默断掉。
  8. 因为 _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),而加法有类型提升,一声不吭地跑完。
  9. 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() 也不会坏;不过仍建议先搬设备再建优化器,因为优化器状态是照参数当时的设备创建的。
  10. 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

打卡记录保存在你的浏览器里,首页能看到总进度