跳到主要内容
极客日志极客日志面向AI+效率的开发者社区
首页博客GitHub 精选镜像AI 生图工具UI配色美学隐私政策关于联系
搜索内容 / 工具 / 仓库 / 镜像...⌘K搜索
注册
博客列表
C++AI算法

ONNXRuntime CUDA 版源码编译与 C++ 部署 PyTorch 模型教程

ONNXRuntime CUDA 版本源码编译流程包含环境配置、构建脚本执行及系统安装步骤。通过 Python 将 PyTorch 训练好的 SAC 策略模型导出为 ONNX 格式后,利用 C++ 代码调用 ONNX Runtime 进行推理验证,可实现高性能的模型部署。

CodeArtist发布于 2026/2/8更新于 2026/7/2365 浏览
ONNXRuntime CUDA 版源码编译与 C++ 部署 PyTorch 模型教程

1 ONNXRuntime 介绍

近年来,机器学习与深度学习领域迎来了应用层面的激增。例如 SkLearn、PyTorch、TensorFlow、Caffe…众多框架可供选择,模型训练的工具选项已变得极为丰富。与此同时,部署目标的多样性也在不断扩展——涵盖移动设备、桌面 CPU、GPU、TPU 等多种硬件。这一趋势带来了核心挑战:如何为将某一框架训练的模型部署到特定目标设备选择合适的工具。这一点至关重要,它直接关系到模型性能的优化——必须考虑所选框架与目标设备之间的兼容性。

在这里插入图片描述

一种可行的方案是使用 ONNX 将模型转换为目标应用适配的框架。例如,若要将 PyTorch 模型部署到 iPhone,可先将模型转换为 ONNX 格式,再从 ONNX 转换为 Core ML(苹果开发的机器学习框架)。另一种方案是利用 ONNX Runtime——微软开发的推理引擎。它以 ONNX 计算静态图为输入,在推理阶段运行该图。具体来说,ONNX Runtime 接收 ONNX 模型作为输入,将完整的计算静态图保留在内存中;随后,将该图划分为可独立管理的子图;最终,将这些子图分配给执行提供者(如 CPU、GPU 等)。

2 ONNXRuntime 编译安装

按照以下步骤执行:

下载指定版本源码

git clone --recursive https://github.com/Microsoft/onnxruntime
cd onnxruntime/
git checkout v1.12.0

这里需要考虑 CUDA 的版本来选择合适的 ONNXRuntime 分支,详情可以参考 ONNXRuntime 文档。

执行编译脚本

./build.sh --skip_tests --use_cuda --config Release --build_shared_lib --parallel --cuda_home /usr/local/cuda-11.3 --cudnn_home /usr/local/cuda-11.3

其中 use_cuda 表示使用 CUDA 版本 ONNXRuntime,cuda_home 和 cudnn_home 均指向 CUDA 安装目录。

执行安装命令

cd /build/Linux/Release
sudo make install

即可安装到系统目录,输出示例如下:

Install the project... -- Install configuration: "Release" -- Installing: /usr/local/include/onnxruntime/core/common ...
-- Installing: /usr/local.so.. -- : .so -- : _test_runner
/lib/libonnxruntime
1.5
2
Installing
/usr/local
/lib/libonnxruntime
Installing
/usr/local
/bin/onnx

3 Pytorch 模型导出 .onnx

以 SAC 算法的部署为例,核心代码如下所示:

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
policy = SACPolicy(
    obs_dim=16,
    action_dim=2,
    device=device,
    **params["policy"],
)
model_path = os.path.abspath(os.path.join(__file__, "../sac_policy/sac_model.pth"))
policy.load_state_dict(torch.load(model_path, map_location=device))

# convert to ONNX model
policy.eval()
dummy_obs = torch.randn(1, policy.obs_dim).to(device)
input_names = ["observation"]
output_names = ["actions"]
output_path = os.path.abspath(os.path.join(__file__, "../sac_model.onnx"))
torch.onnx.export(
    model=policy,
    args=dummy_obs,
    f=output_path,
    input_names=input_names,
    output_names=output_names,
    dynamic_axes={"observation": {0: "batch_size"}, "actions": {0: "batch_size"}},
    verbose=False,
)

可以使用以下脚本验证:

ort_session = ort.InferenceSession(output_path)
dummy_obs_np = dummy_obs.cpu().numpy()
ort_inputs = {ort_session.get_inputs()[0].name: dummy_obs_np}
ort_outputs = ort_session.run(None, ort_inputs)
torch_outputs = policy(dummy_obs)
print(ort_outputs)
for i, (torch_out, ort_out) in enumerate(zip(torch_outputs, ort_outputs)):
    diff = np.abs(torch_out.detach().cpu().numpy() - ort_out)
    print(f"Output {i} 最大差异:{diff.max():.1e}")
print(f"ONNX 模型已成功导出至 {output_path},并验证通过!")

4 C++ 调用 .onnx

以 SAC 算法的部署为例,核心代码如下所示:

