尧图精选

PyTorch TorchScript 序列化格式完全解析:从 `.pt` 归档结构到 `save`/`load` 内部机制

🕒 发布时间:2026/9/12 4:58:09 📁 来源:尧图网络
PyTorch TorchScript 序列化格式完全解析从.pt归档结构到save/load内部机制【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch导读本文以 PyTorch 仓库中 torch/csrc/jit/docs/serialization.md 为骨架系统讲解 TorchScript 模型的序列化格式与torch::jit::save()/torch::jit::load()的完整调用路径。你将掌握.pt文件作为一个 ZIP 归档的内部目录结构、code/中 Python 源码的生成规则PythonPrint、data.pkl与data/中对象状态与张量的分离存储、constants.pkl解决代码先于数据加载的循环依赖问题、__getstate__/__setstate__的自定义状态控制以及CompilationUnit的代码对象所有权与命名mangling语义。读完本文你可以直接unzip打开任意 TorchScript 模型进行人工排障也能理解序列化相关 C 源码的实现动机。概览.pt是一个模仿 Python 生态的 ZIP 归档一个序列化后的模型通常叫model.pt本质是一个ZIP 归档里面包含多个文件。你可以直接对其执行unzip来查看内部结构$ unzip model.pt Archive: model.pt extracting ... $ tree model/ ├── code/ │ ├── __torch__.py │ ├── __torch__.py.debug_pkl │ ├── foo/ │ │ ├── bar.py │ │ ├── bar.py.debug_pkl ├── data.pkl ├── constants.pkl └── data/ ├── 0 └── 1注意到归档中同时存在.py和.pkl文件——这正是序列化格式的设计精髓它刻意模仿 Python。所有代码类信息方法、模块、类、函数以人类可读、且语法合法的.pyPython 源码存储所有数据类信息属性、对象等用Python pickle 协议的一个子集进行序列化。一个模型本质上是一个顶层模块附带若干子模块、参数等。因此data.pkl中保存的是被 pickle 的顶层模块反序列化模型时只需对data.pkl调用unpickle()即可恢复模块对象状态并按需加载其关联代码。设计原则Design Notes仓库文档明确列出了四条设计准则是理解后续所有实现决策的出发点Do what Python does跟随 Python 的行为所有序列化代码必须是合法 Python所有 pickle 对象应能被 Python 反 pickle。这保证了格式的可调试性你可以随时拆开归档用 Python 把玩以及对熟悉 Python 但不熟悉 TorchScript 的开发者更友好。Human readable人类可读序列化出的代码要尽量可读保留作者书写的变量名、适当内联短表达式方便对序列化代码进行调试。No jitter无抖动执行保存→加载→再保存→再加载后两次加载结果应当完全一致。该性质有助于捕捉序列化过程中的 bug避免模型因保存/加载次数不同而产生漂移m MyModule() m.save(foo.pt) m_loaded torch.load(foo.pt) m_loaded.save(foo2.pt) m_loaded2 torch.load(foo2.pt)Initial load should be fast首次加载要快load()对人而言应近乎瞬时任何耗时操作如读取张量数据都应惰性完成。code/代码如何被序列化代码序列化在高层分为两步先把ClassType与Function统称 code object转换成 Python 源码再把源码写入模型 ZIP 归档的code/目录。用PythonPrint把代码对象打印成 Python 源码PythonPrint是完成代码对象 → Python 源码转换的核心函数其声明位于 torch/csrc/jit/serialization/python_print.h。它以ClassType或Function为输入输出 Python 源码。ScriptModule以类类型class type实现因此其方法与属性同样会被序列化。PythonPrint的实现方式是遍历GraphClassType方法或原始Function的 IR 表示并逐节点发射对应的 Python 代码核心逻辑位于 python_print.cpp。除逐条语句的简单规则外它还会额外跟踪以下几类信息类依赖Class dependencies遍历图时记录图中使用到的所有类加入当前代码对象的依赖列表。例如打印一个Module时它会依赖其子模块以及方法、属性中用到的所有类。张量常量引用Uses of tensor constants大多数常量字符串、整数等会作为字面量直接内联。但张量可能非常大因此PythonPrint遇到常量张量时不会内联而是发射对全局CONSTANTS表的引用形如foo CONSTANTS.c0。源码中的实际发射点见 python_print.cppss CONSTANTS.c getOrAddConstant(v);。在导入阶段importer 通过查张量表把该引用解析为真实张量CONSTANTS.c0表示constants.pkl中张量元组的第 0 个张量。该解析逻辑由 import_source.cpp 中的ConstantTableValue完成——它把CONSTANTS.cN中的下标N解析为常量表向量中的偏移量并插入为图常量同时还会做类型去特殊化unshapedType避免张量类型特化破坏类型关系。原始源码区间记录Original source range records为辅助调试PythonPrint会记住所发射代码的原始用户书写位置。这样用户调试加载回来的模型时诊断信息会指向其真正书写的代码而不是PythonPrint发射的代码。这些记录被 pickle 后存入与代码同名的.debug_pkl文件可视为序列化代码源码区间 ↔ 原始用户代码的映射。相关生成与读取路径见 export_module.cpp 与 import_export_helpers.cpp。模块信息Module information模块有几处特殊处理Parameter标记部分模块属性实为Parameter具有特殊性质。为跟踪哪些属性是参数PythonPrint会在类体中发射特殊赋值语句class MyModule(Module): __parameters__ [foo, bar, ] foo : Tensor bar : Tensor attribute_but_not_param : Tensor属性枚举模块通常在 Python 中构造而__init__()方法并不会被编译。为保证静态类型PythonPrint必须显式枚举模块的全部属性如上面代码所示而不能依赖编译__init__()来推断属性。非合法标识符属性nn.Sequential等模块可能含有不是合法 Python 标识符的属性名。下面的写法是非法的# wrong! class MyModule(Module): 0 : ASubmodule 1 : BSubmodule虽然 Python 允许这种属性名但这不是合法 Python 语法。解决办法是直接写入__annotations__字典class MyModule(Module): __annotations__ [] __annotations__[0] ASubmodule __annotations__[1] ASubmodule把源码放进归档所有代码对象被PythonPrint打印为源码字符串后需要确定其存放位置。这涉及两个关键概念CompilationUnit与给定模型关联的所有代码对象的拥有者容器。加载时所有代码对象被载入到同一个CompilationUnit。QualifiedName代码对象的全限定名类似 Python 的限定名形如foo.bar.baz。在同一个CompilationUnit内每个代码对象拥有唯一的QualifiedName。导出器依据代码对象的QualifiedName决定其在code/目录中的位置方式与 Python 类似例如QualifiedName为foo.bar.Baz的类Baz会被放在code/foo/bar.py中名为Baz。位于层级根部的类会被加上__torch__前缀作为限定名以便放进__torch__.py。为什么不叫__main__因为 pickle 对__main__中的对象有特殊规则。此外还有一层额外逻辑在单个文件内类会按逆依赖序排放从而保证叶子依赖先于依赖者被编译。数据如何被序列化模型本质是带任意数量子模块、参数、属性的顶层ScriptModule。我们实现了 pickle 格式中序列化模块对象所需的子集。选择 pickle 格式的原因文档明确列出用户友好属性文件可在 Python 中直接用pickle加载无整体大小限制Protobuf 等格式对总消息大小有限制而 pickle 的限制仅针对单个值例如字符串不能超过 4 GB标准格式pickle 是 Python 标准模块格式合理简单本质是供栈式虚拟机消费的一段程序内置 memoization支持共享引用类型Tensor、字符串、列表、字典的复用自描述理解被 pickle 的数据无需额外定义文件与 eager 模式一致torch.save()本身产出的就是 pickle 归档用同样方式保存属性可避免再引入一种格式。data.pkl模块对象状态如何被序列化除张量见下节外所有数据都写入data.pkl。数据指模块对象状态的全部组成部分如属性、子模块等。在底层pickle 写入由 pickler.cpp / pickle.cpp 完成pickle.cpp中调用writeArchiveAndTensors(data, ...)将 pickle 字节与张量分别写入data.pkl与data/0、data/1、data/2…… 每个张量一个文件见 pickle.cpp。PyTorch 在 torch/jit/_pickle.py 中定义了一组函数用于标记特殊数据类型如张量表索引、特殊列表类型build_intlist、build_tensorlist、build_doublelist、build_boollist、build_tensor_from_id以及restore_type_tag。其中restore_type_tag用于把列表等容器类型的完整静态类型在重新加载时恢复旧版build_*函数仅保留作向后兼容。data/张量如何被序列化导出期间会构建模型中所有张量的列表张量可来自模块参数或Tensor类型的属性。张量之所以与其它数据走标准 pickle 流程不同处理原因有三文档明确列出张量经常超过 pickle 的文件大小限制希望直接mmap张量避免一次性读入内存希望保持与常规 PyTorch 序列化格式的兼容性。constants.pkl代码中的常量pickle 格式强制代码与数据分离TorchScript 序列化通过code/与data.pkl tensors/来体现这一分离。然而 TorchScript 会把常量即prim::Constant节点直接内联进code/。这对张量常量构成了问题张量无法方便地以字符串形式表示。同时张量常量不能放进data.pkl——因为源码必须在data.pkl之前加载把张量常量放进data.pkl会形成循环加载依赖。解决方案是创建独立的 pickle 文件constants.pkl存放代码中引用的全部张量常量。加载顺序在下一节说明。加载端对应实现见 import.cpp先readArchive(constants)得到张量常量元组再逐个 push 进constants_table_。torch::jit::load()的加载流程加载过程分两步反 pickleconstants.pkl得到代码中引用的全部张量常量构成的元组反 pickledata.pkl为顶层Module并返回。反 pickle 过程本质上是对data.pkl中模块对象的单次 unpickle 调用。Unpickler被赋予一个回调用于把遇到的任何限定名解析为ClassType先把限定名解析到code/中对应的文件再编译该文件并返回相应ClassType。这正是给代码对象在CompilationUnit中分配唯一限定名至关重要的原因——这样Unpickler遇到的每个类在code/中都有确定性的存放位置。Unpickler还负责把张量引用解析为真实的at::Tensor在反 pickle 过程中通过查张量表偏移实现文档注明该机制即将替换为与其它数据一致的 pickle 策略。从源码结构看加载端核心是ScriptModuleDeserializer见 import.cpp它持有CompilationUnit与PyTorchStreamReader并使用code/前缀定位代码归档ObjLoaderFuncWithVersion同文件 import.cpp则负责把 pickle 状态还原为对象若类实现了__setstate__则调用之否则把字典状态按属性名逐一写入对象槽位。__getstate__与__setstate__与 Python 的pickle一致用户可通过实现__getstate__()与__setstate__()方法自定义类或模块的 pickling 行为。Pickler与Unpickler会透明地处理对这两个方法的调用序列化过程本身无需过多操心。值得指出的是编译器实现了一些特殊的类型推断行为以弥补用户目前无法为Module添加类型注解的缺陷__getstate__和__setstate__不需要类型注解对__getstate__编译器可完全根据用户返回的属性推断返回类型对__setstate__编译器直接查__getstate__的返回类型作为其输入类型。例如class M(torch.nn.Module): def __init__(self) - None: self.a torch.rand(2, 3) self.b torch.nn.Linear(10, 10) def __getstate__(self): # Compiler infers that this is a tuple of (Tensor, Linear) return (self.a, self.b) def __setstate__(self, state): # Dont need to annotate this, we know what type state is! self.a state[0] self.b state[1]从源码实现看加载端会在类存在合法的__setstate__时走专门分支创建一个空对象然后关闭图优化器GraphOptimizerEnabledGuard guard(false)再调用__setstate__避免在类尚未初始化前就尝试特化调用后还会执行postSetStateValidate校验所有非Optional/None/Union类型的属性都已被初始化否则报错提示 The field {} was left uninitialized after setstate见 import.cpp。附录CompilationUnit与代码对象所有权CompilationUnit承担两项职能在 C 意义上拥有全部代码对象构成一个代码对象名称必须唯一的命名空间。每次调用torch::jit::load()时都会创建一个CompilationUnit来存放新反序列化的代码对象在 Python 中则存在一个持有 Python 内定义的所有代码对象的全局CompilationUnit。所有权语义参与所有权模型的实体如下CompilationUnit拥有代码对象并为其命名的容器每个代码对象在其中的限定名唯一。Function带执行器的Graph。Graph可能拥有ClassType因为某些Value目前持有指向其类型的shared_ptr也可能通过函数调用弱引用其它Function。ClassType类型的定义可指用户自定义 TorchScript 类或ScriptModule拥有其属性类型包括其它ClassType弱引用类的方法Function。Object特定类的实例拥有持有其ClassType的CompilationUnit。这保证了用户在 C 中传递对象时其全部代码持续存活、方法可被调用。Module对ClassType与其状态Object的视图负责把非限定名如forward()转换为限定名如__torch__.MyModule.forward以在所属CompilationUnit中查找拥有Object进而传递性地拥有CompilationUnit。Method(Module, Function)二元组。代码对象命名CompilationUnit维护一个所有代码对象ClassType与Function唯一命名的命名空间。这些名字本身没有特别含义只在序列化/反序列化期间唯一标识代码对象。基本命名方案一切从__torch__命名空间开始类命名与 Python 的模块命名平行foo.py中的类Bar变为__torch__.foo.Bar方法挂在模块的命名空间下Bar.forward()变为__torch__.foo.Bar.forward。存在几点注意事项无前缀的CompilationUnit出于测试及内部目的偶尔需要名字不带前缀此时一切只是CompilationUnit内的裸名。用户无法构造这种形态的CompilationUnit。名称 mangling名字重整在 Python 中可能构造出限定名相同的代码对象主要有两种情况每个ScriptModule在 JIT 中都是单例类用户构造多个ScriptModule会生成多个同名ClassType受 Python 限制嵌套函数也会导致限定名冲突。这些情况下代码对象在放入全局 PythonCompilationUnit前会被 mangling。规则很简单以限定名__torch__.foo.Bar为例__torch__.foo.Bar # first time, unchanged __torch__.foo.__torch_mangle_0.Bar # second time, when we request a mangle __torch__.foo.__torch_mangle_1.Bar # and so on注意 mangling 作用在Bar之前的命名空间段上——这样 pretty-print 代码时非限定名Bar保持不变该性质使 trace-checking 等对 mangling 无感知从而保持稳定。实战建议手动拆包排障综合全文当你面对一个无法加载或行为异常的.pt文件时可按如下思路排查unzip model.pt查看归档是否完整、code/、data.pkl、constants.pkl、data/是否齐备检查code/__torch__.py等源码是否可读、__parameters__与__annotations__是否与预期一致CONSTANTS.cN引用是否越界加载端会在越界时报 constant index is out of bounds见 import_source.cpp若对象状态异常检查是否实现了__getstate__/__setstate__以及__setstate__是否初始化了全部非可选属性关注.debug_pkl是否存在它决定了加载后诊断信息能否映射回原始用户代码。通过理解上述归档结构与加载顺序你可以把序列化/反序列化黑盒变成完全可控、可调试的过程这也正是该格式跟随 Python、人类可读、无抖动、快速加载四大设计原则的价值所在。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
上一篇/下一篇内容由系统自动关联 返回资讯列表 →