🏠 总目录📚 本教程 10 · 把模型交出去 ← →
📑 本页目录(点开跳转)

10 · 把模型交出去:保存、TorchScript、ONNX、torch.export、compile

⏱ 132 分钟 | ⭐ 实测:torch.save(model) 存出来的文件里,真的写着 mymodel Net 这个模块路径 —— 换个目录加载就是 ModuleNotFoundError: No module named 'mymodel'


🎯 一句话

交付一个模型永远是交付两样东西:【权重】和【怎么算】。四种格式的差别只在于「怎么算」这一半装在哪里 —— 装在你的 .py 文件里、装在 pickle 的一根引用里、还是装进导出的那张图里。 ⚠️ 这一章只管模型怎么变成一个可交付的文件。文件到手之后怎么起服务、怎么做健康检查、怎么灰度,是 《AI 全栈》 和 《模型上线之后》 的事。

本章所有实验用同一个玩具模型,它住在 mymodel.py 里:

import torch
import torch.nn as nn


class Net(nn.Module):
    def __init__(self, hidden=8):
        super().__init__()
        self.fc1 = nn.Linear(4, hidden)
        self.fc2 = nn.Linear(hidden, 2)

    def forward(self, x):
        return self.fc2(torch.relu(self.fc1(x)))

⚠️ 它只有 58 个参数(fc1 40 个 + fc2 18 个),所以下面所有文件体积基本全是元数据,不是权重。别拿这些字节数去比较格式谁更省空间 —— 它们比的是「装了多少结构信息」。


🧩 一、torch.save(model) 存的不是模型,是一根引用

import os
import torch
from mymodel import Net

torch.manual_seed(0)
m = Net()

torch.save(m, "whole.pt")                 # ① 整个模型对象
torch.save(m.state_dict(), "sd.pt")       # ② 只存权重

for f in ("whole.pt", "sd.pt"):
    print(f"{f:10} {os.path.getsize(f):6d} 字节")

print("state_dict 的 key:", list(m.state_dict().keys()))
print("类型:", type(m.state_dict()).__name__)

实跑输出:

对照

whole.pt 3401 字节

sd.pt 2411 字节

state_dict 的 key: ['fc1.weight', 'fc1.bias', 'fc2.weight', 'fc2.bias']

类型: OrderedDict

差的那 990 字节里装了什么?.pt 文件其实是个 zip,里面的 data.pkl 就是一段 pickle,可以直接拆开看它引用了哪些类:

import io
import pickletools
import zipfile


def globals_in(path):
    with zipfile.ZipFile(path) as z:
        name = [n for n in z.namelist() if n.endswith("data.pkl")][0]
        blob = z.read(name)
    out = io.StringIO()
    pickletools.dis(blob, out)
    return sorted({ln.split("GLOBAL")[1].strip().strip("'")
                   for ln in out.getvalue().splitlines() if "GLOBAL" in ln})


print("whole.pt 里引用的类:")
for g in globals_in("whole.pt"):
    print("   ", g)
print("sd.pt 里引用的类:")
for g in globals_in("sd.pt"):
    print("   ", g)

实跑输出:

对照

whole.pt 里引用的类:

__builtin__ set

collections OrderedDict

mymodel Net

torch FloatStorage

torch._utils _rebuild_parameter

torch._utils _rebuild_tensor_v2

torch.nn.modules.linear Linear

sd.pt 里引用的类:

collections OrderedDict

torch FloatStorage

torch._utils _rebuild_tensor_v2

⭐ mymodel Net —— 你的模块名和类名,一字不差地写在文件里了。 但类的代码没有进去,进去的只是这个名字。加载的时候 pickle 会照着这个名字去 import mymodel,找不到就完。

把两个文件拷到一个没有 mymodel.py 的目录里:

import torch

print("--- ① 整个模型 ---")
try:
    m = torch.load("whole.pt", weights_only=False)
    print("加载成功:", type(m).__name__)
except Exception as e:
    print(type(e).__name__ + ": " + str(e))

print("--- ② 只有权重 ---")
sd = torch.load("sd.pt")
print("加载成功:", list(sd.keys()))

实跑输出:

要点

--- ① 整个模型 ---

ModuleNotFoundError: No module named 'mymodel'

--- ② 只有权重 ---

加载成功: ['fc1.weight', 'fc1.bias', 'fc2.weight', 'fc2.bias']

💀 torch.save(model) 把你的项目结构腌进了权重文件。 它会在这些时候咬你:

✅ 规矩就一条:torch.save(model.state_dict()),加载时先自己构造一个同样结构的模型,再 load_state_dict。 state_dict 里没有任何类名,它只是一个 OrderedDict[str, Tensor] —— 纯数据,跨项目、跨版本、跨人都能读。 (state_dict 里到底装了什么、load_state_dict 为什么老报 key 对不上,在 06 · state_dict 到底装了什么。)


🔒 二、PyTorch 2.6 改了 torch.load 的默认值

即使 mymodel.py 就在手边,直接 torch.load("whole.pt") 现在也会失败:

import torch
from mymodel import Net          # ⭐ 类定义在,也照样试试

m = torch.load("whole.pt")       # torch 2.13 的默认参数

实跑输出:

