pytorch转torchscript仅trace和script两种可靠方式:trace适用于无控制流的前向模型,通过示例输入录制计算图,但会固化动态分支;script需代码可静态分析,要求类型注解严格、禁用非支持库;混合使用时应分段处理并验证输出一致性。

PyTorch模型转TorchScript只有两种可靠路径:torch.jit.trace 和 torch.jit.script,选错模式会导致运行时崩溃或逻辑错误,不是“都能跑通”就代表转换正确。
trace 模式:适合纯前向、无控制流的模型
torch.jit.trace 通过记录一次前向执行的张量操作来生成计算图,本质是“录制”,不分析 Python 代码逻辑。
- 必须提供一个真实输入(
example_inputs),且 shape/dtype 要与后续推理一致;动态 batch 或变长序列需用torch.jit.trace+torch.jit.fork或改用script - 遇到
if x.sum() > 0:这类依赖张量值的分支,trace 会固化为单条路径(比如永远走 true 分支),后续输入即使满足 false 条件也不会切换 - 不支持
print()、assert、Python 字典遍历等运行时行为;报错通常是RuntimeError: Encountered a call to... - 示例:
traced_model = torch.jit.trace(model, torch.randn(1, 3, 224, 224))
script 模式:需要模型代码可静态分析
torch.jit.script 直接解析 Python AST,要求所有分支、循环、属性访问都可被类型推导器理解,对写法约束更严但语义保真度高。
- 模型类中不能有未注解的
self.xxx属性(如self.device需写成self.device: torch.device) - 函数内不能调用未被
@torch.jit.export标记的普通 Python 函数,也不能用numpy、time.sleep等非 TorchScript 支持库 - 字符串、列表推导式、字典键访问(
dict["key"])需显式标注类型,否则报TraitError或UnsupportedNodeError - 推荐先用
torch.jit.script尝试,失败再看是否 trace 更合适;示例:scripted_model = torch.jit.script(model)
混合使用 trace + script 的典型场景
很多实际模型(如带条件 ROI 处理的检测头)无法全靠一种方式转换,常见折中方案是分段处理:
- 主干网络(CNN/Transformer)用
trace,因其结构固定、无动态逻辑 - 后处理模块(NMS、坐标变换)重写为带类型注解的
torch.jit.script函数,并用@torch.jit.export导出 - 避免在 traced 模块里调用未 script 化的子模块,否则会触发
RuntimeError: Cannot insert a Tensor that requires grad... - 最终拼接用
torch.nn.Sequential不可靠,建议封装为新nn.Module类并整体script
验证和部署前必须检查的三件事
导出后直接 .save() 上线极易翻车,务必做轻量级验证:
- 用相同输入对比原始模型和 TorchScript 模型输出,误差应
(注意 <code>.eval()和torch.no_grad()一致性) - 检查
model.graph或用torch.jit.show_graph_for_ir()看是否含prim::PythonOp—— 出现即表示有未编译的 Python 调用,运行时必崩 - 在目标部署环境(如 Jetson 或 LibTorch C++)用
torch.jit.load()加载并 warmup 一次,有些 CUDA kernel 兼容性问题只在首次运行暴露
最常被忽略的是 trace 的输入 shape 绑定和 script 的类型注解完整性——前者让模型在移动端 batch=4 时静默输出错误结果,后者让模型在 C++ 加载时报 “schema mismatch”,都不是异常中断,而是沉默失效。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











