🏠 总目录📚 本教程 06 · state_dict 与存取 ← →
📑 本页目录(点开跳转)

06 · state_dict 到底装了什么,以及加载为什么老是报错

⏱ 68 分钟 | ⭐ 实测:state_dict() 拿到的张量和模型参数是同一块内存 —— 这就是早停回滚必须 .clone() 的原因


🎯 一句话

state_dict() 返回的是一个普通的 OrderedDict,key 是「模块路径 + 属性名」拼出来的字符串,value 是参数张量本身的引用(不是拷贝)。 所以它有两副面孔:存盘的时候它是一份清单,在内存里传来传去的时候它是一堆活的引用 —— 后一副面孔坑过的人比前一副多得多。


🧩 一、先把它打开看看

站内有 6 处代码在用 state_dict(ml_md/15 的早停、ml_md/10 的回滚、rl_md/08 的目标网络同步、ai_md/11 的分布式 checkpoint),每一处都是直接用,没有一处打开看过它长什么样。先看:

import torch
import torch.nn as nn


class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc = nn.Linear(3, 2)
        self.register_buffer("running", torch.zeros(2))
        self.plain = torch.zeros(2)              # ⚠️ 普通属性

    def forward(self, x):
        return self.fc(x)


m = Net()
sd = m.state_dict()
print("type      :", type(sd).__name__)
print("keys      :", list(sd.keys()))
print("值的类型  :", type(sd["fc.weight"]).__name__)
print("requires_grad:", sd["fc.weight"].requires_grad)

实跑输出:

对照

type : OrderedDict

keys : ['running', 'fc.weight', 'fc.bias']

值的类型 : Tensor

requires_grad: False

三件事值得记住:

观察 意思
类型是 OrderedDict 它就是个字典,你可以 del、可以改 key、可以自己拼一个塞进去
plain 不在里面 普通属性不进 state_dict(上一章 nn.Module 内部讲过为什么),存盘会丢
requires_grad 是 False 值是 Tensor 不是 Parameter ——state_dict() 出门前做了一次 detach

⚠️ detach 过 ≠ 拷贝过。这是下一节整节的内容。


🧨 二、💀 它装的是引用,不是快照

ml_md/15 的早停模板里有这么一行,站内没有任何一页解释过为什么要 .clone():

best_state = {k: v.clone() for k, v in model.state_dict().items()}

把 .clone() 去掉会发生什么,实测:

import torch
import torch.nn as nn

m = nn.Linear(2, 1)
with torch.no_grad():
    m.weight.fill_(1.0)

best_wrong = m.state_dict()                                     # ❌ 直接拿
best_right = {k: v.clone() for k, v in m.state_dict().items()}  # ⭐ 拷一份

with torch.no_grad():                                           # 继续训练,权重变坏
    m.weight.fill_(7.0)

print("现在的权重      :", m.weight[0].tolist())
print("❌ 没 clone 的备份:", best_wrong["weight"][0].tolist())
print("⭐ clone 过的备份 :", best_right["weight"][0].tolist())

m.load_state_dict(best_wrong)
print("用没 clone 的回滚后:", m.weight[0].tolist(), "← 什么都没回滚")
m.load_state_dict(best_right)
print("用 clone 的回滚后  :", m.weight[0].tolist())

实跑输出:

对照

现在的权重 : [7.0, 7.0]

❌ 没 clone 的备份: [7.0, 7.0]

⭐ clone 过的备份 : [1.0, 1.0]

用没 clone 的回滚后: [7.0, 7.0] ← 什么都没回滚

用 clone 的回滚后 : [1.0, 1.0]

底层证据是同一个指针:

import torch
import torch.nn as nn

m = nn.Linear(3, 2)
sd = m.state_dict()
print("是同一块存储吗:", sd["weight"].data_ptr() == m.weight.data_ptr())

输出 是同一块存储吗: True。

⭐ 判据:state_dict() 落到磁盘的那一刻才产生快照。留在内存里的 state_dict,权重一变它就跟着变。 所以 torch.save(model.state_dict(), path) 不需要 clone(存盘本身就是一次拷贝), 而 「先存一份最好的、以后回滚」必须 clone。