查看完整报错:安全加载拒绝了自定义类(不要直接照提示关闭安全限制)
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, but it can result in arbitrary code execution. Do it only if you got the file from a trusted source.
    (2) Alternatively, to load with `weights_only=True` please check the recommended steps in the following error message.
    WeightsUnpickler error: Unsupported global: GLOBAL mymodel.Net was not an allowed global by default. Please use `torch.serialization.add_safe_globals([mymodel.Net])` or the `torch.serialization.safe_globals([mymodel.Net])` context manager to allowlist this global if you trust this class/function.

⭐ 原因写在报错里了:pickle 反序列化等于执行任意代码。 一个 .pt 文件可以在你 torch.load 的那一刻删你的家目录 —— 而 HuggingFace 上的权重是别人上传的。 2.6 之后默认 weights_only=True:只允许张量和一小撮内置类型,遇到别的一律拒绝。

⚠️ 这条会以一种意外的方式咬你:连你自己存的 checkpoint 也可能过不了。

import argparse
import numpy as np
import torch
from mymodel import Net

torch.manual_seed(0)
m = Net()
opt = torch.optim.Adam(m.parameters())
opt.step()

ckpt = {"model": m.state_dict(), "optim": opt.state_dict(), "epoch": 3, "best": 0.91}
torch.save(ckpt, "ckpt.pt")
back = torch.load("ckpt.pt")
print("① 标准 checkpoint:", list(back.keys()), " epoch =", back["epoch"])

torch.save({"cfg": argparse.Namespace(lr=1e-3), "sd": m.state_dict()}, "ckpt2.pt")
try:
    torch.load("ckpt2.pt")
except Exception as e:
    print("② 塞了个 Namespace ->", type(e).__name__)
    print("   ", [ln.strip() for ln in str(e).splitlines() if "Unsupported global" in ln])

torch.save({"arr": np.arange(3), "sd": m.state_dict()}, "ckpt3.pt")
try:
    torch.load("ckpt3.pt")
except Exception as e:
    print("③ 塞了个 numpy 数组 ->", type(e).__name__)
    print("   ", [ln.strip() for ln in str(e).splitlines() if "Unsupported global" in ln])

实跑输出:

结果对照

① 标准 checkpoint: ['model', 'optim', 'epoch', 'best'] epoch = 3
② 塞了个 Namespace -> UnpicklingError
['WeightsUnpickler error: Unsupported global: GLOBAL argparse.Namespace was not an allowed global by default. …']
③ 塞了个 numpy 数组 -> UnpicklingError
['WeightsUnpickler error: Unsupported global: GLOBAL numpy._core.multiarray._reconstruct was not an allowed global by default. …']

⭐ 三条结论:

你存了什么 结果
state_dict + 优化器 state_dict + int / float / str / list / dict ✅ 直接能读
一个 argparse.Namespace(配置对象) ❌ Unsupported global: GLOBAL argparse.Namespace
⚠️ 一个 numpy 数组 ❌ Unsupported global: GLOBAL numpy._core.multiarray._reconstruct

第三条最容易中招 —— 谁都往 checkpoint 里塞过一个 np.array 的类别权重或者混淆矩阵。

✅ 三种修法,按推荐顺序:

  1. ⭐ 别往 checkpoint 里塞对象。 配置存成 dict,numpy 数组存之前 torch.from_numpy(...) 或者 .tolist()。这样文件对所有人都是安全可读的。
  2. with torch.serialization.safe_globals([argparse.Namespace]): torch.load(...) —— 明确地放行你认识的那几个类。
  3. torch.load(..., weights_only=False) —— 只在文件是你自己产的时候用。⚠️ 别把它当成「修好了」,它是「我确认这个文件可信」。

📐 三、光有权重还不够:交付清单

state_dict 是纯数据,好处是跨项目可读,代价是「怎么算」全部丢在你这边。所以交出去的时候还得附上:

要一起交的 不交会怎样
模型代码(或者一个导出的图,见第四节起) 对方构造不出结构,load_state_dict 无从谈起
⭐ 构造超参(hidden=8、层数、num_classes) 结构对不上 → size mismatch for fc1.weight 那一串(见 06 章)
⭐ 预处理(归一化的均值/方差、tokenizer、resize 尺寸) 💀 不报错,只是精度莫名其妙地低 —— 训练和推理的预处理不一致是线上事故的头号来源
类别顺序 同上,argmax 出来的下标对到了别的类
torch 版本、opset 版本 老文件在新版本上加载出奇怪的 warning 或行为差异
模型卡 / 评测数字 对方不知道这个模型什么时候能用、什么时候不能用

⭐ 这就是为什么「导出成一张图」有价值:它把前四行里的「怎么算」那部分焊进了文件。 ⚠️ 但注意预处理通常仍在图外面 —— 除非你特意把归一化写进 forward。💡 一个很实用的习惯:把归一化那两行搬进模型的 forward,这样导出的图自带预处理,交付面就少一个能出错的地方。

(训练/推理预处理不一致的完整后果和排查法,在 《模型上线之后》04 · 训练推理一致性。)


🛑 读到这里可以停 —— 前半章讲完了(约 22 分钟):torch.save 的两种写法、weights_only 这个新默认值、以及交付清单。 后半章还有:TorchScript(trace 会把分支写死)· ONNX(换个运行时跑)· torch.export(现在的正路)· torch.compile(它不是交付格式)· 怎么选。 回来的时候不用重读,直接从下一节接着看就行。


🧊 四、TorchScript:把「怎么算」也装进文件

