TorchScript 是 PyTorch 官方给出的模型中间表示,它把 nn.Module 从依赖 Python 解释器的动态图形式转换成可静态分析、可序列化的图结构。这样得到的脚本可以脱离 Python 运行,也能被 C++ 程序加载,或者进一步导出到移动端和服务端部署框架。要理解 TorchScript,核心问题不是记住某个 API,而是搞清楚 torch.jit.trace 和 torch.jit.script 两种转换方式分别在做什么。

简单来说,TorchScript 提供了两条生成路径:一条是运行时追踪,另一条是源代码编译。运行时追踪会执行一次模型,把所有涉及张量计算的算子记录成图;源代码编译则直接解析 forward 方法的 Python 语法,把支持的部分翻译成静态图。两种方式并不完全等价,选错方式比写错代码更容易造成隐蔽问题。
一、TorchScript 是什么:PyTorch 模型的可部署中间表示
PyTorch 在训练阶段依赖动态计算图,每一次前向传播都可以根据 Python 控制流动态改变结构。这种灵活性对研究和调试非常友好,但部署时会带来性能与可移植性问题。TorchScript 的目标就是把模型锁定成一种可独立加载、可静态优化的图结构。转换后的模型可以保存为独立文件,不需要包含 Python 解释器即可在 C++ 环境运行。
生成 TorchScript 有两种典型方式。第一种是 torch.jit.trace,它需要你提供一组示例输入,然后把示例输入在模型中实际执行一遍,记录所有执行过的张量操作。第二种是 torch.jit.script,它不执行模型,而是读取你的 Python 源码,将 TorchScript 支持的语言子集编译成静态图。大多数情况下,简单的线性前馈网络用 trace 就能工作;一旦模型里有条件分支、循环或复杂数据结构,就需要优先考虑 script。
下面是一个最简转换示例,先把普通 nn.Module 通过 trace 生成 TorchScript,并查看生成的图代码。
import torch
import torch.nn as nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(4, 2)
def forward(self, x):
y = self.linear(x)
return torch.relu(y)
model = SimpleModel().eval()
example = torch.randn(1, 4)
traced = torch.jit.trace(model, example)
print(traced.code)
traced.save("simple_ts.pt")
这里 traced.code 能展示生成的 TorchScript 代码。如果模型只包含矩阵乘法、激活函数和形状变换,trace 通常可以稳定工作。但对于带条件判断的模型,trace 会丢失未执行到的分支,这是接下来要重点区分的内容。
二、trace 与 script 怎么选:控制流决定转换方式
选 torch.jit.trace 还是 torch.jit.script,本质上取决于模型前向逻辑中是否存在依赖输入数据的控制流。trace 只执行一次,它看到的是“哪条路走了”,而不是“为什么要走这条路”。因此当 forward 中写了 if 判断时,trace 会把这次执行到的分支固化下来,后续无论如何调用,都不会再切到另一个分支。
以下模型包含一个布尔参数 flag,在 trace 时会暴露分支固化问题。
import torch
import torch.nn as nn
class BranchModel(nn.Module):
def __init__(self):
super().__init__()
self.a = nn.Linear(4, 4)
self.b = nn.Linear(4, 4)
def forward(self, x, flag: bool):
if flag:
return self.a(x)
else:
return self.b(x)
model = BranchModel().eval()
example = (torch.randn(1, 4), True)
traced = torch.jit.trace(model, example)
# 即使传入 False,仍然执行 True 分支
print(traced(torch.randn(1, 4), False).shape)
运行上面代码后,虽然调用时传入 False,模型实际仍走 self.a 分支。因为 trace 在构建阶段只看到了 flag=True 的执行路径。对于这种情况,使用 torch.jit.script 才能保留完整控制流。
scripted = torch.jit.script(model) x = torch.randn(1, 4) print(scripted(x, True).shape) print(scripted(x, False).shape)
script 会解析 Python 源码,把 if flag 翻译成 TorchScript 的条件节点,因此不同输入会走不同分支。不过 script 对 Python 语法的支持是受限子集,不能随意使用可变参数、生成器、某些第三方库操作等特性。
可以用一个简单规则判断:模型结构固定、没有数据依赖控制流,优先用 trace;模型中有 if、for 依赖运行时判断,或者使用了复杂数据结构,优先用 script。对于既有固定子模块又有动态控制流的大型模型,也可以采用混合方式:对静态部分使用 trace,对动态部分用 torch.jit.script 或 torch.jit.script_if_tracing 装饰。
三、注意事项与避坑建议:输入形状、Python 特性和保存加载
第一个常见问题是输入形状与数据类型。trace 生成图时会记录示例输入的形状,如果生产环境输入尺寸发生变化,老版本或某些算子组合下可能报错,也可能静默产生错误。即使部分框架允许动态 batch,也建议明确测试不同形状。对于 script,虽然控制流可以保留,但 shape 推断仍然需要保证输入满足模型内部操作约束。
第二个常见问题是把 Python 标量与张量混在一起使用。例如在 forward 里写 if x.sum() > 0:,在即时模式下会被 Python 解释为对整个张量求布尔值,容易抛异常。正确的做法是显式使用 TorchScript 支持的条件,例如把阈值作为标量参数传入,或者在编写模型时避免对张量直接做布尔判断。script 模式下也建议给参数添加类型注解,例如 flag: bool 或 x: torch.Tensor,这能减少类型推断失败。
第三个问题是保存和加载。TorchScript 模型保存要使用 torch.jit.save,加载使用 torch.jit.load,而不是直接保存 state_dict。加载后要调用 eval() 确保推理模式下 dropout、batch norm 等行为正确。
scripted = torch.jit.script(model)
torch.jit.save(scripted, "branch_scripted.pt")
loaded = torch.jit.load("branch_scripted.pt")
loaded.eval()
x = torch.randn(1, 4)
print(loaded(x, True))
print(loaded(x, False))
此外,还应避免在 TorchScript 脚本中使用 Python 独有的动态行为,比如依赖任意第三方对象、使用 *args、**kwargs 或者动态修改类属性。若转换失败,可以先从最小单元开始测试,把复杂模块拆成几个小函数分别 script,定位是哪一段 Python 语法不被支持。
四、调试与验证建议:转换后如何确认模型正确性
生成 TorchScript 后,不能只看转换是否成功,还要验证结果是否与原始 PyTorch 模型一致。推荐使用 torch.testing.assert_close 对比输出。选择多组代表性输入,尤其是边界条件、不同 batch 大小和不同设备上的输入,避免只测试一组示例就上线。
import torch
model = BranchModel().eval()
scripted = torch.jit.script(model)
scripted.eval()
x = torch.randn(1, 4)
flag = False
with torch.no_grad():
eager_output = model(x, flag)
scripted_output = scripted(x, flag)
torch.testing.assert_close(scripted_output, eager_output)
print("输出一致")
转换后的 TorchScript 还提供了 .code 属性,可以像阅读源码一样检查图结构。如果发现控制流节点缺失、分支被合并或某些算子没有出现,往往说明 trace 阶段已经丢掉了关键逻辑。对复杂模型,可以把生成的图代码打印出来,和原始 forward 做对照。
另一个容易忽略的调试点是设备。模型在 GPU 上训练后如果直接在 CPU 上加载 TorchScript,可能出现设备不一致。建议在转换和保存前统一使用 eval(),并尽量在目标推理设备上完成一次完整测试。对于跨设备部署,可以在转换为脚本前调用 .cpu() 或 .to(device),保存后再加载到目标设备。
理解 trace 的记录机制和 script 的编译机制,是稳定使用 TorchScript 的关键。只要在转换前分清模型控制流类型,转换后做多组输入对比验证,就能大幅减少部署阶段的形状错误、分支丢失和数值不一致问题。
PyTorch TorchScriptTrace与Script模型部署修改时间:2026-08-30 07:21:53