torch.jit.trace加速不明显甚至变慢,因其仅记录单次前向路径,遇动态控制流、变长输入或Python对象时退化为trace fallback,无法融合算子;适合输入shape固定、纯张量运算的场景。

torch.jit.trace 为什么有时加速不明显甚至变慢
因为 torch.jit.trace 只记录一次前向执行路径,遇到动态控制流(比如 if x.sum() > 0:)、不同 shape 输入或 Python 对象(如 list、dict)参与计算时,会退化为“trace fallback”,实际仍走 Python 解释器,融合失效。它适合输入 shape 固定、逻辑纯张量运算的模型。
实操建议:
- 用
torch.jit.script替代torch.jit.trace,尤其当模型含条件分支或循环——script能静态分析 Python 控制流,生成真正可融合的图 - 确保 trace 时输入是典型 batch size 和 shape,例如
torch.randn(1, 3, 224, 224),别用(0, 3, 224, 224)或全零张量(某些算子会短路) - trace 后调用
model.graph_for(*inputs)查看 IR 图,确认 conv + relu + bn 是否被合并成一个aten::conv2d节点;若看到大量prim::If或prim::ListConstruct,说明没融合成功
torch.jit.script 报错 “Tracing failed” 或 “Unsupported type List[Tensor]”
这是最常卡住的地方:torch.jit.script 默认不支持任意 Python 类型,比如返回 List[torch.Tensor] 的检测头、带字典输出的分割模型,或用了 **kwargs 的 forward 函数。
实操建议:
- 把返回值显式标注类型,用
typing.Tuple或typing.Optional替代 list/dict,例如-> Tuple[torch.Tensor, torch.Tensor] - 避免在 forward 中直接 unpack list:不要写
a, b = feats,改用a = feats[0]; b = feats[1] - 如果必须用 dict,先用
@torch.jit.export标记该方法,并在类定义中声明__constants__ = ["keys"];更稳妥的是改输出为 tuple,后处理再转 dict - 遇到
Not supported: built-in function range,把for i in range(N):改成for i in torch.arange(N):(需 N 是 tensor 或 int 常量)
导出的 TorchScript 模型在 CPU 上快了,GPU 上却没提升
常见于没关掉 cudnn benchmark 或没预热。JIT 编译后的 kernel 仍依赖 cuDNN 的自动算法选择,首次运行会花时间搜最优卷积实现,且不同 batch size 可能触发不同策略。
实操建议:
- 推理前固定 cuDNN 行为:
torch.backends.cudnn.benchmark = False,torch.backends.cudnn.deterministic = True - GPU 上务必做 warmup:用相同 shape 输入跑 3–5 次 forward,再开始计时;否则第一次耗时含 kernel 编译和 memory allocation
- 检查是否启用了
torch.jit.fuser("fuser2")(PyTorch ≥ 1.12 默认启用),它比旧版 fuser 更激进地融合算子;可通过torch._C._jit_set_texpr_fuser_enabled(True)强制开启(已弃用,仅作兼容) - 确认没混用
.cuda()和.to(device)—— JIT 模型 device 绑定后不能跨设备调用,否则回退到慢路径
部署时加载 .pt 文件报错 “Expected object of scalar type Float but got scalar type Half”
这是精度不一致导致的隐性失败:训练时用了 AMP(torch.cuda.amp.autocast),但 trace/script 时输入是 float32,而部署时传入了 half,JIT 图里所有算子都按 float32 构建,无法自动降级。
实操建议:
- trace/script 阶段就用目标精度输入,比如部署用 FP16,则用
torch.randn(..., dtype=torch.float16).cuda()去 trace - 更推荐统一在模型导出后做转换:
traced_model.half().cuda(),但注意 BN 层权重仍是 float32,需额外调用model.eval().half()并确保输入也是 half - 避免在 JIT 模型内部调用
.half()—— 它不是可追踪操作,会导致 graph 断裂 - 生产环境建议用
torch.jit.load(..., map_location="cpu")加载后再统一.to(device).half(),别依赖保存时的 device 状态
真正卡住性能的往往不是 fusion 本身,而是 trace 输入的 shape 和 dtype 和线上请求对不上,导致 runtime fallback 到解释执行。多打一行 print(model.graph_for(...)) 比反复调参更管用。


















