You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.09 16:30:49