torch.jit 有两条路,它们的失败方式完全不同,这是本节唯一要记住的事:

import torch
import torch.nn as nn


class Cond(nn.Module):
    def forward(self, x):
        if x.sum() > 0:
            return x * 2
        return x * -1


m = Cond()
pos = torch.tensor([1.0, 2.0])
neg = torch.tensor([-1.0, -2.0])

print("原模型      :", m(pos).tolist(), m(neg).tolist())

traced = torch.jit.trace(m, pos)          # ⚠️ 只喂了正数
print("trace 之后  :", traced(pos).tolist(), traced(neg).tolist())

scripted = torch.jit.script(m)
print("script 之后 :", scripted(pos).tolist(), scripted(neg).tolist())

print("--- trace 出来的代码 ---")
print(traced.code)
print("--- script 出来的代码 ---")
print(scripted.code)

实跑输出:

关键信息

TracerWarning: Converting a tensor to a Python boolean might cause the trace to be incorrect. We can't record the data flow of Python values, so this value will be treated as a constant in the future. This means that the trace might not generalize to other inputs!
if x.sum() > 0:
原模型 : [2.0, 4.0] [1.0, 2.0]
trace 之后 : [2.0, 4.0] [-2.0, -4.0]
script 之后 : [2.0, 4.0] [1.0, 2.0]
--- trace 出来的代码 ---
def forward(self,
x: Tensor) -> Tensor:
return torch.mul(x, CONSTANTS.c0)
--- script 出来的代码 ---
def forward(self,
x: Tensor) -> Tensor:
if bool(torch.gt(torch.sum(x), 0)):
_0 = torch.mul(x, 2)
else:
_0 = torch.mul(x, -1)
return _0

💀💀 负数输入下,原模型给 [1.0, 2.0],trace 出来的给 [-2.0, -4.0]。 看 traced.code 就明白了:if 整个消失了,只剩一句 torch.mul(x, CONSTANTS.c0) —— tracer 只看见了「这次乘了 2」,它根本不知道有另一条分支存在。

⭐ trace 的失效模式是「悄悄给出错答案」,而不是报错。 它只给了一条 TracerWarning,而 warning 混在训练日志里没人会看见。 ⚠️ 凡是 if x.sum() > 0 / while / len(x) 参与控制流 / .item() 之后走分支,tracer 全部当成常量写死。

trace script
怎么工作 跑一遍,记录算子 编译你的 Python 源码
数据相关的 if / for 💀 写死成常量,不报错 ✅ 保留下来
支持的 Python 任意(因为根本没在读你的代码) ⚠️ 只支持一个子集,写法太动态会编译不过
典型用法 结构固定的 CNN 有控制流的模型(RNN、beam search)

装进文件之后长什么样

import os
import torch
from mymodel import Net

torch.manual_seed(0)
m = Net().eval()
x = torch.randn(1, 4)
print("原模型输出:", [round(v, 4) for v in m(x)[0].tolist()])

torch.jit.save(torch.jit.script(m), "scripted.pt")
torch.save(x, "x.pt")
print("scripted.pt", os.path.getsize("scripted.pt"), "字节")

实跑输出:

要点

原模型输出: [-0.1947, -0.1502]

scripted.pt 5758 字节

把 scripted.pt 和 x.pt 拷到没有 mymodel.py 的目录里:

import torch                     # ⭐ 注意:这个目录里没有 mymodel.py

m = torch.jit.load("scripted.pt")
x = torch.load("x.pt")
print("类型:", type(m).__name__)
print("输出:", [round(v, 4) for v in m(x)[0].tolist()])
print("还能看到源码:")
print(m.code)

实跑输出:

关键信息

类型: RecursiveScriptModule
输出: [-0.1947, -0.1502]
还能看到源码:
def forward(self,
x: Tensor) -> Tensor:
fc2 = self.fc2
fc1 = self.fc1
_0 = (fc2).forward(torch.relu((fc1).forward(x, )), )
return _0

⭐ 和第一节形成了完整对照:同一个目录、同样没有 mymodel.py,torch.load("whole.pt") 报 ModuleNotFoundError,而 torch.jit.load("scripted.pt") 跑出了一模一样的 [-0.1947, -0.1502]。

「怎么算」这一半真的进了文件 —— 而且 m.code 还能把它打印出来。这就是 TorchScript 的全部卖点:一个不需要你的 Python 代码就能跑的模型(C++ 的 LibTorch 也能直接加载它)。

🗓️ 但 TorchScript 现在处于维护状态,PyTorch 的力量都投在 torch.export 上(第六节)。⭐ 新项目从 torch.export 起步;torch.jit 主要是维护存量、或者你要交给一个只认 TorchScript 的老部署环境。


🕸️ 五、ONNX:交给另一个运行时

前面几种格式的读者都还是 PyTorch。ONNX 换的是运行时 —— onnxruntime、TensorRT、CoreML、各种边缘芯片的 SDK 都吃它。

⚠️ 先说一条会当场绊住你的。PyTorch 2.9 起 torch.onnx.export 默认走新的 torch.export 导出器,它需要一个额外的包:

ModuleNotFoundError: No module named 'onnxscript'

✅ 两条路:装 onnxscript,或者用 dynamo=False 回到老的 TorchScript 导出器(本机实验走的是后者,因为环境里没有 onnxscript):