💀 这个 bug 的可怕之处在于它不报错:早停照常触发、日志照常打印「回滚最佳权重」、最终指标只是比预期低一点。 ml_md/15 的实测里早停发生在第 51 轮、回滚到第 10 轮的权重;如果漏了 .clone(),你回滚到的其实是第 51 轮的权重,而验证集指标是拿第 10 轮那次算的 —— 两个数字对不上,但没有任何东西会告诉你。


🏷️ 三、key 的名字是拼出来的

命名规则只有一条:从根模块往下走,每一层的属性名用 . 连起来,最后接参数名。

import torch.nn as nn

a = nn.Sequential(nn.Linear(3, 4), nn.ReLU(), nn.Linear(4, 2))
print(list(a.state_dict().keys()))

输出 ['0.weight', '0.bias', '2.weight', '2.bias']。

⚠️ nn.Sequential 和 nn.ModuleList 里的子模块没有名字,只有下标。于是:

你改了什么 key 变成什么 后果
把 self.fc1 改名成 self.head fc1.weight → head.weight 旧 checkpoint 全部 missing
在 Sequential 中间插一层 2.weight → 3.weight 插入点之后的全错位
把 Sequential 换成显式写属性 数字 → 名字 全部对不上
用 DistributedDataParallel 包一层 全部多出 module. 前缀 全部 unexpected

⭐ 所以「绝对不要重编号」这条纪律不只适用于教程文件名,也适用于你的模型属性名。 真要改,写个映射把旧 key 翻译成新 key,比改模型便宜得多 —— state_dict 就是个字典,改 key 是一行推导式:

import torch.nn as nn

plain = nn.Linear(3, 2)
ddp_like = {"module." + k: v for k, v in plain.state_dict().items()}
print("DDP 存出来的 key:", list(ddp_like.keys()))

stripped = {k.removeprefix("module."): v for k, v in ddp_like.items()}
print("剥掉前缀后:", plain.load_state_dict(stripped))

实跑输出:

要点

DDP 存出来的 key: ['module.weight', 'module.bias']

剥掉前缀后: <All keys matched successfully>


🧯 四、加载报错的三种原文(都是实跑抄下来的)

load_state_dict 只会抛 RuntimeError,但信息里的关键词有三种,对应三个完全不同的问题。

① 形状不匹配

import torch.nn as nn

a = nn.Sequential(nn.Linear(3, 4), nn.ReLU(), nn.Linear(4, 2))
b = nn.Sequential(nn.Linear(3, 8), nn.ReLU(), nn.Linear(8, 2))
b.load_state_dict(a.state_dict())
RuntimeError: Error(s) in loading state_dict for Sequential:
    size mismatch for 0.weight: copying a param with shape torch.Size([4, 3]) from
    checkpoint, the shape in current model is torch.Size([8, 3]).

⭐ 读法:from checkpoint 是文件里的,in current model 是你现在这个模型的。 两个数字一比就知道是哪边的超参写错了 —— 上面这例是隐藏层宽度 4 和 8 对不上。

② 多了 key(Unexpected)

RuntimeError: Error(s) in loading state_dict for Sequential:
    Unexpected key(s) in state_dict: "2.weight", "2.bias".

checkpoint 里有、模型里没有。典型来源:DDP 前缀、模型被裁掉了几层、换了个更小的骨干。

③ 少了 key(Missing)

RuntimeError: Error(s) in loading state_dict for Sequential:
    Missing key(s) in state_dict: "2.weight", "2.bias".

模型里有、checkpoint 里没有。典型来源:迁移学习时新加了分类头、或者你忘了这层参数其实是 register_buffer 的。

⚠️ 两种可以同时出现。前面那个 DDP 例子如果直接加载,报的就是两条一起:

Error(s) in loading state_dict for Linear:
    Missing key(s) in state_dict: "weight", "bias".
    Unexpected key(s) in state_dict: "module.weight", "module.bias".

⭐ 看到 missing 和 unexpected 数量一样多、而且名字长得像,八成就是前缀问题,不是模型结构问题。