SACPolicy::SACPolicy(const std::string& model_path, bool use_gpu) : env_(ORT_LOGGING_LEVEL_WARNING, "SACPolicyInference") {
    Ort::SessionOptions session_options;
    session_options.SetIntraOpNumThreads(1);
    if (use_gpu) {
        OrtCUDAProviderOptions cuda_options;
        cuda_options.device_id = 0;
        cuda_options.cudnn_conv_algo_search = OrtCudnnConvAlgoSearchExhaustive;
        cuda_options.arena_extend_strategy = 0;
        cuda_options.do_copy_in_default_stream = 0;
        session_options.AppendExecutionProvider_CUDA(cuda_options);
    }
    session_ = std::make_unique<Ort::Session>(env_, model_path.c_str(), session_options);
}

SACPredictResult SACPolicy::predict(const std::vector<double>& observation) {
    Ort::AllocatorWithDefaultOptions allocator;
    const auto& input_tensor_info = session_->GetInputTypeInfo(0);
    const auto& input_shape = input_tensor_info.GetTensorTypeAndShapeInfo().GetShape();
    Ort::MemoryInfo memory_info = Ort::MemoryInfo::CreateCpu(OrtAllocatorType::OrtArenaAllocator, OrtMemType::OrtMemTypeDefault);
    
    std::vector<float> input_data(observation.begin(), observation.end());
    std::vector<Ort::Value> input_tensors;
    Ort::Value input_tensor = Ort::Value::CreateTensor<float>(memory_info, input_data.data(), input_data.size(), input_shape.data(), input_shape.size());
    input_tensors.push_back(std::move(input_tensor));
    
    std::vector<const char*> input_names = {"observation"};
    std::vector<const char*> output_names = {"actions"};
    std::vector<Ort::Value> output_tensors = session_->Run(
        Ort::RunOptions{nullptr},
        input_names.data(),
        input_tensors.data(),
        input_tensors.size(),
        output_names.data(),
        output_names.size()
    );
    
    float* output_ptr = output_tensors[0].GetTensorMutableData<float>();
    SACPredictResult result;
    result.linear_velocity = output_ptr[0];
    result.angular_velocity = output_ptr[1];
    return result;
}

目录

  1. 1 ONNXRuntime 介绍
  2. 2 ONNXRuntime 编译安装
  3. 下载指定版本源码
  4. 执行编译脚本
  5. 执行安装命令
  6. 3 Pytorch 模型导出 .onnx
  7. convert to ONNX model
  8. 4 C++ 调用 .onnx
  • 免费图片AI生成工具免费生成了解详情
  • Magick API 一键接入全球大模型注册送1000万token查看
  • 免费图片视频在线生成30秒,将你的创意变成现实开始设计
  • X/Twitter免费视频下载器免登陆无限额度免费视频解析下载了解详情
  • 100+免费在线小游戏爽一把
极客日志微信公众号二维码

微信扫一扫,关注极客日志

微信公众号「极客日志V2」,在微信中扫描左侧二维码关注。展示文案:极客日志V2 zeeklog

更多推荐文章

查看全部
  • 国产数据库新路径:电科金仓融合架构与 AI 实践
  • 程序员转行大模型:核心高薪岗位解析与技能要求
  • Flutter与Web混合开发实践
  • 本地电脑搭建 PyTorch 深度学习环境指南
  • Spring Boot 微服务架构设计与实战
  • NFT 元数据去中心化存储与智能合约集成实战
  • 通义灵码 AI 程序员实操指南:从 IDE 安装到全栈开发落地
  • 双指针算法专题:快乐数与盛水最多的容器
  • Python 基础:列表与元组的区别及操作
  • 如何从零开始学习信息安全与网络安全
  • 前端高频面试题:TypeScript 核心考点与实战
  • OpenClaw 安装百度网页搜索技能 (baidu-web-search)
  • ERNIE-4.5-0.3B 超轻量模型部署与实战测评
  • LLM 驱动的智能体(Agent)应用与实践指南
  • AI 学术写作与查重工具功能解析
  • AI Agent 实战:生产级框架搭建与落地指南
  • C++ 的关键基础:命名空间、引用与函数重载
  • Arduino BLDC 四足仿生穿越机器人设计与控制
  • Rust 与 WebAssembly 实战:在浏览器与 Node.js 运行高性能代码
  • Ollama 底层原理:llama.cpp 与 GGUF 格式解析

相关免费在线工具

  • 加密/解密文本

    使用加密算法(如AES、TripleDES、Rabbit或RC4)加密和解密文本明文。 在线工具,加密/解密文本在线工具,online

  • RSA密钥对生成器

    生成新的随机RSA私钥和公钥pem证书。 在线工具,RSA密钥对生成器在线工具,online

  • Mermaid 预览与可视化编辑

    基于 Mermaid.js 实时预览流程图、时序图等图表,支持源码编辑与即时渲染。 在线工具,Mermaid 预览与可视化编辑在线工具,online

  • 随机西班牙地址生成器

    随机生成西班牙地址(支持马德里、加泰罗尼亚、安达卢西亚、瓦伦西亚筛选),支持数量快捷选择、显示全部与下载。 在线工具,随机西班牙地址生成器在线工具,online

  • Gemini 图片去水印

    基于开源反向 Alpha 混合算法去除 Gemini/Nano Banana 图片水印,支持批量处理与下载。 在线工具,Gemini 图片去水印在线工具,online

  • Base64 字符串编码/解码

    将字符串编码和解码为其 Base64 格式表示形式即可。 在线工具,Base64 字符串编码/解码在线工具,online