DeprecationWarning: You are using the legacy TorchScript-based ONNX export. Starting in PyTorch 2.9, the new torch.export-based ONNX exporter has become the default.
import os
import torch
from mymodel import Net

torch.manual_seed(0)
m = Net().eval()
x = torch.randn(1, 4)

torch.onnx.export(
    m, (x,), "net.onnx",
    input_names=["x"], output_names=["y"],
    dynamic_axes={"x": {0: "batch"}, "y": {0: "batch"}},   # ⭐ 声明哪一维是可变的
    dynamo=False,
)
print("net.onnx", os.path.getsize("net.onnx"), "字节")
print("原模型输出:", [round(v, 4) for v in m(x)[0].tolist()])
torch.save(x, "x1.pt")

实跑输出:

要点

net.onnx 675 字节

原模型输出: [-0.1947, -0.1502]

导出的东西可以直接拆开看:

import numpy as np
import onnx
import onnxruntime as ort
import torch

g = onnx.load("net.onnx").graph
print("输入 :", [(i.name, [d.dim_param or d.dim_value
                          for d in i.type.tensor_type.shape.dim]) for i in g.input])
print("输出 :", [(o.name, [d.dim_param or d.dim_value
                          for d in o.type.tensor_type.shape.dim]) for o in g.output])
print("节点 :", [n.op_type for n in g.node])
print("权重 :", [(w.name, list(w.dims)) for w in g.initializer])

sess = ort.InferenceSession("net.onnx")
x = torch.load("x1.pt")
print("onnxruntime 输出:", [round(v, 4) for v in sess.run(None, {"x": x.numpy()})[0][0].tolist()])

big = np.random.randn(5, 4).astype(np.float32)
print("换成 batch=5 :", sess.run(None, {"x": big})[0].shape)

实跑输出:

要点

输入 : [('x', ['batch', 4])]

输出 : [('y', ['batch', 2])]

节点 : ['Gemm', 'Relu', 'Gemm']

权重 : [('fc1.weight', [8, 4]), ('fc1.bias', [8]), ('fc2.weight', [2, 8]), ('fc2.bias', [2])]

onnxruntime 输出: [-0.1947, -0.1502]

换成 batch=5 : (5, 2)

⭐ 三件事值得盯一会儿:

  1. nn.Linear 没了,变成了 Gemm。 ONNX 里只有算子,没有你的类、没有 nn.Module 的层级结构 —— 它是一张纯粹的计算图。Relu 也不再是一个「层」,就是图上一个节点。
  2. 权重名字还留着(fc1.weight 等),且权重就在这个文件里(initializer),所以 ONNX 是自包含的:一个文件 = 结构 + 权重。
  3. ⭐ dynamic_axes 声明过的那一维显示成 'batch'(字符串)而不是数字,所以 batch=5 能跑通。没声明的维会被写成固定数字,换个尺寸就直接报错。

💀 同一个坑:数据相关的分支

老的 ONNX 导出器是基于 trace 的,所以第四节那个坑原样存在:

import onnx
import onnxruntime as ort
import torch
import torch.nn as nn


class Cond(nn.Module):
    def forward(self, x):
        if x.sum() > 0:
            return x * 2
        return x * -1


m = Cond().eval()
pos = torch.tensor([[1.0, 2.0]])
torch.onnx.export(m, (pos,), "cond.onnx", input_names=["x"], output_names=["y"],
                  dynamic_axes={"x": {0: "b"}, "y": {0: "b"}}, dynamo=False)
print("节点:", [n.op_type for n in onnx.load("cond.onnx").graph.node])

sess = ort.InferenceSession("cond.onnx")
print("正数 pytorch:", m(pos).tolist(), " onnx:", sess.run(None, {"x": pos.numpy()})[0].tolist())
neg = torch.tensor([[-1.0, -2.0]])
print("负数 pytorch:", m(neg).tolist(), " onnx:", sess.run(None, {"x": neg.numpy()})[0].tolist())

实跑输出:

对照

节点: ['Constant', 'Mul']

正数 pytorch: [[2.0, 4.0]] onnx: [[2.0, 4.0]]

负数 pytorch: [[1.0, 2.0]] onnx: [[-2.0, -4.0]]

⚠️ 正样例对上了,负样例差得离谱。 这解释了一条老生常谈的部署纪律:导出之后必须拿一批真实数据逐条比对 PyTorch 和 ONNX 的输出,只测一条样例正好是最容易漏掉这个 bug 的做法。

💀💀 更狠的一种:整张图变成一个常数

上一章那个「前向绕出去用 numpy 算」的 Function(09 · 自己写一个 autograd 算子),导出时会发生这个:

import numpy as np
import onnx
import onnxruntime as ort
import torch
import torch.nn as nn


class NpSqrt(torch.autograd.Function):
    @staticmethod
    def forward(ctx, x):
        out = torch.from_numpy(np.sqrt(x.detach().numpy()))
        ctx.save_for_backward(out)
        return out

    @staticmethod
    def backward(ctx, g):
        (out,) = ctx.saved_tensors
        return g / (2 * out)


class M(nn.Module):
    def forward(self, x):
        return NpSqrt.apply(x)


m = M().eval()
x = torch.tensor([[1.0, 4.0, 9.0]])
torch.onnx.export(m, (x,), "np.onnx", dynamo=False)

