# Checkpoint 与 ONNX 详解:区别、关系、转换方式与原理
一、核心结论速览
| 维度 | Checkpoint (.pth) | ONNX (.onnx) |
|---|---|---|
| 本质 | 权重的序列化文件(多为 Python pickle) | 计算图 + 权重的跨框架标准格式 |
| 结构从哪来 | 依赖 Python 代码(如 nn.Module 定义) | 图结构写在 ONNX 文件内 |
| 运行依赖 | 需要 PyTorch + 定义模型的代码 | 只需 ONNX Runtime 等,无需 PyTorch |
| 典型用途 | 训练、断点恢复、在 PyTorch 内推理 | 部署、跨平台、量化、嵌入式/服务端推理 |
关系:Checkpoint 是「权重的来源」;ONNX 是「图 + 权重的导出结果」。转换时,先用 checkpoint 把权重加载进 PyTorch 模型,再通过一次前向追踪把「图 + 当前权重」一起写入 ONNX。
二、Checkpoint(.pth)是什么
2.1 定义与内容
Checkpoint 是 PyTorch 中常用的持久化格式,由 torch.save() 写出、torch.load() 读入。文件扩展名多为 .pth 或 .pt,本质是 Python pickle 序列化后的字节流。
典型保存内容:
- 仅 state_dict:
torch.save(model.state_dict(), "model.pth")- 内容是一个 Python 字典:键为参数字符串(如
"conv1.weight","fc.bias"),值为torch.Tensor。 - 文件中没有「有哪些层、层之间如何连接」的信息,只有「名字 → 张量」的映射。
- 内容是一个 Python 字典:键为参数字符串(如
- 完整对象:
torch.save(model, "model.pth")- 会序列化整个
nn.Module对象(含子模块结构),但依赖 pickle 反序列化时的类定义存在,且不利于跨版本、跨机器,实践中更多用 state_dict。
- 会序列化整个
state_dict 里具体有什么:
- 每个键:完整参数量路径的字符串,例如
"backbone.layer4.2.bn2.weight"。- 若用
DataParallel包装后保存,键会带"module."前缀,如"module.backbone.layer4.2.bn2.weight"。 - 这些字符串会原样写入 pickle,长键名会重复占用空间。
- 若用
- 每个值:一个
torch.Tensor。- Pickle 会序列化 Tensor 的 dtype、shape、stride、device 等元数据,以及原始数值的二进制;
- 同时会带上「这是
torch.Tensor类型」的类信息、以及 PyTorch 模块路径(如torch._tensor),以便反序列化时重建对象。 - 因此同一组数值在 .pth 里占用的字节数 ≥ 纯数值本身。
不包含什么:
- 计算图:没有「先 Conv 再 ReLU 再 Linear」的拓扑,只有「这一坨权重叫什么名字」。
- 输入/输出约定:没有「模型接受什么 shape、输出什么」的显式声明;若想知道,只能看代码里
forward的签名和实现,或根据某些权重的 shape 反推(如最后一层 Linear 的 out_features 当作类别数)。
因此,.pth = 「权重的快照」,必须配合「定义结构的代码」才有完整含义;换一台机器或另一个框架,若没有同一份代码,仅靠 .pth 无法还原出可执行的模型。
2.2 使用时的依赖关系
用 checkpoint 做推理时,必须同时具备:
- 定义模型的 Python 代码:例如
class MyModel(nn.Module): ...,在forward里规定输入输出和计算顺序; - 构建模型的逻辑:根据配置或命令行参数实例化
MyModel(...),得到「图」在内存中的表示; - Checkpoint 文件:
model.load_state_dict(torch.load(path))把 .pth 里的权重填进已构建好的图。
缺少任一项都无法得到「可前向计算的模型」。因此 checkpoint 推理链路可以概括为:
┌─────────────────────────────────────────────────────────────────┐
│ Checkpoint 推理链路(PyTorch) │
│ │
│ 模型代码 ──► 实例化 model = MyModel(...) ──► 内存中有一张「图」 │
│ │ ▲ │
│ │ │ load_state_dict(torch.load(.pth)) │
│ └────────────────────┴────────── .pth 只提供权重数值 │
│ │
│ 输入 ──► model(x) ──► 输出 (输入/输出由 forward 的代码决定) │
└─────────────────────────────────────────────────────────────────┘
2.3 通用代码示例:保存与加载
import torch
import torch.nn as nn
# 定义结构(只有代码里有)
class SimpleNet(nn.Module):
def __init__(self, in_dim=10, num_classes=5):
super().__init__()
self.fc = nn.Linear(in_dim, num_classes)
def forward(self, x):
return self.fc(x)
# 训练后只保存权重
model = SimpleNet(in_dim=10, num_classes=5)
# ... 训练 ...
torch.save(model.state_dict(), "model.pth")
# 之后在另一处加载:必须先「再建一张同样的图」,再填权重
model2 = SimpleNet(in_dim=10, num_classes=5) # 必须知道 in_dim、num_classes
model2.load_state_dict(torch.load("model.pth"))
model2.eval()
out = model2(some_input) # 输入/输出 shape 由 forward 决定,.pth 里没有写
说明:.pth 里没有写 in_dim、num_classes,也没有写「输入是 (N, 10)、输出是 (N, 5)」;这些要么来自配置/约定,要么从权重 shape 反推(例如 state_dict["fc.weight"].shape 得到 (num_classes, in_dim))。
三、ONNX(.onnx)是什么
3.1 定义与内容
ONNX(Open Neural Network Exchange)是一种与框架无关的模型表示,用 Protocol Buffers 描述计算图,用二进制存储权重与常量。
文件中大致包含:
- 计算图(Graph)
- 节点(Node):每个节点对应一个算子,如
Conv,Relu,MatMul,BatchNormalization。 - 节点有:
op_type(算子类型)、input(输入张量名列表)、output(输出张量名列表)、attribute(如kernel_shape: [3, 3],strides: [1, 1])。 - 边由「张量名」体现:节点 A 的某个 output 名字与节点 B 的某个 input 名字相同,即表示数据从 A 流到 B。
- 图不包含 Python 或任何框架特有的控制流,只有「数据流图」:张量经过哪些算子、顺序如何,都在节点与边的拓扑里写死。
- 节点(Node):每个节点对应一个算子,如
- 初始器(Initializer)
- 图中作为「常量」出现的张量(例如卷积核、BN 的 weight/bias、Linear 的 weight/bias)会列在 initializer 里。
- 每个 initializer 有:名字(与图中某条边的名字一致)、dtype、dims、以及原始数值的二进制块(如 float32 按行优先存储)。
- 推理时,这些常量被加载到内存,前向时直接使用,不再从外部读入。
- 输入/输出声明(ValueInfo)
- 图的入口和出口会显式声明:名字、元素类型(如 float)、shape(可以是具体数字,也可以是符号维度,如
batch_size)。 - 例如输入
"image" : float32[batch, 3, 224, 224],输出"logits" : float32[batch, 1000]。 - 这样契约明确:任何运行时只要按名字和 shape 喂入/取出即可,无需看训练代码。
- 图的入口和出口会显式声明:名字、元素类型(如 float)、shape(可以是具体数字,也可以是符号维度,如
不包含:任何 Python 源码、PyTorch 的 nn.Module 定义、或训练框架特有的 API。因此,.onnx = 「图结构 + 权重 + 输入/输出契约」,可在不依赖训练框架的环境下推理(例如 ONNX Runtime、TensorRT、OpenVINO 等)。
3.2 使用时的依赖关系
用 ONNX 做推理时只需:
- 一个 .onnx 文件(图 + 权重已内嵌);
- 支持 ONNX 的运行时(如 ONNX Runtime:
ort.InferenceSession); - 与声明一致的输入:按输入名传入符合 shape 和 dtype 的数组;
- 按输出名取出结果。
无需模型定义代码、无需 PyTorch。例如(伪代码):
import onnxruntime as ort
session = ort.InferenceSession("model.onnx")
# 输入/输出名、shape 可从 session.get_inputs() / get_outputs() 查询
outputs = session.run(None, {"input_name": input_array})
3.3 为何需要「导出时固定」输入/输出
训练时的 PyTorch 模型往往有多个输入(例如图像 + 标签、或图像 + 相机 ID)、以及训练/推理分支(如 if self.training: ... else: ...)。ONNX 是静态图,通常只支持「固定的一组输入、固定的一组输出」;且很多部署场景希望接口简单(例如单输入「图像」、单输出「特征」或「logits」)。因此导出前会做两件事:
- 在 Python 里包一层:写一个
nn.Module,其forward只接受「真正要暴露的输入」(如图像),在内部构造或固定其余输入(如置 0、或从配置读),再调用原始model,并只返回需要的那一个张量。这样 ONNX 看到的就只有一个输入、一个输出。 - 在 export 时指定名字:
torch.onnx.export(..., input_names=["image"], output_names=["feature"]),这样 ONNX 图里的入口/出口名字就是"image"和"feature",与代码里的变量名解耦,便于运行时按名绑定。
四、两者的关系与在流水线中的角色
- Checkpoint:训练产出的「权重快照」,在 PyTorch 侧用于继续训练或推理;不包含图,图由代码定义。
- ONNX:由「当前代码定义的图 + 从 checkpoint 加载的权重」导出得到,是图与权重的一体化表示,用于部署。
- 关系:Checkpoint 是权重的来源;ONNX 是导出结果,写入 ONNX 后,.onnx 自包含,不再依赖 .pth。
五、转换方式(通用流程)
5.1 逻辑步骤
- 加载 checkpoint:
state_dict = torch.load(checkpoint_path, map_location="cpu")。 - 用代码构建图:根据配置或从 state_dict 推断出的超参(如 num_classes),实例化模型,例如
model = MyModel(num_classes=num_classes)。 - 把权重填进图:
model.load_state_dict(state_dict)(若键带module.前缀,可先遍历 state_dict 去掉前缀再 load)。 - (可选)包一层 Wrapper:若原始模型有多输入/多输出或训练分支,写一个只接受「要导出的输入」、只返回「要导出的输出」的
nn.Module,内部调用model。 - 构造 dummy 输入:与真实推理时 shape 一致(如
(1, 3, 224, 224)),dtype 与设备一致。 - 调用导出:
torch.onnx.export(wrapper_or_model, dummy, onnx_path, input_names=[...], output_names=[...], opset_version=..., ...)。 - (可选)量化:用 ONNX Runtime 的
quantize_dynamic等对 .onnx 做权重量化,得到另一份更小的 .onnx。
5.2 通用导出代码示例
import torch
import torch.nn as nn
class Wrapper(nn.Module):
"""单输入单输出,便于 ONNX 导出。"""
def __init__(self, model):
super().__init__()
self.model = model
def forward(self, x):
# 若原模型还有其它输入,在此固定(如全 0)
return self.model(x) # 只返回需要的那一个输出
# 构建模型并加载权重
model = MyModel(num_classes=1000)
model.load_state_dict(torch.load("model.pth", map_location="cpu"))
model.eval()
wrapped = Wrapper(model)
dummy = torch.randn(1, 3, 224, 224)
torch.onnx.export(
wrapped,
dummy,
"model.onnx",
input_names=["image"],
output_names=["logits"],
opset_version=17,
do_constant_folding=True,
dynamic_axes={"image": {0: "batch"}, "logits": {0: "batch"}},
)
六、转换原理:从「跑一次前向」到「写出 ONNX」
6.1 不是「把 .pth 直接转成 ONNX」格式
转换不是对 .pth 文件做「格式翻译」或「解析 state_dict 再重写成 ONNX 语法」。而是:
- 用 Python 代码 在内存里构建一张「图」(nn.Module 的 forward 定义的计算);
- 用 checkpoint 把权重填进这张图;
- 用 dummy 输入 跑一次前向,导出器在这一次执行中追踪所有参与计算的算子及数据流;
- 导出器把追踪到的图 + 当前模型参数(作为常量) 一起写成 ONNX。
因此:ONNX 里的权重 = 当时已加载进内存的那份权重的副本;写出 ONNX 后,.onnx 不再依赖 .pth 或 Python 代码。
6.2 追踪(Tracing)在做什么
- 执行一次:
output = wrapper(dummy)会真实执行 Conv、BN、ReLU、Linear 等。 - 记录算子:PyTorch 的导出器会在这条执行路径上,记录每个参与计算的算子(op)及其输入/输出张量、属性(如 kernel_size)。
- 得到静态图:记录结果是一张「有向无环图」:节点是算子,边是张量。只有 dummy 实际走到的分支会被记录;
if self.training里未走到的分支不会出现在 ONNX 里。 - 参数当常量:前向中用到的
weight、bias等,在追踪时已经是具体数值;导出器把这些 Tensor 的数值写入 ONNX 的 initializer,图中用这些名字作为「常量输入」。
所以:图 = 一次前向所经路径的「快照」;权重 = 当时内存里参数的快照。
6.3 数据流示意

6.4 导出参数的具体含义
-
model, dummy, onnx_path
- 用
dummy执行一次model(dummy)以追踪图;权重在model里已加载,导出器会把这些参数当作常量写入 ONNX。
- 用
-
input_names / output_names
- 给 ONNX 图的入口/出口张量起名。推理时运行时按这些名字绑定输入、读取输出。
- 可与 Python 里变量名不同;只要导出时与图中实际张量对应即可。
-
opset_version
- ONNX 的算子集版本号(如 17、18)。版本越高,支持的算子越多、语义越新;运行时也需支持该 opset,否则可能报错或不兼容。
-
do_constant_folding
- 「常量折叠」:在导出前,把图中「只依赖常量的子图」先算成常数。
- 例如:某节点是「常量 A 与常量 B 相乘」,导出前会直接算成常量 C,图中只保留 C,减少节点数、有时还能和后续算子融合,有利于推理优化和减小图体积。
-
dynamic_axes
- 指定哪些维度是「符号维度」(可变)。例如
{"image": {0: "batch"}}表示输入image的第 0 维是 batch,推理时可以是 1、4、32 等; - ONNX 里该维会变成符号(如
batch),而不是写死为 1。若不设 dynamic_axes,则所有维度按 dummy 的 shape 固定(如 batch=1)。
- 指定哪些维度是「符号维度」(可变)。例如
-
dynamo / legacy 导出器
- PyTorch 2.x 起,默认可能用基于 TorchDynamo 的导出路径,对复杂控制流、动态 shape 支持更好;
- 旧版用 TorchScript 的 trace 方式:严格按一次执行路径记录,控制流或动态 shape 复杂时容易失败,此时可关闭 dynamo 用旧导出器,并配合固定 batch(不设 dynamic_axes)提高成功率。
七、为什么 ONNX 文件反而更小?
ONNX 包含「图 + 权重」,理论上比「只有权重」的 .pth 多了一张图,但实际常见情况是 .onnx 比 .pth 更小,原因在于存储格式与元数据开销,而不是 ONNX 少存了权重。
7.1 Checkpoint(pickle)的膨胀来自哪里
- 长键名:state_dict 的键可能是
"backbone.layer4.2.bn2.weight"这类字符串。Pickle 会把这些字符串完整写入;若有很多层,键名总长度可达数百 KB。 - 每个 Tensor 的 Python 包装:pickle 要能反序列化出
torch.Tensor,会写入类型信息、模块路径(如torch)、以及 dtype、shape、stride、device 等;数值部分虽然也是二进制,但外面裹了一层对象描述。 - 嵌套结构:state_dict 是 dict,其值可能是嵌套结构;pickle 会递归编码,带来额外边界与引用信息。
- 协议与兼容性:pickle 协议会带版本和兼容信息,以便不同 Python 版本能读,这也占一点空间。
所以:同一组权重的「纯数值」在两边体积接近;.pth 多出来的是「键名 + Python 对象元数据 + pickle 结构」。
7.2 ONNX 的存储为何更紧凑
- 图结构:用 Protocol Buffers 描述,节点只需算子类型、输入/输出名字(短字符串)、属性(小整数、小数组);整张图的描述通常只有几 KB 到几 MB,远小于权重本身。
- 权重(Initializer):每个 initializer 用「名字 + dtype + dims + 原始二进制」存储,没有 Python 类型、没有设备信息;名字在图中可复用,且通常较短。
- 无嵌套:图是扁平的「节点列表 + initializer 列表」,没有 Python 的嵌套 dict/list 结构。
因此:同一份权重的数值,在 ONNX 里占用的字节数 ≈ 纯 float32 体积;在 .pth 里还要加上 pickle 与键名的开销,所以总文件 .onnx 可以更小。
7.3 量化后的 ONNX(如 INT8)
若对 ONNX 做权重量化(例如 INT8):
- 权重从 float32(4 字节/参数)变为 int8(1 字节/参数),约为原来的 1/4;
- 图结构几乎不变,但权重大幅缩小,整体会明显小于未量化的 .onnx 和 .pth。
- 量化的作用、原理与实现细节见下一节。
八、量化:作用、原理与实现
8.1 量化的作用
量化(Quantization) 把模型中的浮点权重(及可选地激活)从高精度(如 float32)映射到低精度(如 int8),主要带来三方面收益:
| 作用 | 说明 |
|---|---|
| 减小体积 | float32 每参数 4 字节,int8 每参数 1 字节,权重量化后约 1/4,模型文件显著变小,便于分发与存储。 |
| 加速推理 | 整数运算在 CPU/GPU/NPU 上通常比浮点更省周期、更易做 SIMD;内存带宽压力下降,有利于提高吞吐、降低延迟。 |
| 降低内存与带宽 | 推理时权重与中间激活占用更少内存,对嵌入式、移动端或高并发服务更友好。 |
代价是精度损失:量化会引入舍入误差,可能带来少量精度下降;通过合理选择量化方式(动态/静态、per-tensor/per-channel)和可选校准,通常可将影响控制在可接受范围内。
8.2 量化的原理
8.2.1 从浮点到整数的映射
把 float32 数值映射到 int8(-128~127)的常用公式为:
val_float = scale × (val_int8 - zero_point)
即:
- scale:缩放因子,正实数,表示「一个 int8 单位对应多少 float」。
- 例如权重范围约为 [-2.5, 2.5],若用对称量化(zero_point=0),可取
scale = 2.5 / 127,则val_int8 = round(val_float / scale)落在 [-127, 127]。
- 例如权重范围约为 [-2.5, 2.5],若用对称量化(zero_point=0),可取
- zero_point:浮点零对应的整数值。
- 很多算子(如 Conv 的 padding、ReLU 的零)依赖「精确表示 0」;若量化后 0 没有唯一整数值对应,会带来误差,因此量化方案里会保证
float(0) = scale × (zero_point - zero_point)或等价地让 0 落在可表示的整数上。
- 很多算子(如 Conv 的 padding、ReLU 的零)依赖「精确表示 0」;若量化后 0 没有唯一整数值对应,会带来误差,因此量化方案里会保证
推理时:权重量化后以 int8 存储;计算前按 scale、zero_point 反量化为浮点再算,或直接在整数域做「量化感知」的矩阵乘/卷积(由运行时实现)。
权重量化:只对权重做上述映射,权重文件里存 int8 + 每层或每张量的 scale/zero_point;激活可以仍用 float32(动态量化),或在静态量化中同样量化为 int8。
8.2.2 动态量化 vs 静态量化
-
动态量化(Dynamic Quantization)
- 权重:在导出/离线阶段就量化为 int8(并写入 ONNX);
- 激活:在推理时按当前输入的统计范围临时算 scale/zero_point,再量化,不依赖校准集。
- 优点:无需准备校准数据,流程简单;适合权重占主导、激活分布随输入变化较大的场景。
- 缺点:激活仍在 float 上算(或运行时再量化),加速与体积收益主要来自权重的 1/4 与带宽降低。
-
静态量化(Static Quantization)
- 权重与激活:都在离线阶段量化为 int8,激活的 scale/zero_point 通过校准集前向得到统计后确定。
- 优点:推理时整算可做到底,加速和体积收益更大。
- 缺点:需要代表性校准数据与校准流程,实现和调参更复杂。
下面给出的实现示例是 权重的动态量化(权重 → int8,激活不预先量化):即 ONNX Runtime 的 quantize_dynamic,只量化权重,激活仍在运行时以 float 参与计算。
8.3 实现:ONNX Runtime 权重量化
ONNX Runtime 提供 quantize_dynamic:读入 float32 的 .onnx,将权重量化为指定类型(如 QInt8),写出新的 .onnx。
推理时由 ONNX Runtime 按量化信息反量化或做整型运算。
典型调用方式:
from onnxruntime.quantization import quantize_dynamic, QuantType
quantize_dynamic(
model_input="model.onnx", # 原始 float32 ONNX
model_output="model_quant.onnx",
weight_type=QuantType.QInt8, # 权重量化为有符号 8 位整型
optimize_model=True, # 先做图优化再量化,有利于融合与常量折叠
)
- weight_type=QuantType.QInt8:所有权重从 float32 变为 int8(约 1/4 体积)。
- optimize_model=True:量化前对图做一轮优化(如常量折叠、算子融合),再量化,通常能得到更小或更快的模型。
- 不传校准数据,因此是动态量化:只量化权重;激活的量化若需要则由运行时在推理时处理(此处主要为权重量化带来的体积与带宽收益)。
使用方式:先导出 float32 的 .onnx,再在离线阶段调用 quantize_dynamic 得到量化后的 .onnx(如 model_quant.onnx)。推理时用量化模型与用原始模型的接口一致(同一套 input/output 名字);ONNX Runtime 会在内部按量化参数处理权重,通常能获得更小的内存占用和一定的加速。
九、对比小结
9.1 一句话对照
| 概念 | 说明 |
|---|---|
| Checkpoint | 存的是权重;结构和输入/输出都依赖 Python 代码;换环境需带代码和依赖。 |
| ONNX | 存的是整张推理图 + 权重,输入/输出在导出时固定,不依赖训练框架,便于部署与量化。 |
9.2 推理时的「必须指定权重」vs「不需权重」
- 用 PyTorch + checkpoint 推理:必须指定 checkpoint 路径,因为要
load_state_dict才能得到完整模型;图由代码在运行时构建。 - 用 ONNX 推理:只需 .onnx 路径;图与权重都在文件内,运行时按输入/输出名喂入和取出即可,不再需要 .pth 或模型定义代码。
十、总结
- Checkpoint:权重的快照(多为 state_dict 的 pickle);结构和 I/O 依赖代码;用于训练与 PyTorch 内推理。
- ONNX:图 + 权重的标准表示(Protocol Buffers + 二进制权重);I/O 在导出时固定;用于部署、跨平台、量化。
- 关系:Checkpoint 提供权重;转换时用代码建图并加载该权重,再通过一次前向追踪把「图 + 当前权重」写入 ONNX。
- 转换方式:加载 .pth → 代码实例化模型并 load_state_dict → 可选 Wrapper 收口输入/输出 → dummy 前向 + torch.onnx.export;可选量化得到更小的 .onnx。
- 为何 ONNX 更小:同一份权重在 ONNX 中以紧凑二进制存储且无 pickle/键名开销;若再量化则为 INT8,体积进一步下降。
- 量化:权重量化为 INT8 可减小约 1/4 体积并利于推理加速;原理为 scale/zero_point 的浮点–整型映射;实现上可用 ONNX Runtime 的 quantize_dynamic,无需校准数据。
更多推荐



所有评论(0)