如何解决Libtorch(C++)加载MNIST模型推理时‘forward未定义’错误?
问题:LibTorch保存MNIST模型后加载提示“forward方法未定义”
问题背景
我基于官方MNIST的C++示例代码训练了模型,添加以下代码保存模型:
string model_path = "model.pt"; torch::serialize::OutputArchive output_archive; model.save(output_archive); output_archive.save_to(model_path);
随后在另一个cpp文件中尝试加载并推理:
torch::jit::script::Module model; model=torch::jit::load("model.pt"); model.to(device); std::vector<torch::jit::IValue> inputs; inputs.push_back(torch::ones({1, 1, 28, 28})); at::Tensor output = model.forward(inputs).toTensor(); std::cout << output << std::endl;
代码编译通过,但运行时抛出错误:
terminate called after throwing an instance of 'c10::Error' what(): Method 'forward' is not defined. Exception raised from get_method at .../include/torch/csrc/jit/api/object.h:111 (most recent call first): frame #0: c10::Error::Error(c10::SourceLocation, std::string) + 0x57 (0x7f8e0163f897 in .../lib/libc10.so) frame #1: c10::detail::torchCheckFail(char const*, char const*, unsigned int, std::string const&) + 0x64 (0x7f8e015efb25 in .../lib/libc10.so) frame #2: <unknown function> + 0xa729 (0x560eaa26e729 in ./example-app) frame #3: <unknown function> + 0xa8b1 (0x560eaa26e8b1 in ./example-app) frame #4: <unknown function> + 0x5355 (0x560eaa269355 in ./example-app) frame #5: __libc_start_main + 0xe7 (0x7f8d8eea0c87 in /lib/x86_64-linux-gnu/libc.so.6) frame #6: <unknown function> + 0x4bca (0x560eaa268bca in ./example-app)
我尝试过用TORCH_MODULE创建模块持有者保存模型,但未成功,请问该如何解决?
解决方案
问题核心是保存与加载的方式不匹配:你用torch::serialize::OutputArchive保存的是普通Module的状态字典,而torch::jit::load仅支持加载TorchScript格式的模型,两者无法兼容,因此会报forward方法未定义的错误。
以下是两种可行的解决方法:
方案1:转换为TorchScript模型保存(推荐用于推理场景)
训练完成后,将模型转换为TorchScript模块再保存,这样加载时就能直接用torch::jit::load调用:
// 训练完成后执行保存逻辑 string model_path = "model.pt"; // 用随机输入trace模型,生成TorchScript模块 torch::jit::script::Module ts_module = torch::jit::trace(model, torch::ones({1, 1, 28, 28})); // 保存TorchScript模型 ts_module.save(model_path);
加载代码可直接沿用你原来的实现,不会再出现错误。
方案2:状态字典匹配加载(适合继续训练或保持原模块结构)
如果不需要转换为TorchScript,需要保持保存与加载的状态字典方式一致:
保存代码(与你原代码一致):
string model_path = "model.pt"; torch::serialize::OutputArchive output_archive; model.save(output_archive); output_archive.save_to(model_path);
加载代码:
必须先定义和训练时完全一致的模型结构(比如原示例中的Net类),再加载状态字典:
// 重新定义训练时的MNIST模型结构 Net model; model.to(device); // 加载状态字典 torch::serialize::InputArchive input_archive; input_archive.load_from("model.pt"); model.load(input_archive); // 直接调用模型的forward方法推理 std::vector<torch::jit::IValue> inputs; inputs.push_back(torch::ones({1, 1, 28, 28}).to(device)); at::Tensor output = model.forward(inputs).toTensor(); std::cout << output << std::endl;
关于TORCH_MODULE的正确用法
如果之前尝试TORCH_MODULE未成功,是因为缺少完整的Holder类操作流程:
- 在
Net类定义完成后添加Holder声明:
TORCH_MODULE(NetHolder);
- 保存时用Holder包裹模型:
NetHolder model_holder(model); torch::serialize::OutputArchive output_archive; model_holder.save(output_archive); output_archive.save_to(model_path);
- 加载时通过Holder恢复模型:
NetHolder model_holder; torch::serialize::InputArchive input_archive; input_archive.load_from("model.pt"); model_holder.load(input_archive); // 取出实际模型对象 Net model = *model_holder; model.to(device);
这种方式本质还是状态字典的保存逻辑,需要依赖相同的模型结构,不如TorchScript方案直接通用。
内容的提问来源于stack exchange,提问作者John
相关产品推荐
相关产品推荐