g = onnx.load("np.onnx").graph
print("图里的节点:", [n.op_type for n in g.node])
print("图的输入  :", [i.name for i in g.input])
sess = ort.InferenceSession("np.onnx")
print("ort 要的输入:", [i.name for i in sess.get_inputs()])
print("跑一次    :", sess.run(None, {})[0].tolist())
print("pytorch 换个输入:", m(torch.tensor([[16.0, 25.0, 36.0]])).tolist())

实跑输出(省去两条 TracerWarning):

对照

图里的节点: ['Constant']

图的输入 : []

ort 要的输入: []

跑一次 : [[1.0, 2.0, 3.0]]

pytorch 换个输入: [[4.0, 5.0, 6.0]]

💀💀💀 导出的模型没有输入。整个计算被折成了一个常数,永远返回 [[1.0, 2.0, 3.0]],而 PyTorch 换个输入给的是 [[4.0, 5.0, 6.0]]。

torch.onnx.export 没有报错,只给了两条 TracerWarning(Converting a tensor to a NumPy array might cause the trace to be incorrect 和 torch.from_numpy results are registered as constants in the trace)。

⭐ 由此得到一条最省事的自检:导完先看一眼图的输入列表和节点数。输入是空的、或者节点少得不像话,就说明有大段计算被折成常量了。

⚠️ 顺便回答上一章留的问题:导出的是推理图,你写的 backward 会被整个丢掉(推理不需要它,所以这没关系)。真正的问题只在于 forward 里有没有绕出框架。


🛑 第二个休息点 —— 中段讲完了(约 40 分钟)。 最后一段还有:torch.export:现在的正路 · torch.compile 不是一种交付格式 · 选哪一个 这一章确实长,分三次读完全没问题 —— 回来直接从下一节接着看。


🛑 读到这里可以停 —— 已经读了约 76 分钟。 最后一段还有(约 31 分钟):torch.export:现在的正路 · torch.compile 不是一种交付格式 · 选哪一个 回来的时候不用重读,直接从下一节接着看就行。


🧭 六、torch.export:现在的正路

torch.export 是 PyTorch 2.x 给出的答案,位置介于 trace 和 script 之间:它用符号执行走一遍你的代码,产出一张 ATen 级别的图。

import os
import torch
from mymodel import Net

torch.manual_seed(0)
m = Net().eval()
x = torch.randn(1, 4)

ep = torch.export.export(m, (x,))
print("类型:", type(ep).__name__)
print(ep.graph_module.code)
torch.export.save(ep, "net.pt2")
print("net.pt2", os.path.getsize("net.pt2"), "字节")
print("原模型:", [round(v, 4) for v in m(x)[0].tolist()])

实跑输出:

对照

类型: ExportedProgram

def forward(self, p_fc1_weight, p_fc1_bias, p_fc2_weight, p_fc2_bias, x):

linear = torch.ops.aten.linear.default(x, p_fc1_weight, p_fc1_bias); x = p_fc1_weight = p_fc1_bias = None

relu = torch.ops.aten.relu.default(linear); linear = None

linear_1 = torch.ops.aten.linear.default(relu, p_fc2_weight, p_fc2_bias); relu = p_fc2_weight = p_fc2_bias = None

return (linear_1,)

net.pt2 11581 字节

原模型: [-0.1947, -0.1502]

⭐ 注意 forward 的签名:参数从 4 个权重开始,x 排在最后 —— 权重被「提」成了函数的入参(这叫 functionalization)。图里没有任何状态,是一个纯函数。ONNX 里 nn.Linear 变成 Gemm,这里变成 torch.ops.aten.linear.default。

同样,拷到没有 mymodel.py 的目录里:

import torch                      # ⭐ 这个目录里没有 mymodel.py

ep = torch.export.load("net.pt2")
x = torch.load("x.pt")        # ④ 存的那份输入,和 net.pt2 一起拷过来
print("类型:", type(ep).__name__)
print("输出:", [round(v, 4) for v in ep.module()(x)[0].tolist()])

实跑输出:

要点

类型: ExportedProgram

输出: [-0.1947, -0.1502]

⭐ 它和 trace 最关键的一处不同:宁可报错,也不悄悄写死

同样那个 Cond 模型:

import torch
import torch.nn as nn


class Cond(nn.Module):
    def forward(self, x):
        if x.sum() > 0:
            return x * 2
        return x * -1


torch.export.export(Cond().eval(), (torch.tensor([[1.0, 2.0]]),))

实跑输出(截断):

GuardOnDataDependentSymNode: Could not guard on data-dependent expression Eq(u0, 1) (unhinted: Eq(u0, 1)).  (Size-like symbols: none)

consider using data-dependent friendly APIs such as guard_or_false, guard_or_true and statically_known_true.

⭐ 这是本节最值得记住的一句:trace 给你一个错模型,torch.export 给你一个报错。 报错难看,但它出现在你的机器上;错模型不难看,它出现在线上。

形状也是同一个道理

import torch
from mymodel import Net

torch.manual_seed(0)
m = Net().eval()
x2 = torch.randn(2, 4)
x5 = torch.randn(5, 4)

ep = torch.export.export(m, (x2,))                       # 只给了 batch=2 的样例
try:
    print("① 不声明动态维,喂 batch=5:", ep.module()(x5).shape)
except Exception as e:
    print("① 不声明动态维,喂 batch=5 ->", type(e).__name__ + ": " + str(e).strip()[:220])