🛑 读到这里可以停 —— 前半章讲完了(约 25 分钟)。 后半章还有:strict=False 到底放过了什么 · 优化器也有 state_dict,而且 step 之前是空的 · torch.save 存的到底是什么 · 一份能真的恢复训练的 checkpoint 回来的时候不用重读,直接从下一节接着看就行。


📋 五、strict=False 到底放过了什么

大多数人对 strict=False 的印象是「忽略所有加载错误」。错。它只放过 key 的问题,不放过形状的问题。

import torch.nn as nn

a = nn.Sequential(nn.Linear(3, 4), nn.ReLU(), nn.Linear(4, 2))
c = nn.Sequential(nn.Linear(3, 4))
r = c.load_state_dict(a.state_dict(), strict=False)
print("返回:", r)
print("missing:", r.missing_keys, " unexpected:", r.unexpected_keys)

b = nn.Sequential(nn.Linear(3, 8), nn.ReLU(), nn.Linear(8, 2))
try:
    b.load_state_dict(a.state_dict(), strict=False)
except RuntimeError as e:
    print("形状不匹配照样报:", str(e)[:80])

实跑输出:

要点

返回: _IncompatibleKeys(missing_keys=[], unexpected_keys=['2.weight', '2.bias'])

missing: [] unexpected: ['2.weight', '2.bias']

形状不匹配照样报: Error(s) in loading state_dict for Sequential:

size mismatch for 0.weight

⭐ load_state_dict 有返回值,而几乎没人接。 它是一个 _IncompatibleKeys 具名元组, strict=False 之下唯一能告诉你「到底有多少层没加载上」的东西就是它。

💀 典型事故形态:微调时写 model.load_state_dict(ckpt, strict=False), 拼错了一个前缀 → 所有 key 都进 missing_keys → 一个权重都没加载,模型是随机初始化的, 而程序一声不吭地开始训练。指标看起来「就是不太好」,你会去调学习率,调三天。

✅ 写成这样,成本一行:

def load_checked(model, sd, allow_missing=()):
    r = model.load_state_dict(sd, strict=False)
    bad = [k for k in r.missing_keys if not k.startswith(tuple(allow_missing))]
    if bad:                                   # ⭐ 只允许你点名的那几层缺
        raise RuntimeError(f"这些层没加载上: {bad[:5]} ... 共 {len(bad)} 个")
    print(f"加载完成:跳过 {len(r.missing_keys)} 个,忽略 {len(r.unexpected_keys)} 个")
    return r

(微调时 allow_missing=("head.", "classifier.") —— 新加的分类头本来就该是随机的,其它任何一层缺失都是 bug。)


⚙️ 六、优化器也有 state_dict,而且 step 之前是空的

只存模型不存优化器,恢复训练时 loss 会跳一下 —— 因为动量和二阶矩全丢了。

import torch
import torch.nn as nn

m = nn.Linear(3, 2)
opt = torch.optim.Adam(m.parameters(), lr=1e-3)
print("step 之前 state =", opt.state_dict()["state"])

m(torch.randn(1, 3)).sum().backward()
opt.step()
st = opt.state_dict()["state"]
print("step 之后 state 的 key =", list(st.keys()))
print("每个参数存了:", list(st[0].keys()))

实跑输出:

要点

step 之前 state = {}

step 之后 state 的 key = [0, 1]

每个参数存了: ['step', 'exp_avg', 'exp_avg_sq']

三个观察:


💾 七、torch.save 存的到底是什么

torch.save 就是 pickle + 一个张量存储的旁路。存什么进去,就得有什么才能读出来。

① 存整个 model 对象 = 把类的引用路径腌进去

在 save_side.py 里定义 MyNet 并 torch.save(m, "mynet_whole.pt"), 换到 load_side.py(没有 MyNet 这个类)里加载:

AttributeError : Can't get attribute 'MyNet' on <module '__main__' from '...load_side.py'>

⚠️ 注意它说的是 __main__ —— pickle 存的是「__main__.MyNet」这个名字,不是类的代码。 所以下面每一件事都会让文件失效:改类名、把类挪到别的文件、把定义脚本改名、从别的入口跑。

② 存 state_dict = 只存纯张量

同一个脚本里 torch.save(m.state_dict(), "mynet_sd.pt"),在没有类定义的文件里加载:

