torch.export:模型导出新范式
本教程共 60 篇 · 第 50 篇 · 更新于 2026-08-17 · 约 4 分钟阅读
本节目标:搞懂「导出」是什么,学会 torch.export 的基本用法,知道动态形状怎么声明,并了解导出的下一步去向。
为什么要导出
训练时我们离不开 Python:动态图、调试器、各种库,怎么舒服怎么来。但部署环境常常没有 Python——手机 App 是 Java/Kotlin,服务器服务可能是 C++,边缘设备资源只有一点点。这时候就要把模型「导出」(Export):把训练好的模型抓成一张计算图,脱离 Python 也能跑。
打个比方。训练像在厨房里现做现吃,什么工具都有。导出就是把菜谱写成一张标准流程图,交给外面的流水线批量生产。
注意区分「保存模型」和「导出模型」。第 25 章讲的 state_dict 保存,存的是参数值,运行还是靠 Python 加载再跑。导出存的是结构加参数:一张完整的计算图,谁拿到都能跑。前者像存了食材和配方,后者像交付了一份标准作业流程。
老路已弃用:TorchScript
曾经 PyTorch 的标准答案是 TorchScript,就是 torch.jit.trace 和 torch.jit.script 那套。它们从 PyTorch 2.10 起已弃用,官方明确不再推荐。
WarningTorchScript(
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 还提供 map、while_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 语音模型就靠这一招。
两个高频报错,提前认识:
Cannot mutate tensors with frozen storage——模型里直接改张量元素(如x[0] = 1)被禁止,先clone()一份再改。原因也好理解:导出时张量是被追踪的「证人」,随便改数据会破坏图的可靠性;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——另一套通用格式,模型跨框架的通行证。