ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

PyTorch模型部署:TorchScript中的trace与script选型实战指南

2026/9/16 1:45:36 拓冰建站 浏览量
PyTorch模型部署:TorchScript中的trace与script选型实战指南 早几年做PyTorch模型上线最怕的不是精度不够而是环境不兼容。训练机上跑得飞快的模型挪到生产机器上总能遇到一堆破事Python版本、torch版本、自定义的预处理依赖、还有各种奇奇怪怪的动态库。后来TorchScript成为默认方案至少能把模型打包成一个不依赖Python脚本的独立文件。但真正开始用的时候文档里那两个函数最让人纠结torch.jit.script和torch.jit.trace。名字长得像功能也像似乎很多场景都能互换但内部逻辑完全不同。选错了轻则告警刷屏重则模型导出后行为悄悄变异推理结果和训练时对不上。这篇文章就把这两个入口的机制、边界和实战选型拆开讲清楚把我自己踩过的坑一并放进来给要上TorchScript的同学做个参考。1. 先看trace真实跑一遍代码然后把路径“烙”进图里1.1 trace到底在做什么torch.jit.trace的核心逻辑其实特别朴素你给它一个模型再给一个示例输入它就拿这个输入真的把forward跑一遍。跑的过程中PyTorch会把所有在Tensor上发生的算子操作记录下来形成一个有向无环图。之后你再用这个trace出来的模块做推理它就不再执行原来的Python代码了而是按照记录好的图一路算下去。这么说可能有点抽象我拿个最简单的例子import torch class IfModel(torch.nn.Module): def __init__(self): super().__init__() self.fc torch.nn.Linear(4, 4) def forward(self, x): if x.shape[0] 2: return torch.relu(self.fc(x)) return torch.sigmoid(self.fc(x)) model IfModel().eval() dummy torch.randn(2, 4) traced torch.jit.trace(model, dummy) print(traced.code)注意我传入的dummy第一维是2所以x.shape[0] 2这个条件在trace执行时是False。打印出来的code大概长这样graph(%self.1 : __torch__.IfModel, %x : Tensor): %fc.weight : Tensor ... %fc.bias : Tensor ... %3 : Tensor aten::linear(%x, %fc.weight, %fc.bias) %4 : Tensor aten::sigmoid(%3) return (%4)看到了吗代码里那个if判断不见了只剩下linear和sigmoid。这是因为trace根本不去分析你写的if逻辑它只是真实跑了一遍代码把执行过的那条路径记录下来。至于if条件怎么判断、另一条分支是什么它完全不关心。这个特性带来一个很大的好处trace对模型代码几乎没有要求。哪怕forward里用了各种Python语法、第三方库、奇怪的动态操作只要示例输入能跑通trace基本都能过。因为它不解析源码只关注实际发生的Tensor计算。1.2 最容易踩的坑if分支和循环被悄悄固化上面那个例子已经暴露了trace最大的坑逻辑被固化。如果你以为模型是带if判断的然后拿trace结果去服务不同shape的输入比如来了个batch size为3的请求模型依然会走sigmoid分支因为图里根本没有relu这个节点。循环也是一样的情况。PyTorch的nn.LSTM内部有循环如果拿一个定长序列去trace它导出的图会把循环完全展开成固定步数。后面换一个不同序列长度的输入要么报错要么直接按固定步数算这是个非常隐蔽的线上生产问题。再看一个更细的坑trace执行时Python层面的对象、标量、切片索引都会变成“当时的值”。比如forward里有一段逻辑依赖某个Python字典的key字典内容在训练和推理时不同那trace出来的行为就是错的。原理上trace只会把“发生在Tensor上的算子”记录进图而Python原生对象的计算发生在图之外结果一旦参与了Tensor运算就会被当作常量熔进图里。所以本质问题是trace适合“数据结构清晰、运算路径固定”的模型不适合“运算路径依赖运行时输入”的模型。另一个实际教训dropout和BatchNorm这种带训练模式的层trace前一定要先确认模型处于eval状态。不然trace执行时会把训练模式的随机行为一并“记录成默认状态”导出后的模型表现会非常奇怪。这个坑我见过不止一次很多同学一上来就torch.jit.trace完全忘了model.eval()。1.3 trace能做的和不能做的边界trace的优点很明显能处理几乎任何代码工程上省心不需要为了导出而反复改模型。它适合的对象是那些“输入输出固定、内部没有运行时条件分支”的模型尤其是CNN堆叠、Transformer encoder主体这类纯计算管线。你给它一组示例输入导出完成后它基本就是一份可预测、高性能的静态计算图。但trace的边界也很明确一旦forward里含有依赖Tensor内容的分支、循环步数不固定、或者shape在推理时会变化trace就会把这些逻辑“悄悄吃掉”。结果就是模型看起来可以保存、可以加载、可以推理但行为和Python原版不一致而且这种不一致特别难排查因为报错很少只有线上数据分布变了才会冒出来。所以如果你问“trace怎么避免踩坑”我的回答不是去学各种技巧绕开问题而是先判断模型到底适不适合trace。不适合就转用script或者把容易变化的部分挪到模型外面去处理。2. 再看script不是执行代码而是“读懂”代码2.1 script的编译机制与TorchScript方言torch.jit.script走的是另一条路。它不会真的拿示例输入去跑而是解析你的Python源码把它翻译成TorchScript方言再编译成中间表示。这个过程大约相当于编译器前端的工作读语法树、做类型推断、判断哪些变量是Tensor、哪些是int、哪些是bool然后生成一份带有显式控制流节点的图。这个机制决定了script对代码的要求比trace高得多。你说“forward里有个if x.shape[0] 2”没问题script会把这个判断编译成图里的一个分支节点之后对任意shape的输入模型都会在推理时真正判断shape[0]是否大于2。这也正是script比trace强的地方控制流是活着的动态shape是支持的。一个典型例子scripted torch.jit.script(IfModel()) print(scripted.code)输出里会保留if结构类似graph(%self.1 : __torch__.IfModel, %x : Tensor): %2 : int aten::size(%x, 0) %3 : bool aten::gt(%2, 2) %4 : Tensor prim::If(%3) block0(): %5 : Tensor ... %6 : Tensor aten::relu(%5) - (%6) block1(): %7 : Tensor ... %8 : Tensor aten::sigmoid(%7) - (%8) return (%4)这样即使线上来了batch size为5的输入模型也会正确走relu分支。对比一下trace的结果是不是差异立刻就出来了。2.2 script怎么处理控制流和动态shape说到这里就要提一个script最大的价值场景模型内部有真正的运行时逻辑。比如推理时需要按输入长度做不同处理输出需要根据某些条件组合或者循环次数来自输入Tensor的值。这类模型如果用trace基本等于把模型“绑死”在示例输入上而script则能把逻辑原封不动地带进图里。循环的例子也很直观torch.jit.script def loop_sum(x: torch.Tensor, n: int) - torch.Tensor: for i in range(n): x x 1 return x在script编译后这段循环会变成图里的Loop节点n会被当作运行时入参不同的n进来会得到不同的循环次数。如果用trace去跑n就会被当作示例输入里的常量固化。这就是script最大的价值模型在运行时依然是“活”的而不是一份快照。当然scrtipt也不是没有代价。它能读懂Python但只支持TorchScript方言规定的Python子集。你可以在torch.jit.script函数里写for/while/if可以写torch.cat、torch.stack这些Tensor算子可以用torch.Tensor的方法但如果你用了一些自定义的Python对象、继承自list的自定义类、复杂的元组解包、或者某些动态属性赋值编译就会报错。2.3 script的编译约束以及报错时该怎么读第一次用script的人大概率会被一堆编译报错淹没。这是因为TorchScript对类型推断很严格不像普通Python那样容忍变量在不同分支变成不同类型。比如torch.jit.script def bad_example(x: torch.Tensor, flag: bool): if flag: result x 1 else: result not a tensor return result这段代码在普通Python里完全合法但script编译时会直接报类型冲突分支返回的类型不一样。TorchScript要求一个变量的类型在静态分析阶段能确定下来不能一会儿Tensor一会儿str。这种限制其实很合理因为TorchScript最终要生成高效的图类型不确定就没法正确分配算子和内存。还有一类常见报错是关于torch.jit注解的。TorchScript支持Optional[torch.Tensor]、List[int]、Tuple[torch.Tensor, torch.Tensor]这类注解但如果你不写或者写错编译器就推断不出来。这时候最常见的解法就是给变量补上torch.jit.annotate或者在函数签名上标注类型。报错信息本身其实是很好的调试线索。它通常会精确到文件名、行号、列号并告诉你期望的类型和实际的类型。我遇到编译不过的情况一般先看报错涉及的变量属于哪一类再决定是加注解还是把那块Python逻辑挪到模型外面。script编译不过不代表模型不好多半是“TorchScript不支持这种Python写法”这时候要做的是拆解而不是硬怼。3. trace和script放在一起看三大分歧决定选型3.1 分歧一控制流是否完整保留这是两者最核心的分歧上面已经反复提到。trace关注的是“实际执行过的路径”script关注的是“代码本身的语义”。因此如果你的模型forward里没有任何依赖Tensor内容的if/else也没有不固定次数的循环那trace完全够用而且省心。一旦有运行时分支或者循环次数动态变化就必须用script。但控制流的判断也要看得细一点。有些分支条件是依赖Python标量的比如训练阶段和推理阶段走不同分支这种分支在部署时其实是固定路径用trace反而更稳因为不会把训练模式误带进来。有些分支条件是依赖Tensor值的比如输入长度是否超过阈值这种就必须script。所以不能迷信“含If就用script”要看if依赖的是什么。3.2 分歧二动态shape的支持程度动态shape在线上推理里太常见了。NLP模型输入句子有长有短检测模型输入图片不一定是固定尺寸推荐模型batch size也可能在运行时不固定。这部分trace的表现是示例输入一旦确定shape就半固定了。注意我说的是“半固定”。很多情况下trace出来的模型在运行时会有shape的检查PyTorch会检查输入shape是否与示例输入匹配不同shape往往会报错。有同学会试出“能跑通”但那得看模型内部的张量操作是不是碰巧支持任意shape比如纯卷积网络输入尺寸不固定也能trace。可一旦forward里有reshape、view、transpose这类依赖固定shape的操作trace就会在导出后把shape焗死换一个输入尺寸就崩。script对动态shape的支持要自然很多。只要逻辑里没有把shape強行硬编码成常量脚本编译后的图会保留运行时的维度信息。比如x.view(x.shape[0], -1)这种写法script能正确处理不同batch sizetrace却很可能把x.shape[0]当作示例输入的值固定下来。所以如果你的服务请求batch size会变优先考虑script。3.3 分歧三Python对象与副作用的处理方式这个分歧在项目里最容易被忽略但往往是你线上出Bug的元凶。trace是“执行源码”所以Python对象、Python函数、文件IO、外部库调用在trace那一刻都会被真实执行。执行的结果如果参与了Tensor计算就会被固化成常量。但后续部署推理时这些Python逻辑不会再执行。也就是说trace把“运行一次”后的世界冻结了下来。副作用则更隐蔽如果你在forward里调用了random.random()trace时随机数会被当常量带进图里之后每次推理都返回同一个随机结果——这不一定是你要的。script是“编译源码”所以它不允许你在图内调用任意Python库。像random、numpy、os.path这些在script编译阶段基本都会报错或者要求你把它们限制在图外。表面上看script限制更多但换个角度它反而逼着你把“真正的Tensor计算”和“外围数据处理”分开这对于部署可维护性来说是好事。3.4 一张表看透两个工具维度torch.jit.tracetorch.jit.script生成方式示例输入真实跑一遍forward记录算子静态分析源码并编译TorchScript方言控制流冻结只保留执行到的分支和步数保留if/for/while等逻辑动态shape依赖示例输入通常被约束相对友好支持运行时推断Python库调用执行期可调用结果固化或逃逸图内基本不支持需拆到图外对源码要求低跑得通就行高需符合TorchScript子集编译报错一般不会报行为问题靠排查报错多但定位相对清晰典型适用场景纯Tensor计算堆叠的静态模型含动态分支、循环、变化的模型实际选型时我通常用一句话总结能用script就尽量script但不要把模型里所有东西都塞进script。如果script编译实在太痛苦再根据模型实际运行路径考虑trace但必须把trace的行为边界确认清楚。4. 实测中的混合编译把trace和script组合起来用4.1 整体拆解预处理、模型、后处理分层很多人以为TorchScript导出就是“一个命令搞定”拿到项目里一跑才发现根本不是那回事。真实项目里模型往往不是孤立的forward里可能塞进了大量Python预处理逻辑读取列表、拼接字符串、检查文件是否存在、调用第三方库。这些逻辑script基本不支持trace虽然能跑通但会把Python逻辑的执行结果固化成常量同样不靠谱。所以第一个实用思路是分层把整个推理流程拆成预处理、模型计算、后处理三层。预处理的字符串处理、IO、字典操作留在Python侧或C侧模型计算部分单独提出来做成一个纯Tensor运算的Module或函数后处理如果涉及复杂Python对象也放在图外。TorchScript要管的只是中间那一层“干净”的Tensor计算这样无论用trace还是script成功率都会高很多。拿我的一个线上项目举例模型是文本序列到向量的编码器里面除了nn.TransformerEncoder还有一大段token长度的判断、mask构造逻辑。最初把整块forward丢给script编译报错几十行后来把mask构造和padding逻辑挪到调用侧只导出纯编码器部分问题立刻消失还顺手把延迟降了一个档次。很多模型不是“导不出”而是“塞太多”。4.2 先trace后script处理难编译的子模块即便我们把模型拆分到只剩纯Tensor计算偶尔还是会遇到script编译不通过的子模块比如某些写法比较动态的第三方实现。这时候可以考虑“先trace后script”的组合用trace处理那些不好编译的子模块再在父模块里用script把这些trace好的子模块串起来。这样做是有技术依据的。torch.jit.trace返回的TracedModule本质上也是一个具备TorchScript能力的模块可以被script后的父模块调用。具体操作是对子模块用trace导出替换掉原模型里的同名子模块再对整体执行torch.jit.script。父模块在script编译时只需要拿到子模块的调用接口和数据流而不需要再解析子模块内部的Python源码。class SubModel(torch.nn.Module): def forward(self, x, index): return x[index] class BigModel(torch.nn.Module): def __init__(self): super().__init__() self.sub SubModel() self.fc torch.nn.Linear(2, 2) def forward(self, x, index): y self.sub(x, index) return self.fc(y) big_model BigModel().eval() example_x torch.randn(4, 2) example_idx torch.tensor([0, 1]) traced_sub torch.jit.trace(SubModel().eval(), (example_x, example_idx)) big_model.sub traced_sub scripted torch.jit.script(big_model)这种模式特别适合处理“局部有动态结构、整体还是固定流程”的模型。你把真正难以script的Python特性缩小到一个很小的子模块用trace容忍下来剩下的主干逻辑用script保留控制流。需要提醒的是trace子模块时同样要传一个能代表线上分布的示例输入不然trace子模块自己也容易把某些shape或分支固定错。4.3 trace和script前的固定操作与常见坑不管最后选哪条路有几个固定操作我会先做好避免后面吃暗亏。第一model.eval()。前面提过训练模式下trace会把dropout的随机行为带进图里。对script也一样虽然它能保留train/eval属性但如果你在训练状态下编译行为依然可能不对。部署场景一律先eval()再导出。第二输入要用元组包好。torch.jit.trace(model, (x1, x2))这种形式最稳。如果模型forward接收多个参数第二个参数传一个tuple而不是list。detectron、transformers里的模型经常有复杂入参这里尤其容易踩坑有的同学直接把dict传进去结果警告一大堆或者类型转换出错。第三确认输出类型。TorchScript支持的输出类型包括Tensor、Tuple、List、Dict、基本类型。如果你的模型输出是一个自定义的Python对象trace可能侥幸把一些字段固化成常量script直接用不了。这种情况下必须改造模型把输出改成TorchScript支持的类型组合。第四strictFalse要慎用。torch.jit.trace有个strict参数默认True会在存在“可能改变数据流”的Python操作时发出警告。很多人因为警告烦人就把它关掉这其实是把问题藏起来了。我建议先保留strictTrue跑一次认真读警告内容确认为什么是可接受的再决定是不是要关。4.4 保存加载与冻结优化导出之后的操作也很关键。模型保存成.pt文件加载用torch.jit.load这部分基本没有坑scripted.save(model.pt) loaded torch.jit.load(model.pt)如果目标是移动端或嵌入式设备可以额外保存一份Lite Interpreter版本这样能在资源受限环境里加载。这个操作在torch.jit里都有对应API实测下来对减小包体和启动时间有帮助。还有一个常用的优化手段是torch.jit.freeze。如果是script后的模型推理前可以先冻结权重把参数变成常量再做一轮优化。不过要注意冻结后的模型不再保留可训练参数不是所有场景都适合。我的操作习惯是部署推理脚本里先用torch.jit.freeze(scripted.eval())再看性能和显存占用如果能满足需求就用冻结版保留原版作为回滚方案。5. 选型判断与调试方法项目实操中的总结5.1 决策思路从模型结构出发聊了这么多原理和案例回到最朴素的选型问题上我到底该用哪个我的决策顺序一般是这样先看模型forward里的控制流。如果依赖Tensor值的if/else、不定次数的循环基本没有优先考虑trace因为它省事、报错少适合快速落地。如果存在动态控制流直接用script别想着用trace偷懒。即使你当前的示例输入固定了路径线上变化时一定会出问题。如果script编译报错不要立刻转trace先尝试把报错的Python逻辑拆到模型外面。大多数时候拆完就过了。如果拆完依然编译不过再考虑对局部子模块使用trace用混合方式绕开难点。如果混合也搞不定考虑改造模型结构把复杂逻辑收敛到更小的范围而不是硬着头皮继续。这套思路最大的好处是优先保证TorchScript导出后的模型行为与Python端完全一致再考虑省事。上线模型最重要的是确定性零星的告警可以忍行为不一致才是最致命的。5.2 调试工具与技巧graph、code、is_tracing和is_scripting实际调试TorchScript问题时不要光靠肉眼看输出对不对。有几个工具很实用scripted.code和traced.code可以打印导出的代码这是第一个该看的。如果里面有不该出现的常量、缺失的分支、或者意外固化的shape一眼就能发现。scripted.graph能打印更底层的图结构当你要定位某个模块的输入输出是否被错误重排时很好用。配合torch.jit.freeze前后的graph对比可以看到优化器做了什么操作。还有一个相对小众但很关键的APItorch.jit.is_tracing()和torch.jit.is_scripting()。这两个可以在模型forward内部判断当前是在trace、script还是普通Python执行。有时候我会在模型代码里临时加一段比如def forward(self, x): if torch.jit.is_tracing(): # trace模式下的特殊处理 pass return x这能帮助我确认某个分支在导出时是否真的被执行相当于给导出过程加了个探针。等确认完毕再移除。这种方式在调试trace“固化分支”问题时非常有效强烈建议试试。5.3 我常用的“最小化复现”排查法TorchScript的报错有时候嵌套层级很深直接看一堆traceback容易头大。我自己的习惯是如果大模型script编译报错先把它拆成最小的可复现函数单独用torch.jit.script编译。比如报错在第12行某个变量类型推断失败那我就把包含这行的函数单独拿出来用一个最小的输入测一遍直到能独立编译通过为止。这种“剥离法”比在一整个大模型里辗转找问题高效得多。还有一种情况是trace导出后行为与Python不一致这种最难查。我一般的流程是构造一组与线上分布接近的输入分别在Python模型、trace模型上跑一遍对比输出。找一个会产生差异的输入逐步缩小范围逐模块比较中间张量找到第一个出现差异的模块。对该模块分别执行trace和script用graph和code对比差异来源。这个过程类似于调试梯度反传时定位NaN本质上是二分查找。坚持做下来绝大多数行为不一致的问题都能找到根因。另外提醒一句TorchScript导出后并不一定能加速推理到让人惊艳的程度。它带来的最大价值是部署形态的简化和环境的隔离性能提升需要配合torch.jit.freeze、算子融合、线程配置等手段才能看到明显效果。不要因为一张图慢就否定TorchScript先看看是不是冻结和优化没做。我在实际项目里的体会是trace和script不是“谁替代谁”的关系它们只是TorchScript这棵树上不同的两个入口。一个解决“怎么省事地导出”一个解决“怎么保真地导出”。搞清楚每个入口的边界按模型的真实结构去选比背一堆API参数重要得多。最后再分享一个小经验新项目里我习惯先用script跑一遍能过就直接用过不了就带着报错信息去看代码把那些看似聪明的动态写法改成TorchScript能理解的显式写法。绝大多数情况下改完之后模型反而更规整后续维护也更省心。