ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

手撕 .pt:从 ZIP 魔数到 Pickle 协议,彻底搞懂 PyTorch 序列化

手撕 .pt:从 ZIP 魔数到 Pickle 协议,彻底搞懂 PyTorch 序列化 引言我们欠 .pt 文件一次“解剖”在深度学习的日常开发中.pt或.pth模型文件是我们最熟悉的陌生人。我们习惯了torch.save()和torch.load()这对组合拳但当报错信息铺满终端时这个文件的内部结构依然是一个黑盒。这篇文章带你亲手拆开它。我们不满足于知道“怎么用”而是追问三个核心问题它到底是什么格式—— 打开文件头看看不就知道了里面装了什么—— 反序列化后逐层展开。为什么有些做法是错的有些是对的—— 复现错误亲眼见证崩溃和修复的全过程。第一章、撕开外壳.pt 文件到底是什么格式保存一个样本用最原始的方式读取它先存一个最简单的张量这大概是我们能构建的最小.pt文件了。import torch torch.save(torch.tensor([1, 2, 3, 4, 5]), sample.pt)现在我们像读任何普通文件一样以二进制模式打开它只看前 20 个字节。with open(sample.pt, rb) as f: raw f.read(20) print(f原始字节: {raw}) print(f十六进制: {raw.hex()})运行后终端会打印类似这样的内容看到开头的PK了吗50 4B是 ZIP 文件的通用魔数来源于 PKZIP 软件。也就是说我们保存的.pt文件本质上是一个 ZIP 压缩包根本不是某种私有的二进制格式。为了进一步确认直接用 Python 的zipfile模块把它当作压缩包打开python import zipfile with zipfile.ZipFile(sample.pt, r) as z: print(z.namelist())输出注意所有文件都带上了sample/前缀data.pkl被放在了sample/子目录下。这是 PyTorch 新版 ZIP 序列化格式的特征自 1.6 版本起默认启用——它将每个张量单独存储为二进制文件如sample/data/0并用data.pkl描述对象结构从而实现大模型的分块加载和零拷贝读取避免把几百 GB 的模型一次性全部解压到内存中。但无论结构如何变化核心目标不变找到那个data.pkl它才是整个序列化的入口。读取 data.pkl验证 Pickle 协议用zipfile找到data.pkl并读取其文件头import zipfile with zipfile.ZipFile(sample.pt, r) as z: pkl_candidates [name for name in z.namelist() if name.endswith(data.pkl)] if pkl_candidates: with z.open(pkl_candidates[0]) as f: header f.read(10) print(fdata.pkl 路径: {pkl_candidates[0]}) print(f文件头: {header}) print(f十六进制: {header.hex()}) else: print(未找到 data.pkl)输出这短短 10 个字节揭示了三个关键信息字节含义\x80\x02Pickle Protocol 2— Python 序列化协议版本 2ctorch._Pickle 反序列化时需要加载torch._模块PyTorch 的 C 扩展这是张量重建的标志至此铁证如山PyTorch 的.pt ZIP 外壳 Pickle 内核。如果你看到的是b\x80\x04\x95...那是 Pickle Protocol 4Python 3.4 默认。不同的协议版本不影响“它是 Pickle 序列化”这一结论只说明序列化时的 Python 版本和参数差异。第二章、解剖内部.pt 文件里到底装了什么知道了它是 ZIP Pickle但 Pickle 反序列化之后的东西长什么样我们构造一个真实的训练场景来看看。保存一个完整的 checkpoint看它还原成什么定义一个两层的小网络配上优化器打个包。import torch import torch.nn as nn import torch.optim as optim class TinyNet(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(10, 5) model TinyNet() optimizer optim.SGD(model.parameters(), lr0.01) checkpoint { epoch: 10, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), loss: 0.23, lr: 0.01 } torch.save(checkpoint, checkpoint.pt)现在加载回来看看torch.load()还原出了什么类型的对象import torch loaded torch.load(checkpoint.pt) print(type(loaded)) print(loaded.keys())输出它就是一个普普通通的 Python 字典。字典里塞着整数、浮点数还有嵌套的权重字典和优化器状态字典。把权重的名字和形状打出来看看import torch loaded torch.load(checkpoint.pt) for name, param in loaded[model_state_dict].items(): print(f{name} - {param.shape})输出看到这里你可能已经隐约理解了为什么加载时需要先定义模型结构state_dict里存的是“参数名 → 张量”的映射它依赖模型有同名、同形状的层来接收这些参数。这个在下一章会亲眼验证。第三章、三种保存方式的对比自己“踩坑”才最深刻PyTorch 社区一直强调用state_dict而不是保存整个模型。亲手复现一遍就全明白了。保存完整模型PyTorch 2.6 中的加载变化把整个模型对象存下来import torch import torch.nn as nn class MyNet(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(2, 2) net MyNet() torch.save(net, whole_model.pt)在 PyTorch 2.6 及以上版本中直接使用torch.load()加载会失败import torch try: model torch.load(whole_model.pt) except Exception as e: print(f报错{e})终端会打印类似这样的错误信息_pickle.UnpicklingError: Weights only load failed. This file can still be loaded, to do so you have two options, do those steps only if you trust the source of the checkpoint. (1) In PyTorch 2.6, we changed the default value of the weights_only argument in torch.load from False to True. Re-running torch.load with weights_only set to False will likely succeed...场景二显式设置weights_onlyFalse但缺少类定义即使绕过安全限制如果当前环境中没有定义MyNet类加载仍然会失败。下面的代码单独运行时一定会报错import torch # 注意如果当前环境没有定义 MyNet 类这行代码会报错 model torch.load(whole_model.pt, weights_onlyFalse)这是因为whole_model.pt中存储了类的引用路径__main__.MyNetPickle 在反序列化时需要在当前环境中找到这个类的定义。要成功加载需要满足两个条件之一方法一先定义类再加载必须与保存时的类定义一致import torch import torch.nn as nn # 先定义类必须与保存时的类定义完全一致 class MyNet(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(2, 2) # 再加载 model torch.load(whole_model.pt, weights_onlyFalse) print(加载成功)方式二使用add_safe_globals将自定义类加入白名单理论上可行实践中不推荐错误信息中提到了add_safe_globals方案。但在实际操作中会遇到更复杂的情况import torch import torch.nn as nn class MyNet(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(2, 2) # 将 MyNet 加入白名单 torch.serialization.add_safe_globals([MyNet]) try: model torch.load(whole_model.pt) # weights_only 默认为 True print(加载成功) except Exception as e: print(f报错{e})你会看到WeightsUnpickler error: Unsupported global: GLOBAL torch.nn.modules.linear.Linear was not an allowed global by default. Please use torch.serialization.add_safe_globals ([torch.nn.modules.linear.Linear]) ...错误信息表明不仅自定义类MyNet需要加入白名单模型中用到的所有 PyTorch 内部类如Linear也需要逐一加入白名单。对于一个真实的深度模型这几乎是不可能的任务因为你无法预知模型中用到的所有底层 PyTorch 类。因此add_safe_globals的适用场景实际上是适合为自定义的简单数据类如配置类MyConfig开启weights_onlyTrue加载权限不适合为包含复杂 PyTorch 网络结构的完整模型对象开启加载权限方法三推荐使用state_dict保存权重这是 PyTorch 官方推荐的最佳实践可以完全避免上述两类问题import torch import torch.nn as nn class MyNet(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(2, 2) net MyNet() # 只保存权重字典不保存整个模型对象 torch.save(net.state_dict(), weights_only.pt) # 加载时先实例化模型再加载权重 model MyNet() model.load_state_dict(torch.load(weights_only.pt)) # PyTorch 2.6 默认 weights_onlyTrue print(加载成功)保存 state_dict换一个类也能加载这是state_dict方式最大的优势——不依赖类定义。import torch import torch.nn as nn # 定义一个模型类 class MyNet(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(2, 2) net MyNet() # 只保存权重字典 torch.save(net.state_dict(), weights_only.pt) # 定义一个名字不同但结构相同的类 class AnotherNet(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(2, 2) # 加载权重到新模型 weights torch.load(weights_only.pt) new_net AnotherNet() new_net.load_state_dict(weights) print(加载成功)输出因为state_dict只是一个字典不依赖任何类定义。你只需要保证目标模型有相同的层名和形状就能把权重灌进去。断点续训的 checkpoint恢复训练现场如果你训练了一个大模型跑了三天结果断电了靠什么救回来答案是保存优化器状态import torch import torch.nn as nn import torch.optim as optim class TinyNet(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(10, 5) model TinyNet() optimizer optim.SGD(model.parameters(), lr0.01) checkpoint { epoch: 10, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), best_loss: 0.23 } torch.save(checkpoint, resume.pt) # 恢复时 ckpt torch.load(resume.pt) model.load_state_dict(ckpt[model_state_dict]) optimizer.load_state_dict(ckpt[optimizer_state_dict]) start_epoch ckpt[epoch] 1 print(f从第 {start_epoch} 个 epoch 继续训练)输出Adam 或 SGD 的动量状态被完整保留和中断前的训练状态丝滑衔接。这是 checkpoint 方式独有的价值。第四章、加载前看一眼不占显存的“无损检测”生产环境中面对一个几十 GB 的.pt文件我们不能贸然torch.load()到 GPU 里。万一设备不对、版本不兼容或者只是想看一眼参数量有没有更轻量的办法map_location 解决 GPU/CPU 不匹配假设模型是在 GPU 上保存的。在 CPU 机器上加载时不加map_location会报错import torch import torch.nn as nn class TinyNet(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(10, 5) model TinyNet() torch.save(model.state_dict(), gpu_weights.pt) # 在 CPU 机器上加载不加 map_location 会报错 try: data torch.load(gpu_weights.pt) except RuntimeError as e: print(f报错{e})报错信息会提示 CUDA 不可用。加上map_locationcpu就解决了import torch data torch.load(gpu_weights.pt, map_locationtorch.device(cpu)) print(data[fc.weight].device)输出cpumap_location的原理是在 Pickle 反序列化时重写张量的存储位置钩子把原本指向 GPU 内存的指针换成 CPU 内存。这个机制是 PyTorch 专门为跨设备加载设计的。weights_onlyTrue 安全加载自 PyTorch 1.10 起torch.load增加了weights_only参数限制反序列化时只允许张量、字典、列表等安全类型。在 PyTorch 2.6 中这个参数默认值已变为Trueimport torch import torch.nn as nn class TinyNet(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(10, 5) model TinyNet() torch.save(model.state_dict(), weights_only.pt) # PyTorch 2.6 中 weights_only 默认为 True safe_data torch.load(weights_only.pt) print(type(safe_data)) total_params sum(p.numel() for p in safe_data.values()) print(f参数量: {total_params})输出第五章、安全红线亲眼见证 Pickle 的“代码执行”漏洞很多文章会提一句“不要加载未知来源的.pt文件因为它可能执行恶意代码”。这句话值得用代码亲自验证——毕竟“听人劝”远不如“亲眼见”来得震撼。Pickle 允许对象通过__reduce__方法自定义反序列化时的行为。我们构造一个无害但足够说明问题的例子import pickle import torch class Exploit: def __reduce__(self): return (print, (⚠️ 这段代码在 torch.load() 时被执行了,)) with open(exploit.pt, wb) as f: pickle.dump(Exploit(), f) print(即将执行 torch.load(exploit.pt)请注意终端输出) # 必须显式设置 weights_onlyFalse否则 PyTorch 2.6 会拒绝加载 torch.load(exploit.pt, weights_onlyFalse)运行后你会在终端看到在torch.load完成返回之前print就被触发了即将执行 torch.load(exploit.pt)请注意终端输出 ⚠️ 这段代码在 torch.load() 时被执行了这就是所谓的反序列化漏洞。如果__reduce__返回的是(os.system, (rm -rf /,))后果可想而知。.pt文件不是数据文件它是可执行的代码包。这也解释了为什么 PyTorch 2.6 要将weights_only默认值改为True——用默认的安全行为来保护大多数用户。安全最佳实践优先使用safetensors格式Hugging Face 推出只存张量不存代码如果必须使用.pt尽量只分发state_dict权重字典而非完整模型对象加载来源不明的文件前在隔离环境如 Docker 沙箱中执行第六章、部署的真相TorchScript 生成的 .pt 有何不同生产环境部署时我们常常用torch.jit.trace()把模型“编译”成 TorchScript然后保存为.pt文件。它和普通.pt是同一个东西吗import torch import torch.nn as nn import zipfile class TinyNet(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(10, 5) def forward(self, x): return self.fc(x) net TinyNet() scripted torch.jit.trace(net, torch.randn(1, 10)) torch.jit.save(scripted, scripted.pt) with zipfile.ZipFile(scripted.pt, r) as z: print(z.namelist())输出取决于版本可能略有不同但核心特征一致[scripted/data/0, scripted/data/1, scripted/data.pkl, scripted/code/__torch__.py, scripted/code/__torch__.py.debug_pkl, scripted/code/__torch__/torch/nn/modules/linear.py, scripted/code/__torch__/torch/nn/modules/linear.py.debug_pkl, scripted/constants.pkl, scripted/traced_inputs/0, scripted/traced_inputs.pkl, scripted/version, scripted/byteorder, scripted/.data/serialization_id]看到了吗多了一个code/目录。TorchScript 文件在data.pkl之外把模型的计算图逻辑以中间表示的形式打包了进去。这就是为什么 TorchScript 模型可以脱离原始的 Python 类定义直接在 C 运行时中执行。这同时也解释了为什么有些.pt文件看起来“特别大”——它存的不仅是权重还有一份完整的、跨平台的计算图代码。总结用代码验证过的才算真正理解整篇文章我们每一处结论都对应了一段可以亲手运行的代码。以下是核心结论汇总结论对应验证方式.pt本质是 ZIP 文件open().read(20)看到PK魔数ZIP 内部藏着data.pklzipfile.namelist()endswith(data.pkl)data.pkl遵循 Pickle 协议读取文件头识别\x80\x02或\x80\x04新版本 PyTorch 将张量分块存储ZIP 内出现data/0、data/1等文件torch.load还原为 Python 对象type(torch.load())确认类型state_dict是层名→张量的映射遍历.items()打印PyTorch 2.6weights_only默认为True加载完整模型报UnpicklingError加载完整模型需要weights_onlyFalse 类定义先定义类再显式设置参数add_safe_globals不适用于复杂网络结构加载报错提示需添加Linear等内部类state_dict不依赖类定义换一个类仍可加载Checkpoint 能恢复优化器状态optimizer.load_state_dict()map_location解决跨设备加载加与不加对比报错信息torch.load存在代码执行风险__reduce__在加载时触发TorchScript 的.pt包含计算图代码ZIP 内部出现code/目录最终实践建议每一句都有上述代码为依据仅提供参考场景推荐做法日常实验使用model.state_dict()保存权重torch.load()默认即可加载长期训练保存 checkpoint 字典包含epoch、optimizer_state_dict、loss模型分发优先使用safetensors格式如用.pt只分发state_dict权重字典而非完整模型部署上线使用torch.jit.trace()生成 TorchScript 模型实现 Python 环境解耦来源不明文件在隔离环境Docker 沙箱中加载或根本不加载现在你拥有了亲手解剖.pt文件的能力。以后遇到任何加载报错记住它的本质——一个穿了 ZIP 马甲的 Pickle 归档——然后顺着报错栈一路追踪下去真相就在那里。
返回列表