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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 13:05:22