跳转到内容
新建笔记

PyTorch 模型:设备迁移、权重恢复与 ONNX 导出

先区分设备、参数文件与执行格式

跳转到“先区分设备、参数文件与执行格式”

模型迁移到设备、保存权重、恢复训练、导出部署图,是几项相关但不同的操作。原记录把 GPU 保存、完整模型 pickle、state_dict、TorchScript 和 ONNX 混在一起,容易把“写出了文件”误当成“能在目标环境正确执行”。

目的常用方式文件本身没有保证什么
保存参数与持久缓冲区model.state_dict() 配合 torch.save不自动保存模型构造代码、优化器与预处理配置
恢复训练状态自定义 checkpoint 字典是否包含优化器、调度器、随机状态等由保存者决定
保存完整 Python 模型对象torch.save(model, ...)不会把依赖的 Python 类源代码完整嵌入文件
保存 TorchScript 模块torch.jit.script/trace、torch.jit.save脚本/跟踪对 Python 与控制流有各自限制
导出 ONNXtorch.onnx.export不保证目标运行时支持全部算子,也不保证数值等价

GPU Tensor 可以由 torch.save 保存;“GPU 上的模型不能直接保存”不成立。为便于跨设备读取,可以将保存用的参数副本放到 CPU,或加载时使用 map_location="cpu"。model.to(device) 修改该模块注册的参数、缓冲区并返回模块;x.to(device) 则要接住返回 Tensor,不能假设普通 Tensor 变量被原地换了设备。通常在创建优化器前完成模型设备迁移,并让输入与模型处于兼容的设备、dtype。序列化说明、Module.to

原例把一个 23×23 整数 Tensor 移到可用 CUDA 或 CPU;其设备选择可以写成 torch.device("cuda" if torch.cuda.is_available() else "cpu"),再执行 dat = dat.to(device)。这只决定 Tensor 的放置,不自动迁移模型、选择浮点精度或验证所需算子的支持情况。本文为可重复检查显式固定在 CPU。

保存和恢复一个可核对的 checkpoint

跳转到“保存和恢复一个可核对的 checkpoint”

下面使用 CPU、随机数据和带 BatchNorm 的小模型运行一步,随后保存权重与 Adam 状态并恢复。state_dict 包含 BatchNorm 的持久缓冲区,不能只理解为需要梯度的参数。模型结构由 make_model() 显式重建。

from pathlib import Path
from tempfile import TemporaryDirectory
import torch
from torch import nn
def make_model():
return nn.Sequential(nn.Linear(3, 4), nn.BatchNorm1d(4),
nn.ReLU(), nn.Linear(4, 2))
torch.manual_seed(31)
model = make_model().to("cpu")
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
x = torch.randn(6, 3)
y = torch.tensor([0, 1, 0, 1, 0, 1])
model.train()
optimizer.zero_grad(set_to_none=True)
nn.functional.cross_entropy(model(x), y).backward()
optimizer.step()
model.eval()
with torch.inference_mode():
expected = model(x).clone()
checkpoint = {
"format_version": 1,
"epoch": 1,
"model": {k: v.detach().cpu().clone()
for k, v in model.state_dict().items()},
"optimizer": optimizer.state_dict(),
}
assert "1.running_mean" in checkpoint["model"]
assert "1.num_batches_tracked" in checkpoint["model"]
with TemporaryDirectory() as directory:
path = Path(directory) / "checkpoint.pt"
torch.save(checkpoint, path)
loaded = torch.load(path, map_location="cpu", weights_only=True)
assert loaded["format_version"] == 1 and loaded["epoch"] == 1
restored = make_model()
restored.load_state_dict(loaded["model"], strict=True)
restored_optimizer = torch.optim.Adam(restored.parameters(), lr=0.01)
restored_optimizer.load_state_dict(loaded["optimizer"])
restored.eval()
with torch.inference_mode():
torch.testing.assert_close(restored(x), expected)
assert len(restored_optimizer.state) == len(optimizer.state)
# This second file is generated and read locally by this example only.
whole_path = Path(directory) / "whole_model.pt"
torch.save(model, whole_path)
whole = torch.load(whole_path, map_location="cpu", weights_only=False)
whole.eval()
with torch.inference_mode():
torch.testing.assert_close(whole(x), expected)
print("checkpoint and locally generated whole-model roundtrips passed")

本例只证明同环境下结构、权重、缓冲区和优化器状态的往返,不声称断点后训练逐位一致。若需要精确继续实验,还应记录模型/数据配置、样本顺序、随机状态、学习率调度器及混合精度缩放器等。state_dict() 中 Tensor 通常仍关联模型存储:若要在内存保留某一时刻的最佳参数,需要克隆或深复制,不能仅赋给另一个变量。