b = torch.export.Dim("batch")
ep2 = torch.export.export(m, (x2,), dynamic_shapes={"x": {0: b}})
print("② 声明 batch 是动态维,喂 batch=5:", ep2.module()(x5).shape)

实跑输出:

操作步骤

① 不声明动态维,喂 batch=5 -> AssertionError: Guard failed: x.size()[0] == 2
② 声明 batch 是动态维,喂 batch=5: torch.Size([5, 2])

⭐ 默认所有形状都是写死的(Guard failed: x.size()[0] == 2 把这件事说得很直白),要动态就用 torch.export.Dim 明确声明 —— 和 ONNX 的 dynamic_axes 是同一个概念。

⚠️ 样例输入的 batch 不要用 1:0 和 1 会被当成特殊值做特化,拿 batch=1 的样例去声明动态维会得到 Constraints violated (batch)! … your code specialized it to be a constant (1)。用 2 或更大。


⚡ 七、torch.compile 不是一种交付格式

这是最容易被归错类的一个。

import torch
from mymodel import Net

torch.manual_seed(0)
m = Net().eval()
c = torch.compile(m)
print("torch.compile 返回的类型:", type(c).__name__)
print("它还是不是 nn.Module :", isinstance(c, torch.nn.Module))
print("state_dict 的 key    :", list(c.state_dict().keys()))
print("拿得到原模型吗       :", type(c._orig_mod).__name__)

实跑输出:

对照

torch.compile 返回的类型: OptimizedModule

它还是不是 nn.Module : True

state_dict 的 key : ['_orig_mod.fc1.weight', '_orig_mod.fc1.bias', '_orig_mod.fc2.weight', '_orig_mod.fc2.bias']

拿得到原模型吗 : Net

⭐ torch.compile 返回的是一个【包着你的模型的壳】,它没有产出任何文件。编译发生在第一次真正跑起来的时候,产物是这个进程里的一份编译缓存 —— 进程一结束就没了(磁盘上是有缓存的,但那是缓存,不是交付物)。 它和前三节根本不是一类东西:那三个回答「怎么把模型变成文件」,这个回答「怎么让它在你的进程里跑得快」。

💀💀 那个 _orig_mod. 前缀会咬人

import torch
from mymodel import Net

torch.manual_seed(0)
m = Net()
c = torch.compile(m)

torch.save(c.state_dict(), "compiled_sd.pt")     # ⚠️ 存的是编译后那个壳的 state_dict

fresh = Net()
fresh.load_state_dict(torch.load("compiled_sd.pt"))

实跑输出:

RuntimeError:
Error(s) in loading state_dict for Net:
    Missing key(s) in state_dict: "fc1.weight", "fc1.bias", "fc2.weight", "fc2.bias". 
    Unexpected key(s) in state_dict: "_orig_mod.fc1.weight", "_orig_mod.fc1.bias", "_orig_mod.fc2.weight", "_orig_mod.fc2.bias". 

✅ 修法:存原模型的 state_dict:

# 🧩 骨架:接着上一段跑,Net / c / torch 都在那里定义
torch.save(c._orig_mod.state_dict(), "ok_sd.pt")     # ⭐ 穿过那层壳
fresh2 = Net()
print(fresh2.load_state_dict(torch.load("ok_sd.pt")))

实跑输出 <All keys matched successfully>。

⚠️ DistributedDataParallel 会以完全相同的方式加一层 module. 前缀 —— 同一个坑的另一个化身。看到 Unexpected key(s) 里所有 key 都多了同一个前缀,就是这类问题。(load_state_dict 报错的完整分类在 06 章。)

🗓️ 本机没法实跑 torch.compile 的编译部分:调用 c(x) 会抛

InductorError: InvalidCxxCompiler: Compiler: cl is not found.

⭐ 这条报错本身值得知道:Inductor 在 CPU 上是生成 C++ 代码再编译的,所以 Windows 上要装 MSVC(cl.exe);GPU 上则需要 Triton。torch.compile 不是纯 Python 的东西,它对你的机器有编译工具链要求。 (torch.compile 到底做了什么优化、能快多少,在 《AI基础设施》07 · 算子融合与编译 —— 那边讲「快在哪」,这一章只讲「它不产出文件」。)


🛑 读到这里可以停 —— 已经读了约 101 分钟。 最后一段还有(约 30 分钟):选哪一个 · 检查点与走神救援 回来的时候不用重读,直接从下一节接着看就行。


🚦 八、选哪一个

你要干什么 用 为什么
训练中途存档、和自己人交接 ⭐ torch.save(model.state_dict()) 纯数据,跨版本跨项目都能读;代码另外给
交给一个没有你代码的 Python/C++ 环境 ⭐ torch.export + torch.export.save 现在的正路;实测在没有源码的目录里加载,输出完全一致
交给老的 LibTorch / 只认 TorchScript 的部署栈 torch.jit.script + torch.jit.save 🗓️ 维护状态,新项目别从这里起步
交给别的运行时(onnxruntime / TensorRT / CoreML / 端侧芯片) ONNX 一个文件自包含结构 + 权重,跨框架
让它在你自己的进程里跑得快 torch.compile ⚠️ 不是交付格式,它不产出文件

⭐ 一条通用纪律,比选哪个格式重要得多:

导出之后,拿一批(不是一条)真实数据,逐条比对原模型和导出模型的输出。

