PyTorch TorchScript导出时如何为模型定义多入口点?
可以为TorchScript模型设置多个入口点
针对你的问题,答案是可以——在将MyModel导出为TorchScript时,完全可以把update()和predict()都作为模型的入口点调用。下面是两种实现方式:
方法一:使用torch.jit.script脚本化模型
脚本化会直接解析模型的Python代码,保留所有可调用的方法,无需额外配置即可使用多个入口点:
import torch class MyModel(torch.nn.Module): def __init__(self): super().__init__() self.param = torch.nn.Parameter(torch.tensor(0.0)) def update(self): # 参数更新逻辑 self.param.data += 1.0 def predict(self, X): # 预测逻辑 return X * self.param # 脚本化模型 scripted_model = torch.jit.script(MyModel()) # 直接调用两个入口方法 scripted_model.update() pred_output = scripted_model.predict(torch.tensor(2.0)) print(pred_output) # 输出: tensor(2.) # 保存模型 scripted_model.save("my_scripted_model.pt") # 加载后仍可调用两个方法 loaded_model = torch.jit.load("my_scripted_model.pt") loaded_model.update() pred_output = loaded_model.predict(torch.tensor(2.0)) print(pred_output) # 输出: tensor(4.)
方法二:使用torch.jit.trace_module追踪多方法
如果你的方法不含复杂控制流,也可以用trace_module显式指定要追踪的多个方法作为入口点:
# 初始化模型 model = MyModel() # 追踪指定的update和predict方法 traced_model = torch.jit.trace_module( model, {"update": (), # update无参数,传入空元组 "predict": (torch.tensor(1.0),)} # predict需要输入张量作为示例 ) # 调用入口方法 traced_model.update() pred_output = traced_model.predict(torch.tensor(2.0)) print(pred_output) # 输出: tensor(2.) # 保存与加载 traced_model.save("my_traced_model.pt") loaded_traced = torch.jit.load("my_traced_model.pt") loaded_traced.update() pred_output = loaded_traced.predict(torch.tensor(2.0)) print(pred_output) # 输出: tensor(4.)
注意事项
- 脚本化(
script)更适合包含if/else、循环等动态控制流的方法,能完整保留逻辑;追踪(trace)仅记录执行时的张量操作,对动态逻辑支持有限。 - 两种方式导出的模型,加载后都能直接调用
update()和predict(),无需额外配置。
内容的提问来源于stack exchange,提问作者tensor
相关产品推荐
相关产品推荐

