LibTorch中torch::jit::trace无法使用?是代码问题还是功能缺失?
解决LibTorch中torch::jit::trace调用报错问题
问题原因
你遇到的错误并非torch::jit::trace在LibTorch中不可用,而是模型对象的传递方式不符合函数要求。LibTorch的torch::jit::trace需要接收torch::nn::Module的引用或智能指针,而你代码中直接传递的MyNet实例是ModuleHolder类型(由TORCH_MODULE宏生成的包装类),编译器无法匹配到对应的重载,因此报错“不存在trace成员”。
修正方案
修改torch::jit::trace的调用参数,传递模型内部的Module实例即可,有两种写法:
写法1:解引用ModuleHolder
将trace调用改为:
auto traced_model = torch::jit::trace(*model, example_input);
写法2:获取Module的智能指针
也可以通过ptr()方法获取内部指针:
auto traced_model = torch::jit::trace(model.ptr(), example_input);
修正后的完整代码
#include <torch/torch.h> #include <torch/script.h> #include <iostream> struct MyNetImpl : torch::nn::Module { torch::nn::Linear fc{ nullptr }; MyNetImpl(int in_features, int out_features) { fc = register_module("fc", torch::nn::Linear(in_features, out_features)); } torch::Tensor forward(torch::Tensor x) { return torch::relu(fc->forward(x)); } }; TORCH_MODULE(MyNet); int main() { MyNet model(10, 2); model->eval(); // Create an example input auto example_input = torch::randn({ 1, 10 }); // 修正:传递Module实例给trace auto traced_model = torch::jit::trace(*model, example_input); // Save TorchScript module traced_model.save("model_traced.pt"); std::cout << "Traced model saved successfully!\n"; return 0; }
验证Python加载
保存后的模型可以在Python中正常加载,示例代码:
import torch traced_model = torch.jit.load("model_traced.pt") print(traced_model) # 测试推理 input_tensor = torch.randn(1, 10) output = traced_model(input_tensor) print(output)
内容的提问来源于stack exchange,提问作者BIS steve
相关产品推荐
相关产品推荐