本章的三个坑 —— trace 写死分支(负样例 [1.0, 2.0] vs [-2.0, -4.0])、ONNX 同样写死、numpy 那个折成常数 —— 全部会被这一步抓住,也全部会被「只测一条样例」放过。

⚠️ .eval() 别忘了。 导出的是当时那个状态,train() 模式下导出的图会把 Dropout 和 BN 的训练分支一起带走(07 · train() / eval() 到底改了什么)。


🔗 这一章连到哪里

相关的地方 为什么
06 · state_dict 到底装了什么 本章第七节那条 Unexpected key(s)(_orig_mod. 前缀)和第三节的 size mismatch,完整分类在那一章
07 · train() / eval() 到底改了什么 ⚠️ 导出前必须 .eval() —— 导出的是当时那个状态,Dropout / BN 的训练分支会被一起写进图里
09 · 自己写一个 autograd 算子 ⭐ 第五节那个「导出成一个常数」的模型,就是那一章里绕出去用 numpy 算的 Function;也回答了那边留的问题:导出的是推理图,backward 会被丢掉
《AI基础设施》07 · 算子融合与编译 ⭐ torch.compile 快在哪(融合、少读写显存、少 kernel launch)—— 本章只讲「它不产出文件」
《AI基础设施》18 · 量化 量化通常发生在导出前后这条链上,ONNX / TensorRT 都有各自的量化路径
《AI全栈》14 · 容器化与部署 ⭐ 文件到手之后:怎么打进镜像、怎么起服务、怎么发版 —— 本章只到「产出一个文件」为止
《模型上线之后》04 · 训练推理一致性 ⭐ 第三节那份交付清单里最容易漏、后果最重的是预处理;那一章讲不一致会怎样、怎么查
《模型上线之后》17 · 版本回溯与可复现 权重文件之外还要钉住哪些东西才能复现一次训练
《大模型全景导论》07 · 本地部署与开源生态 大模型场景下的部署选型(这一章是通用模型的交付格式)

✅ 检查点

  1. torch.save(model) 存出来的文件里,实测能看到哪些类名?torch.save(model.state_dict()) 呢?
  2. 把 whole.pt 拷到没有 mymodel.py 的目录里加载会怎样?为什么在训练脚本里直接 torch.save(model) 尤其危险?
  3. PyTorch 2.6 改了 torch.load 的什么默认值?为什么改?
  4. 哪三类东西塞进 checkpoint 会让默认的 torch.load 失败?哪一类最容易中招?三种修法的推荐顺序是什么?
  5. 除了 state_dict,交付时还必须一起给的东西里,哪一个漏掉是「不报错、只是精度莫名其妙地低」?
  6. trace 和 script 对数据相关的 if 分别怎么处理?实测三行输出分别是什么?traced.code 里剩下什么?
  7. 在没有 mymodel.py 的目录里,torch.load("whole.pt") 和 torch.jit.load("scripted.pt") 的结果分别是什么?这个对照说明了什么?
  8. 导出的 ONNX 图里,nn.Linear 变成了什么?这说明 ONNX 里没有什么?dynamic_axes 不写会怎样?
  9. 「绕出去用 numpy 算」的模型导出成 ONNX 之后是什么样?为什么说这是本章最狠的一个坑?由此得到什么自检办法?
  10. torch.export 遇到数据相关的 if 时的行为,和 trace 有什么本质区别?用一句话概括。
  11. torch.export 默认怎么处理形状?样例输入的 batch 为什么不能用 1?
  12. torch.compile 为什么不算交付格式?它的 state_dict 有什么特别之处,会引出哪条报错?
  13. 不管选哪种格式,导出之后都必须做的那一步是什么?为什么「只测一条样例」正好会漏掉本章的三个坑?