从 2.6 起,在未传入 pickle_module 时,torch.load 默认采用 weights_only=True;示例仍显式写出选项。完整对象的 pickle 加载需要 weights_only=False,可能执行反序列化逻辑,因此这里只读刚由本程序生成的文件。自定义模型类仍依赖可导入的模块与相应定义,改路径或删类后不能期待旧文件独立工作;weights_only=True 也不应被理解为对任意文件的完整安全保证。torch.load 参数、Python pickle 的对象与模块依赖

原记录还单列了 pickle.dump(model, file)。pickle 协议本身不能简单概括为“不跨平台”,但模型对象会依赖 Python 模块、库版本和相关运行环境;普通 pickle 文件也不是 ONNX 等交换格式。torch.save 本身已经利用 pickle 保存对象结构并专门处理 Tensor 存储,直接另用 pickle.dump 不会消除这些依赖。恢复时要匹配实际保存机制,不能仅改扩展名或互换加载函数。

部分权重加载必须说明跳过了什么

跳转到“部分权重加载必须说明跳过了什么”

strict=False 允许缺失键和多余键,但同名 Tensor 形状不一致仍会报错。原记录只按名字过滤 checkpoint,无法解决分类头尺寸变化。下面的函数针对普通 Tensor 状态字典,明确筛选名称、形状与 dtype;它不处理自定义 get_extra_state()、量化状态或任意嵌套 checkpoint。

import torch
from torch import nn
def load_matching_tensors(model, source):
target = model.state_dict()
accepted, skipped = {}, {}
for name, value in source.items():
if name not in target:
skipped[name] = "unknown key"
elif not isinstance(value, torch.Tensor):
skipped[name] = "not a Tensor"
elif value.shape != target[name].shape:
skipped[name] = "shape mismatch"
elif value.dtype != target[name].dtype:
skipped[name] = "dtype mismatch"
else:
accepted[name] = value
if not accepted:
raise ValueError("no matching tensors; model was not updated")
result = model.load_state_dict(accepted, strict=False)
return {"loaded": sorted(accepted), "skipped": skipped,
"missing": result.missing_keys, "unexpected": result.unexpected_keys}
model = nn.Sequential(nn.Linear(3, 4), nn.Linear(4, 2))
before = {k: v.clone() for k, v in model.state_dict().items()}
source = {
"0.weight": torch.full((4, 3), 7.0),
"0.bias": torch.zeros(5),
"1.weight": torch.zeros((2, 4), dtype=torch.float64),
"retired.weight": torch.zeros(1),
}
report = load_matching_tensors(model, source)
assert report["loaded"] == ["0.weight"]
assert len(report["skipped"]) == 3 and len(report["missing"]) == 3
torch.testing.assert_close(model[0].weight, source["0.weight"])
for key in ("0.bias", "1.weight", "1.bias"):
torch.testing.assert_close(model.state_dict()[key], before[key])
print(report)

dtype 相同是这里主动采用的保守规则,不代表 load_state_dict 在所有情况下都禁止 dtype 转换。需要转换时应显式制定策略并验证精度,不能悄悄把“跳过一半权重”当成完整恢复。原记录中的 checkpoint['state_dict']、checkpoint['epoch'] 也是项目约定,使用前应先确认文件结构与模型版本。load_state_dict

ONNX:导出以后继续检查与运行

跳转到“ONNX:导出以后继续检查与运行”

不存在通用的 torch.onnx.import_onnx(PATH) 来把 ONNX 直接还原成原 Python 模型。onnx.load 读取图结构,ONNX Runtime 等执行器负责运行;这与重新构造 PyTorch 模型再加载 state_dict 不同。

原项目示例使用 lib.config.cfg、make_network、load_network,将 1×3×480×640 的 CUDA 输入导出为 clean_pvnet_green.onnx,输入/输出名为 input0 / output0,opset 为 11。这些是 PVNet 项目的构造器、权重位置和当时部署配置,不能脱离项目直接运行。下面使用独立小网络保留“构造—推理模式—导出—简化—验证”的完整过程;它不证明原 PVNet 权重、480×640 输入或 GPU 部署已经验证。

本例固定组合:PyTorch 2.8.0 CPU、ONNX 1.19.1、ONNX Runtime 1.23.2、onnxscript 0.4.0、onnx_ir 0.1.10、onnxsim 0.4.36。导出依赖之间也有版本兼容关系,不能只固定 torch 然后假定任何新版本都能组合。使用 dynamo=True 明确选择 2.8 的新导出路径,opset_version=18 是本例目标约定,不是所有部署环境的统一要求。PyTorch 2.8 导出器