要点

state_dict 加载成功: ['fc.weight', 'fc.bias']

体积也不一样(同一个 nn.Linear(3,2)):整模型 2517 字节,state_dict 1789 字节。 差的那部分就是类路径、模块结构这些元信息 —— 小模型上看着不多,但它带来的是一整条依赖。

③ ⭐ torch.load 现在默认 weights_only=True

这是 PyTorch 2.6 改的默认值,很多老教程还没跟上。实测(torch 2.13)加载整模型:

UnpicklingError : Weights only load failed. ...
    (1) In PyTorch 2.6, we changed the default value of the `weights_only` argument in
    `torch.load` from `False` to `True`. ...
    WeightsUnpickler error: Unsupported global: GLOBAL torch.nn.modules.linear.Linear
    was not an allowed global by default.

⭐ 这条报错其实是在帮你:pickle 反序列化等于执行文件里的代码,从别人那儿下载的 .pt 直接 torch.load 是真的能被拿去执行任意命令的。 weights_only=True 只允许纯张量通过 —— 而只存 state_dict 的人根本不会撞上这条报错。

你要做的事 怎么写
加载自己存的 state_dict torch.load(p) —— 默认就是安全模式
加载别人给的整模型 ⚠️ 先想清楚信不信任来源,再 weights_only=False
GPU 上存的、CPU 上读 torch.load(p, map_location="cpu") ⭐

⚠️ map_location 不加会怎样:checkpoint 里记着张量原来在 cuda:3, 在只有 2 张卡(或者没有卡)的机器上加载会直接报设备不存在。存的时候是什么设备,读的时候就会去找什么设备 —— 加 map_location="cpu" 一劳永逸。 (🗓️ 未实跑 —— 需要 GPU。本机 CUDA 不可用,无法触发这条报错原文。)


📦 八、一份能真的恢复训练的 checkpoint

「存了 checkpoint」和「能从 checkpoint 接着训」是两回事。要接着训,下面每一项少一样就会有肉眼可见的后果:

存什么 少了会怎样
model.state_dict() ——
⭐ optimizer.state_dict() 动量清零,恢复后前几十步 loss 明显跳一下
scheduler.state_dict() 学习率退回起点,等于又 warmup 了一遍
epoch / global_step 数据顺序对不上,日志曲线断成两截
⭐ 随机数状态 数据增强、Dropout 掩码不可复现
scaler.state_dict()(混合精度) 缩放系数重新摸索,前几步容易溢出成 NaN
import torch
import torch.nn as nn

model = nn.Linear(3, 2)
opt = torch.optim.Adam(model.parameters(), lr=1e-3)

ckpt = {
    "model": model.state_dict(),
    "optim": opt.state_dict(),
    "epoch": 7,
    "rng": torch.get_rng_state(),          # ⭐ 随机数状态也是个张量,能直接存
    "config": {"hidden": 2, "lr": 1e-3},   # ⭐ 把超参一起腌进去
}
torch.save(ckpt, "ckpt.pt")

blob = torch.load("ckpt.pt")
model.load_state_dict(blob["model"])
opt.load_state_dict(blob["optim"])
torch.set_rng_state(blob["rng"])
print("从第", blob["epoch"] + 1, "轮继续;配置 =", blob["config"])

实跑输出:从第 8 轮继续;配置 = {'hidden': 2, 'lr': 0.001}。

⭐ config 那一行是最容易被省掉、也最值钱的一项:没有它,半年后你拿到 ckpt.pt 得靠猜隐藏层宽度才能把模型构造出来 —— 而猜错的表现就是第四节那条 size mismatch。


🔗 这一章连到哪里

