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

如何解决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类操作流程:

  1. 在Net类定义完成后添加Holder声明:
TORCH_MODULE(NetHolder);
  1. 保存时用Holder包裹模型:
NetHolder model_holder(model);
torch::serialize::OutputArchive output_archive;
model_holder.save(output_archive);
output_archive.save_to(model_path);
  1. 加载时通过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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 22:10:00