PyTorch模型(含TorchScript)的二进制格式及C++使用咨询
在C++中部署PyTorch/TorchScript模型的工作流与二进制格式说明
模型二进制格式核心说明
TorchScript模型的标准二进制格式为.pt或.pth(注意:.pth也可用于Python端单独保存权重,但TorchScript导出的文件是包含完整计算图与权重的可执行模型),这类文件通过PyTorch内置的序列化机制生成,本质是序列化后的TorchScript模块实例。
完整工作流代码示例
1. Python端导出TorchScript二进制模型
两种主流导出方式,适配不同模型类型:
import torch import torchvision.models as models # 方式一:Trace(适合无动态控制流的模型,如大部分CNN) resnet = models.resnet18(pretrained=True) resnet.eval() # 构造与模型输入维度匹配的示例张量 example_input = torch.randn(1, 3, 224, 224) # 追踪模型生成TorchScript模块 traced_model = torch.jit.trace(resnet, example_input) # 保存为二进制文件 traced_model.save("resnet18_traced.pt") # 方式二:Script(适合含动态分支、循环等控制流的模型) class DynamicModel(torch.nn.Module): def forward(self, x): if x.mean() > 0: return x * 2 else: return x + 1 model = DynamicModel() # 直接编译模型生成TorchScript模块 scripted_model = torch.jit.script(model) scripted_model.save("dynamic_model_scripted.pt")
2. C++端加载并运行模型
需链接LibTorch(PyTorch的C++前端库),核心代码如下:
#include <torch/torch.h> #include <iostream> int main() { // 加载TorchScript二进制模型 torch::jit::script::Module module; try { module = torch::jit::load("resnet18_traced.pt"); } catch (const c10::Error& e) { std::cerr << "模型加载失败: " << e.what() << std::endl; return -1; } // 构造输入张量(维度需与导出时的示例输入一致) std::vector<torch::jit::IValue> inputs; inputs.push_back(torch::randn({1, 3, 224, 224})); // 执行推理并获取输出 at::Tensor output = module.forward(inputs).toTensor(); std::cout << "输出张量形状: " << output.sizes() << std::endl; return 0; }
格式相关关键注意事项
- TorchScript二进制文件是自包含的,包含计算图结构、权重参数、算子依赖,无需Python环境即可在C++中运行。
- 必须保证导出模型的PyTorch版本与C++端使用的LibTorch版本大版本一致(如PyTorch 2.1导出的模型不能用LibTorch 1.13加载),否则会出现格式不兼容错误。
- 若需排查模型结构问题,可在Python中通过
torch.jit.load("model.pt").graph打印计算图,或用torch.jit.inspect(module)查看模块的层级与参数信息。
内容的提问来源于stack exchange,提问作者M46f988b814
相关产品推荐
相关产品推荐