相关的地方 为什么
05 · nn.Module 内部 上一章讲清了谁会进 state_dict(参数和 buffer 进、普通属性不进);这一章讲进去之后长什么样、怎么存怎么取
07 · train() / eval() 改了什么 BN 的 running_mean 是 buffer,存在 state_dict 里 —— 下一章讲它是什么时候被改的,那正是「加载完模型指标却不对」的另一个来源
10 · 把模型交出去 这一章解决「自己能读回来」,那一章解决「交给别人也能读」:TorchScript / ONNX / 交付清单
《机器学习与深度学习基础》15 · PyTorch 实战手册 早停回滚那段模板的 .clone() 就在那里,为什么必须 clone 在本章第二节
《机器学习与深度学习基础》附录C 第 5 题 手写 BatchNorm 时为什么 running_mean 要用 register_buffer 而不是 Parameter —— 就是为了跟着 state_dict 存盘
《强化学习基础》08 · DQN 目标网络同步写的是 q_target.load_state_dict(q.state_dict());⭐ 那里为什么不需要 .clone()(因为 load_state_dict 是逐个 copy_ 进去的,本身就拷贝了)
《AI基础设施》21 · 训练稳定性与故障恢复 千卡训练怎么存 checkpoint:多久存一次、存哪儿、断点续训的工程账
《模型上线之后》17 · 版本回溯与可复现 第八节那个 config 字段的下一步:一个模型版本除了权重还该记什么

✅ 检查点

  1. state_dict() 返回的 value 是 Parameter 还是 Tensor?requires_grad 是多少?
  2. 为什么早停备份最佳权重必须写 .clone(),而 torch.save(model.state_dict(), path) 不用?
  3. 用什么一行代码能证明 state_dict 里的张量和模型参数是同一块内存?
  4. nn.Sequential(nn.Linear(3,4), nn.ReLU(), nn.Linear(4,2)) 的 state_dict 有哪几个 key?为什么没有 1.*?
  5. 加载报错里 from checkpoint 和 in current model 分别指哪一边?
  6. strict=False 放过什么、不放过什么?它的返回值是什么?
  7. 优化器 state_dict 里 state 的 key 是参数名吗?这带来什么风险?
  8. torch.save(model) 和 torch.save(model.state_dict()) 差在哪?前者在别的脚本里加载会报什么错?
  9. torch.load 从哪个版本开始默认 weights_only=True?为什么要改这个默认值?
👀 答案
  1. 是 Tensor(不是 Parameter),requires_grad 是 False —— state_dict() 出门前 detach 过。
  2. 因为 detach 过不等于拷贝过:state_dict 里的张量和参数共享同一块存储,权重一变备份跟着变。实测「没 clone 的备份」在权重从 1.0 改成 7.0 之后也变成 [7.0, 7.0],回滚等于没回滚。而 torch.save 落盘那一刻产生真正的快照,所以不用 clone。
  3. sd["weight"].data_ptr() == m.weight.data_ptr() —— 实测返回 True。
  4. ['0.weight', '0.bias', '2.weight', '2.bias']。没有 1.* 是因为下标 1 是 nn.ReLU,它一个参数都没有;Sequential 里的子模块只有下标没有名字,所以中间插一层会让后面全部错位。
  5. from checkpoint = 文件里存的形状;in current model = 你当前模型的形状。例:torch.Size([4, 3]) vs torch.Size([8, 3]) 说明隐藏层宽度 4 和 8 对不上。
  6. 放过 missing / unexpected key,不放过 size mismatch(照样抛 RuntimeError)。返回值是 _IncompatibleKeys(missing_keys=[...], unexpected_keys=[...]),几乎没人接,而它是 strict=False 下唯一能告诉你「有多少层根本没加载上」的东西。
  7. 不是,是整数 0, 1, ...,优化器只按 param_groups 里的顺序认参数。风险:恢复时构造优化器的参数顺序变了,动量会安到别的参数上,而且不报错。
  8. 前者 pickle 了类的引用路径(__main__.MyNet),换脚本加载报 AttributeError: Can't get attribute 'MyNet' on <module '__main__' ...>;改类名、挪文件、改文件名都会让它失效。后者只存纯张量,任何地方都能读。体积上同一个 nn.Linear(3,2):整模型 2517 字节 vs state_dict 1789 字节。
  9. PyTorch 2.6。因为 pickle 反序列化等于执行文件里的代码,下载来的 .pt 直接 torch.load 可以被拿去执行任意命令;weights_only=True 只放纯张量过。只存 state_dict 的人根本撞不上这条报错。

🛑 可以停在这里

⚡ 走神救援

先记住这几件事

下一节 👉 07-train和eval改了什么.md

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