from pathlib import Path
from tempfile import TemporaryDirectory
import numpy as np
import torch
from torch import nn
import onnx
import onnxruntime as ort
from onnxsim import simplify
torch.manual_seed(19)
model = nn.Sequential(nn.Flatten(), nn.Linear(12, 4),
nn.ReLU(), nn.Linear(4, 2)).eval()
example = torch.randn(2, 3, 2, 2)
batch = torch.export.Dim("batch", min=1, max=8)
with TemporaryDirectory() as directory:
path = Path(directory) / "model.onnx"
torch.onnx.export(
model, (example,), path,
input_names=["input0"], output_names=["output0"],
dynamo=True, external_data=False, verbose=False,
opset_version=18, dynamic_shapes=({0: batch},),
)
graph = onnx.load(path)
onnx.checker.check_model(graph)
simplified, valid = simplify(graph, check_n=3)
if not valid:
raise RuntimeError("simplification check failed; nothing saved")
onnx.checker.check_model(simplified)
sessions = [
ort.InferenceSession(graph.SerializeToString(),
providers=["CPUExecutionProvider"]),
ort.InferenceSession(simplified.SerializeToString(),
providers=["CPUExecutionProvider"]),
]
maximum_error = 0.0
with torch.inference_mode():
for size in (1, 2, 7):
x = torch.randn(size, 3, 2, 2)
expected = model(x).numpy()
for session in sessions:
actual = session.run(["output0"], {"input0": x.numpy()})[0]
assert actual.shape == (size, 2)
np.testing.assert_allclose(actual, expected, rtol=1e-5, atol=1e-6)
maximum_error = max(maximum_error,
float(np.max(np.abs(actual - expected))))
output = Path(directory) / "model_simplified.onnx"
onnx.save(simplified, output)
assert output.is_file()
print("maximum CPU absolute difference:", maximum_error)

check_model 验证图结构等约束,不代替数值比对;简化器的 valid=True 也不覆盖所有输入,所以示例额外对原图和简化图分别测试 batch 1、2、7。这里仅 batch 维动态,通道及空间尺寸仍固定为 3×2×2。导出时的形状约束不能当作 ONNX Runtime 的完整输入校验器;调用者仍应检查允许的范围和实际输入语义。数值误差阈值应随模型、dtype、算子与后端选择,不能照搬为所有模型的精度保证。ONNX checker、ONNX Runtime Python API、ONNX Simplifier

一个不会在验证失败后仍保存的简化入口

跳转到“一个不会在验证失败后仍保存的简化入口”

原函数在 check=False 时打印失败消息,却继续 onnx.save。下面修正控制流,并避免覆盖输入或已有输出。它限定用于本地可信、自包含的小 ONNX 文件;不处理超过内存限制的大模型或外部权重文件拆分。

import argparse
from pathlib import Path
import onnx
from onnxsim import simplify
def simplify_onnx(input_path, output_path):
source = Path(input_path).resolve(strict=True)
output = Path(output_path).resolve()
if source == output:
raise ValueError("input and output must be different")
if output.exists():
raise FileExistsError(output)
graph = onnx.load(source)
onnx.checker.check_model(graph)
simplified, valid = simplify(graph, check_n=3)
if not valid:
raise RuntimeError("simplification failed; output was not created")
onnx.checker.check_model(simplified)
payload = simplified.SerializeToString()
with output.open("xb") as stream:
stream.write(payload)
return output
def main(argv=None):
parser = argparse.ArgumentParser()
parser.add_argument("--onnx_path", required=True)
parser.add_argument("--output_path", required=True)
args = parser.parse_args(argv)
print(simplify_onnx(args.onnx_path, args.output_path))
if __name__ == "__main__":
main()

例如保存脚本为 simplify_model.py,传入 --onnx_path model.onnx --output_path model_simplified.onnx。xb 还会阻止检查后新出现的同名文件被覆盖;若写入途中磁盘出错,可能留下不完整的新文件,此函数没有承诺事务写入。用于部署时还要加入代表性输入、任务指标和目标执行器的对比;只完成简化检查仍不等于部署验收。

TorchScript、CUDA 与部署环境的边界

跳转到“TorchScript、CUDA 与部署环境的边界”

历史代码中的 torch.jit.trace 根据示例输入跟踪执行路径,依赖数据的 Python 分支可能被固定;script 尝试编译受支持的脚本代码,也不是任意 Python 程序的打包器。已有 TorchScript 项目可按原格式使用 torch.jit.load;新项目要根据目标运行时判断使用 torch.export、ONNX 或其他方案,而不是仅凭 .pt 后缀选择加载函数。PyTorch 2.8 TorchScript

ONNX opset 是算子定义的版本约定,CUDA/cuDNN 是部分执行后端的依赖,二者没有“一个 ONNX 版本对应一个 CUDA 版本”的简单关系。PyTorch 导出端与 ONNX Runtime/TensorRT 等部署端还可能使用不同依赖组合。原笔记列出的 CUDA 12.5、cuDNN 9.2 是当时下载记录,不是本文 CPU 实验的安装要求。设备诊断与配套版本的查询方法见 PyTorch 入口;加载图像前还应保持 torchvision 预处理 与训练时一致。