torch.jit.trace仅记录单次前向的张量操作路径,不捕获if/for等控制流,易因输入变化导致错误;推荐pytorch 2.0+优先使用torch.compile,支持动态shape与完整控制流,无需修改代码。

torch.jit.trace 会漏掉控制流,别在 if/for 里直接用
torch.jit.trace 只记录一次前向执行的 tensor 操作路径,不分析 Python 控制逻辑。如果你模型里有 if x.sum() > 0: 或 for i in range(n):,trace 出来的图只会固化当时那一次的分支和循环次数,换输入就可能报错或结果错乱。
实操建议:
- 只对纯张量运算、无条件跳转的模块用
torch.jit.trace(比如一个固定结构的 CNN backbone) - trace 前务必用典型输入(含边界值)跑一遍,检查输出 shape 和数值是否对得上
- trace 后调用
scripted_model.graph看 IR,确认没出现prim::If或prim::Loop—— 出现了说明控制流没被捕捉,实际已失效
torch.jit.script 要求代码可静态分析,@torch.jit.ignore 是救命符
torch.jit.script 会解析 Python AST,要求所有分支、循环、类型都能在编译期确定。常见报错如 Could not infer dtype of NoneType 或 Unsupported operation: list comprehension,本质是它卡在类型推导或语法支持上。
实操建议:
- 把不能 script 的部分(如 logging、plt.show()、非 tensor 的 dict 操作)用
@torch.jit.ignore装饰器标出 - 避免用 Python list 存 tensor,改用
torch.stack()或torch.cat();循环尽量用torch.arange().for_each替代原生 for - 函数参数必须带类型注解(
def forward(self, x: torch.Tensor) -> torch.Tensor:),否则 script 会失败
PyTorch 2.0+ 推荐优先用 torch.compile,不是 jit 的替代品而是新范式
torch.compile 不生成 TorchScript 图,而是通过 FX Graph + 后端(如 Inductor)做图优化,支持动态 shape、autograd、几乎所有 Python 控制流,且无需修改模型代码。
实操建议:
- PyTorch ≥ 2.0 时,先试
model = torch.compile(model, backend="inductor"),比 trace/script 更省心 - 注意:compile 默认启用
dynamic=True,若想强制静态 shape(比如部署到 TensorRT),要显式传dynamic=False - 首次运行会慢(编译开销),但后续 inference 通常比 trace 快,尤其在 GPU 上
导出为 TorchScript 后,加载时 model() 和 model.forward() 行为不同
用 torch.jit.save 导出的模型,加载后直接调用 model(input) 是安全的;但若调用 model.forward(input),可能触发 Python 解释器回退(fallback),导致控制流失效或性能暴跌。
实操建议:
- 永远用
model(input)形式调用,不要碰.forward方法 - 加载后加一句
model = model.eval(),否则 dropout/batchnorm 行为可能和训练时不一致 - 如果部署到 C++,必须用
torch::jit::load()加载,且输入 tensor 需用torch::kCUDA显式指定设备,CPU tensor 传进去不会自动迁移
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











