📑 本页目录(点开跳转)
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']
三个观察:
- ⭐
state在第一次step()之前是空的。所以「刚建好优化器就存 checkpoint」存了个寂寞。 - ⭐ key 是整数
0, 1,不是参数名。优化器根本不知道参数叫什么,它只按param_groups里的顺序认人。 ⚠️ 于是:恢复训练时构造优化器的参数顺序必须和存的时候完全一致,否则动量会被安到别的参数头上,而且不报错。 exp_avg/exp_avg_sq是 Adam 的一阶/二阶矩,每个都和参数一样大 —— 这就是「Adam 的优化器状态是参数的 2 倍」这句话的出处(显存账在AI基础设施讲)。
💾 七、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 字段的下一步:一个模型版本除了权重还该记什么 |
✅ 检查点
state_dict()返回的 value 是Parameter还是Tensor?requires_grad是多少?- 为什么早停备份最佳权重必须写
.clone(),而torch.save(model.state_dict(), path)不用? - 用什么一行代码能证明
state_dict里的张量和模型参数是同一块内存? nn.Sequential(nn.Linear(3,4), nn.ReLU(), nn.Linear(4,2))的state_dict有哪几个 key?为什么没有1.*?- 加载报错里
from checkpoint和in current model分别指哪一边? strict=False放过什么、不放过什么?它的返回值是什么?- 优化器
state_dict里state的 key 是参数名吗?这带来什么风险? torch.save(model)和torch.save(model.state_dict())差在哪?前者在别的脚本里加载会报什么错?torch.load从哪个版本开始默认weights_only=True?为什么要改这个默认值?
👀 答案
- 是
Tensor(不是Parameter),requires_grad是False——state_dict()出门前 detach 过。 - 因为 detach 过不等于拷贝过:
state_dict里的张量和参数共享同一块存储,权重一变备份跟着变。实测「没 clone 的备份」在权重从 1.0 改成 7.0 之后也变成[7.0, 7.0],回滚等于没回滚。而torch.save落盘那一刻产生真正的快照,所以不用 clone。 sd["weight"].data_ptr() == m.weight.data_ptr()—— 实测返回True。['0.weight', '0.bias', '2.weight', '2.bias']。没有1.*是因为下标 1 是nn.ReLU,它一个参数都没有;Sequential里的子模块只有下标没有名字,所以中间插一层会让后面全部错位。from checkpoint= 文件里存的形状;in current model= 你当前模型的形状。例:torch.Size([4, 3])vstorch.Size([8, 3])说明隐藏层宽度 4 和 8 对不上。- 放过 missing / unexpected key,不放过
size mismatch(照样抛RuntimeError)。返回值是_IncompatibleKeys(missing_keys=[...], unexpected_keys=[...]),几乎没人接,而它是strict=False下唯一能告诉你「有多少层根本没加载上」的东西。 - 不是,是整数
0, 1, ...,优化器只按param_groups里的顺序认参数。风险:恢复时构造优化器的参数顺序变了,动量会安到别的参数上,而且不报错。 - 前者 pickle 了类的引用路径(
__main__.MyNet),换脚本加载报AttributeError: Can't get attribute 'MyNet' on <module '__main__' ...>;改类名、挪文件、改文件名都会让它失效。后者只存纯张量,任何地方都能读。体积上同一个nn.Linear(3,2):整模型 2517 字节 vsstate_dict1789 字节。 - PyTorch 2.6。因为 pickle 反序列化等于执行文件里的代码,下载来的
.pt直接torch.load可以被拿去执行任意命令;weights_only=True只放纯张量过。只存state_dict的人根本撞不上这条报错。
🛑 可以停在这里
⚡ 走神救援
先记住这几件事
- state_dict 按模块路径组织参数和缓冲区,值并不是独立的模型副本。
- detach 不等于复制;继续训练时,直接留下的状态字典可能跟着变化。
- 内存中保存最佳权重时复制张量;存取后核对键、设备与恢复出的结果。
下一节 👉 07-train和eval改了什么.md