先区分设备、参数文件与执行格式
跳转到“先区分设备、参数文件与执行格式”模型迁移到设备、保存权重、恢复训练、导出部署图,是几项相关但不同的操作。原记录把 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 与控制流有各自限制 |
| 导出 ONNX | torch.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 Pathfrom tempfile import TemporaryDirectoryimport torchfrom 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 torchfrom 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"]) == 3torch.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 Pathfrom tempfile import TemporaryDirectoryimport numpy as npimport torchfrom torch import nnimport onnximport onnxruntime as ortfrom 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 argparsefrom pathlib import Pathimport onnxfrom 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 预处理 与训练时一致。