首页 / PyTorch 入门教程 / torch.export:模型导出新范式

PyTorch 入门教程

torch.export:模型导出新范式

本教程共 60 篇 · 第 50 篇 · 更新于 2026-08-17 · 约 4 分钟阅读

PyTorchtorch.export模型导出ExportedProgram动态形状AOTInductor部署

本节目标:搞懂「导出」是什么,学会 torch.export 的基本用法,知道动态形状怎么声明,并了解导出的下一步去向。

为什么要导出

训练时我们离不开 Python:动态图、调试器、各种库,怎么舒服怎么来。但部署环境常常没有 Python——手机 App 是 Java/Kotlin,服务器服务可能是 C++,边缘设备资源只有一点点。这时候就要把模型「导出」(Export):把训练好的模型抓成一张计算图,脱离 Python 也能跑。

打个比方。训练像在厨房里现做现吃,什么工具都有。导出就是把菜谱写成一张标准流程图,交给外面的流水线批量生产。

注意区分「保存模型」和「导出模型」。第 25 章讲的 state_dict 保存,存的是参数值,运行还是靠 Python 加载再跑。导出存的是结构加参数:一张完整的计算图,谁拿到都能跑。前者像存了食材和配方,后者像交付了一份标准作业流程。

老路已弃用:TorchScript

曾经 PyTorch 的标准答案是 TorchScript,就是 torch.jit.tracetorch.jit.script 那套。它们从 PyTorch 2.10 起已弃用,官方明确不再推荐。

Warning

TorchScript(torch.jit.trace / torch.jit.script)已弃用。旧教程里凡是教你这两兄弟的,看一眼原理就行。部署请走新路线:torch.export 导出,服务器端用 AOTInductor,端侧用 ExecuTorch。

基本用法:export 一张图

核心 API 是 torch.export.export()。给它一个模型和一组示例输入,它跑一遍前向,把实际发生的计算记成一张图:

import torch
from torch.export import export

