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

自定义PyTorch模型转TorchScript后C++调用自定义推理方法报错求助

解决TorchScript模型在C++中调用自定义方法的问题

问题原因

torch.jit.trace()仅会追踪并记录forward方法的执行流程,不会自动识别并转换模型中的其他自定义方法(比如custom_forward_for_inference),因此C++加载模型后无法找到该方法。

解决方案

  • 方法1:使用torch.jit.script()转换模型(推荐)
    torch.jit.script()通过静态代码分析识别模型类的所有方法,配合@torch.jit.export装饰器可直接暴露自定义方法:

    import torch
    import torch.nn as nn
    
    class CustomModel(nn.Module):
        def __init__(self):
            super().__init__()
            # 你的模型初始化逻辑,比如:
            # self.conv = nn.Conv2d(3, 16, kernel_size=3)
        
        def forward(self, x):
            # 你的训练时forward逻辑
            return x
        
        @torch.jit.export  # 标记该方法可被TorchScript导出
        def custom_forward_for_inference(self, x):
            # 你的推理专用逻辑
            return x
    
    # 实例化并转换为ScriptModule
    model = CustomModel()
    scripted_model = torch.jit.script(model)
    # 保存模型
    scripted_model.save("custom_model.pt")
    
  • 方法2:基于trace手动注册自定义方法
    如果因模型含动态逻辑必须使用trace,可在trace后将自定义方法脚本化并注册到模型:

    import torch
    import torch.nn as nn
    
    class CustomModel(nn.Module):
        def __init__(self):
            super().__init__()
            # 初始化逻辑
        
        def forward(self, x):
            # forward逻辑
            return x
        
        def custom_forward_for_inference(self, x):
            # 推理逻辑
            return x
    
    model = CustomModel()
    # 先trace forward方法,输入需匹配模型输入形状
    traced_model = torch.jit.trace(model, torch.randn(1, 3, 224, 224))
    # 将自定义方法转换为ScriptFunction并注册到trace后的模型
    traced_model.custom_forward_for_inference = torch.jit.script(model.custom_forward_for_inference)
    # 保存模型
    traced_model.save("custom_model.pt")
    

C++端调用方式

C++中不能直接通过.调用自定义方法,需通过get_method()获取方法对象后执行:

#include <torch/script.h>
#include <iostream>

int main() {
    try {
        // 加载TorchScript模型
        torch::jit::script::Module loadedModule = torch::jit::load("custom_model.pt");
        
        // 准备输入Tensor(需匹配模型输入形状)
        torch::Tensor inputs = torch::randn({1, 3, 224, 224});
        std::vector<torch::jit::IValue> input_args = {inputs};
        
        // 获取自定义方法并调用
        auto custom_method = loadedModule.get_method("custom_forward_for_inference");
        at::Tensor output = custom_method(input_args).toTensor();
        
        std::cout << "Output shape: " << output.sizes() << std::endl;
    } catch (const c10::Error& e) {
        std::cerr << "Error loading or running model: " << e.what() << std::endl;
        return -1;
    }
    return 0;
}

内容的提问来源于stack exchange,提问作者Rhaenys98

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 19:15:02