👀 答案
  1. whole.pt 里实测有 mymodel Net、torch.nn.modules.linear Linear、collections OrderedDict、torch FloatStorage、torch._utils _rebuild_parameter / _rebuild_tensor_v2、__builtin__ set;sd.pt 里只有 collections OrderedDict / torch FloatStorage / torch._utils _rebuild_tensor_v2 —— 没有任何你的类名。关键是:进去的只是名字,类的代码没进去。
  2. ModuleNotFoundError: No module named 'mymodel'(而 sd.pt 照常加载)。训练脚本里直接存尤其危险,是因为脚本跑起来时模块名是 __main__,文件里写的就是 __main__ Net,换任何一个入口都加载不了。
  3. weights_only 从 False 改成了 True。原因写在报错里:pickle 反序列化等于执行任意代码(it can result in arbitrary code execution),而权重文件经常是从网上下的。
  4. ①自定义类(实测 Unsupported global: GLOBAL mymodel.Net)②配置对象(GLOBAL argparse.Namespace)③⚠️ numpy 数组(GLOBAL numpy._core.multiarray._reconstruct)——第三类最容易中招,谁都往 checkpoint 里塞过 np.array。推荐顺序:别往 checkpoint 里塞对象(配置存 dict、数组转 tensor 或 .tolist())→ torch.serialization.safe_globals([...]) 明确放行 → weights_only=False(只在文件是自己产的时候,它的意思是「我确认可信」不是「修好了」)。标准 checkpoint(state_dict + 优化器 state_dict + int/float)实测直接能读。
  5. 预处理(归一化的均值/方差、tokenizer、resize 尺寸)—— 它不会报错,只表现为精度莫名其妙地低。类别顺序也是同一类。💡 实用习惯:把归一化搬进 forward,这样导出的图自带预处理。
  6. trace 把分支当常量写死、不报错;script 编译源码、保留分支。实测:原模型 [2.0, 4.0] / [1.0, 2.0],trace 之后 [2.0, 4.0] / [-2.0, -4.0](负样例错了),script 之后 [2.0, 4.0] / [1.0, 2.0](对)。traced.code 里 if 整个消失了,只剩 return torch.mul(x, CONSTANTS.c0)。只有一条 TracerWarning: Converting a tensor to a Python boolean might cause the trace to be incorrect.
  7. torch.load("whole.pt") → ModuleNotFoundError;torch.jit.load("scripted.pt") → 跑出 [-0.1947, -0.1502],和原模型完全一致,而且 m.code 还能把 forward 打印出来。说明 TorchScript 把「怎么算」这一半也装进了文件,不再需要你的 Python 代码。
  8. 变成了 Gemm(实测节点是 ['Gemm', 'Relu', 'Gemm'])。说明 ONNX 里只有算子,没有你的类、没有 nn.Module 的层级结构,它是一张纯计算图(权重也在文件里,所以是自包含的)。dynamic_axes 声明过的维显示成 'batch' 字符串、能跑 batch=5;不声明就写成固定数字,换尺寸直接报错。
  9. 实测导出的图只有一个 Constant 节点、输入列表是空的,永远返回 [[1.0, 2.0, 3.0]],而 PyTorch 换个输入给的是 [[4.0, 5.0, 6.0]]。狠在导出根本没报错,只有两条 TracerWarning。自检办法:导完先看图的输入列表和节点数 —— 输入是空的、节点少得不像话,就是有大段计算被折成了常量。
  10. trace 悄悄写死,torch.export 直接报错(实测 GuardOnDataDependentSymNode: Could not guard on data-dependent expression)。一句话:trace 给你一个错模型,torch.export 给你一个报错;报错出现在你的机器上,错模型出现在线上。
  11. 默认所有形状都写死:实测用 batch=2 的样例导出、喂 batch=5 会抛 AssertionError: Guard failed: x.size()[0] == 2;用 torch.export.Dim("batch") 声明之后才是 torch.Size([5, 2])。⚠️ 样例 batch 不能用 1,因为 0 和 1 会被特化,会报 Constraints violated (batch)! … your code specialized it to be a constant (1)。
  12. 因为它不产出任何文件 —— 返回的是一个 OptimizedModule 壳(仍是 nn.Module),编译发生在第一次真跑的时候、产物是进程内的缓存。它的 state_dict 的 key 全部多了 _orig_mod. 前缀,直接存下来再往原模型上加载会报 Missing key(s) … "fc1.weight" … + Unexpected key(s) … "_orig_mod.fc1.weight" …;✅ 修法是存 c._orig_mod.state_dict()。⚠️ DistributedDataParallel 会以同样的方式加 module. 前缀。
  13. 拿一批(不是一条)真实数据,逐条比对原模型和导出模型的输出。 因为本章三个坑(trace 写死分支、ONNX 同样写死、numpy 折成常数)全都在样例输入上给出正确答案 —— trace 那个例子里正样例 [2.0, 4.0] 两边一致、负样例才暴露;只测一条样例(尤其是导出时用的那一条)必然全部漏掉。⚠️ 另外别忘了 .eval()。

🛑 可以停在这里

⚡ 走神救援

⭐ 交付模型 = 交付【权重】+【怎么算】,四种格式的差别只是「怎么算」装在哪儿。

💀 直接存整个模型对象,存的不是模型,是一根引用——把文件拆开看,里面写的是你的模块名和类名,⭐ 代码根本没进去。拷到没有那个源文件的目录就找不到类;⚠️ 在训练脚本里直接存尤其毒,因为模块名会变成 __main__。 ✅ 规矩:只存权重字典。

🔒 新版本把加载默认改成了「只读权重」(反序列化等于任意代码执行),⚠️ 于是连你自己的 checkpoint 也可能读不了——最容易中招的是往 checkpoint 里塞了别的库的数组对象。⭐ 修法按序:别塞对象 → 显式声明白名单 → 最后才考虑关掉,而且只在文件确实是自己产的时候。

🧊 ⭐ trace 只记录跑过的算子,遇到数据相关的分支会把它写死——⚠️ 实测换个输入结果就完全错了,而它只给一条 warning。 编译源码的那条路会保留分支。

🕸️ ONNX 换的是运行时:图里只剩通用算子,⭐ 你的类没了;权重自包含;⚠️ 动态维度必须提前声明,否则换个 batch 就跑不了。⚠️ 老导出器基于 trace,同一个分支坑原样存在。

💀💀 最狠的一例:一个前向绕出框架的模块导出之后,⭐ 图里没有输入、只有一个常量节点,永远返回同一个结果——而导出过程没有任何报错。 ⭐ 自检很简单:导完先看图的输入列表和节点数。

⭐⭐ torch.export 是现在的正路,它和 trace 的本质区别是:遇到数据相关分支时它报错——trace 给你一个错模型,它给你一个错误信息。

📐 交付清单里最容易漏、后果最重的是预处理:⚠️ 它不报错,只是精度莫名其妙地低。

下一节 👉 附录A-速查.md

打卡记录保存在你的浏览器里,首页能看到总进度