class MyModule(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.lin = torch.nn.Linear(100, 10)

    def forward(self, x, y):
        return torch.nn.functional.relu(self.lin(x + y))

mod = MyModule()
exported = export(mod, (torch.randn(8, 100), torch.randn(8, 100)))
print(type(exported))                     # torch.export.ExportedProgram
out = exported.module()(torch.randn(8, 100), torch.randn(8, 100))
print(out.shape)                          # torch.Size([8, 10])

返回的对象叫 ExportedProgram,翻译过来是「已导出程序」。exported.module() 拿回一个能直接调用的模块,用法和原模型一模一样。想看图长什么样,直接 print(exported):会打印整张 FX 图,参数、用户输入、每个 aten 算子、输出,一目了然。

值得留意的是图的签名(graph signature)。它把输入分成两类:一类是模型参数(PARAMETER),比如 lin.weight;一类是用户输入(USER_INPUT)。参数跟着图一起存,调用时只喂用户输入即可——这层区分是导出模型能脱离原始模型对象独立运行的关键。

最大的限制:不支持图中断

torch.compile 遇到不支持的代码会「图中断」(graph break),退回 Python 执行,之后接着编译。torch.export 不行——它的目标是纯计算图,一中断就导出失败。

最容易踩的坑是数据依赖的控制流:

class Bad(torch.nn.Module):
    def forward(self, x):
        if x.sum() > 0:        # 条件依赖数据内容 → 导出报错
            return torch.sin(x)
        return torch.cos(x)

注意区分:if x.shape[0] > 2 这类只依赖形状的判断没问题。出错的是条件值来自张量内容的情况,编译器无法确定走哪条分支。

要保留分支逻辑,得改用 torch.cond

class Good(torch.nn.Module):
    def forward(self, x):
        def true_fn(x):
            return torch.sin(x)
        def false_fn(x):
            return torch.cos(x)
        return torch.cond(x.sum() > 0, true_fn, false_fn, [x])

torch.cond(条件, 真分支, 假分支, 操作数) 把 if-else 变成图里显式的一环,导出就成功了。分支函数签名必须和操作数匹配,返回的张量形状要一致。PyTorch 还提供 mapwhile_loop 等图内控制流算子,思路相同。

顺带说一句和 torch.compile 的关系。两者共用 TorchDynamo 抓图,但心态完全不同:torch.compile 是即时编译(JIT),图断了就断,还能继续;torch.export 是静态导出,图必须完整。所以导出前最好先用 torch.compile 跑一遍,用 TORCH_LOGS="graph_breaks" 看看哪里断,改干净了再导出,能少踩一半的坑。

动态形状:默认是静态的

导出时用什么形状的示例输入,图就固化成什么形状。官方教程里的视频分类模型 MViT 用 batch=2 导出,换 batch=4 的输入直接报错:Expected input at *args[0].shape[0] to be equal to 2, but got 4

线上服务的 batch 大小常常会变。解决办法是导出时声明动态维度:

from torch.export.dynamic_shapes import Dim

batch = Dim("batch", min=2, max=16)
exported = export(model, (x,), dynamic_shapes={"x": {0: batch}})

Dim 有三种姿态:Dim.AUTO 交给导出器推断;Dim.DYNAMIC 强制动态,一旦被特化成常量就报错;Dim.STATIC 明确静态。注意 min=2 不是笔误——0 和 1 这两个尺寸会被特殊处理(0/1 特化)。原因很实际:长度 0 的张量往往意味着空输入,长度 1 常被算子当作标量处理,这两类形状会触发很多专门的优化分支,所以导出器默认把它们当常量看待。想放宽这个限制,就给 Dim 显式指定从 2 开始的范围。

Tip

导出报「Guard failed: …shape」这类错,先怀疑形状被写死;报「Could not guard on data-dependent expression」,多半是张量内容参与了判断。用 TORCH_LOGS="+dynamic" 打开日志,能看到每条守卫(guard)是从哪行代码加出来的。

非严格模式与常见报错

默认的严格(strict)模式用 TorchDynamo 做符号分析,遇到不支持的 Python 写法会失败。加 strict=False 改用解释器追踪,兼容性好得多,代价是安全性保证弱一些。官方教程导出 Whisper 语音模型就靠这一招。

两个高频报错,提前认识:

  1. Cannot mutate tensors with frozen storage——模型里直接改张量元素(如 x[0] = 1)被禁止,先 clone() 一份再改。原因也好理解:导出时张量是被追踪的「证人」,随便改数据会破坏图的可靠性;
  2. Expected mod to be an instance of torch.nn.Module——导出对象必须是 nn.Module。想导出一个函数,就把它包进一个 Module 类的 forward 里。

报错信息末尾通常会提示:把 export() 换成 draft_export() 可以一次性列出所有可能的错误。排查导出问题的好帮手。

导出之后:AOTInductor

ExportedProgram 是中间产物,真正部署还要「落地」。服务器端走 AOTInductor(提前编译,Ahead-of-Time):

import torch._inductor

path = torch._inductor.aoti_compile_and_package(exported)  # 生成 .pt2 制品
model = torch._inductor.aoti_load_package(path)            # 加载即用
output = model(x)

.pt2 制品里是编译好的 C++ 共享库和 CUDA 二进制,加载后直接推理,不需要现场编译。官方教程实测:ResNet18 首次推理,AOTInductor 约 3.35 毫秒,torch.compile 首次要 4.3 秒——差距就是省掉了运行时编译的预热开销。对在线服务来说,这个差距意味着启动速度和扩容成本,不是小事。端侧(手机、嵌入式)则交给 ExecuTorch,同一套 ExportedProgram 直接下沉。

整条链路回头看,思路其实很简单:torch.export 负责「把模型变成干净的图」,后面的 AOTInductor、ExecuTorch 负责「把图变成能跑的代码」。导出这一步做对了,下游就顺畅;导出时偷懒留下的坑,下游会加倍还回来。

小结

导出就是把 Python 模型翻译成一张干净的静态图:export() 抓图,torch.cond 处理分支,dynamic_shapes 声明变化,最后交给 AOTInductor 或 ExecuTorch。下一章讲 ONNX——另一套通用格式,模型跨框架的通行证。