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

如何设置ONNX Runtime直接返回Torch Tensor而非NumPy数组?

直接让ONNX Runtime返回Torch Tensor的两种方法

ONNX Runtime默认返回NumPy数组,但可以通过以下两种方式避免手动转换的开销,直接得到Torch Tensor:

方法1:使用ORTModule封装PyTorch模型(推荐)

ORTModule是ONNX Runtime为PyTorch用户提供的官方封装,它会自动将PyTorch模型转换为ONNX格式并优化,推理时直接返回Torch Tensor。

步骤:

  1. 安装依赖:
pip install onnxruntime-training
  1. 加载并封装模型:
import torch
from onnxruntime.training import ORTModule
from super_gradients.training import models

# 加载Super-Gradients的YOLOX-S模型
model = models.get("yolox_s", pretrained_weights="coco")
model = ORTModule(model)
model.eval()

# 推理流程
dataset = MyCostumeDataset(args.path, 'val')
val_dataloader = DataLoader(dataset, batch_size=args.bsize)

for inputs in val_dataloader:
    with torch.no_grad():
        raw_predictions = model(inputs)
        # raw_predictions 是 Torch Tensor 类型,无需额外转换

方法2:直接将ONNX Runtime输出转为Torch Tensor(无额外拷贝)

如果已经有导出好的ONNX模型,可以通过torch.as_tensor直接将ONNX Runtime的输出转为Torch Tensor,尤其是在GPU环境下,能避免数据拷贝开销。

代码示例:

import torch
import onnxruntime as onnxrt

# 创建支持CUDA的Session(CPU环境可去掉CUDAExecutionProvider)
providers = ['CUDAExecutionProvider', 'CPUExecutionProvider']
onnx_session = onnxrt.InferenceSession("yolox_s_640_640.onnx", providers=providers)

dataset = MyCostumeDataset(args.path, 'val')
val_dataloader = DataLoader(dataset, batch_size=args.bsize)

for inputs in val_dataloader:
    # 输入Tensor移到对应设备(GPU/CPU)
    inputs = inputs.to('cuda' if torch.cuda.is_available() else 'cpu')
    onnx_inputs = {onnx_session.get_inputs()[0].name: inputs}
    
    raw_predictions_np = onnx_session.run(None, onnx_inputs)
    # 直接转换为Torch Tensor,共享内存(无额外拷贝)
    raw_predictions = [torch.as_tensor(arr, device=inputs.device) for arr in raw_predictions_np]
    # raw_predictions 是 Torch Tensor 列表

关键说明:

  • torch.as_tensor会根据输入数组的存储位置(CPU/GPU)直接创建对应设备的Tensor,若设备匹配则共享内存,几乎无开销;
  • CPU环境下,torch.as_tensor兼容性比torch.from_numpy更好,支持更多非NumPy数组类型。

内容的提问来源于stack exchange,提问作者Ruslan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 01:15:31