ONNX Runtime GPU推理中动态形状输出张量的内存预分配问题
动态形状输出的ONNX Runtime GPU端内存预分配方案
我用onnxruntime-gpu结合NVIDIA DALI做GPU图像预处理,流程能正常运行,但希望全程让数据留在设备端,避免主机与设备间的数据拷贝瓶颈。
ONNX Runtime的IO绑定支持设备端输入输出绑定,但机制偏静态,没法给RetinaNet这类输出形状动态变化的模型预分配内存。手动创建固定形状的输出张量匹配不上实际输出;临时方案会导致数据回传主机,带来额外开销。更新代码后返回OrtValue对象,问题仍未解决。求动态形状输出张量的设备端内存预分配的正确实现方式。
预处理代码
class ImagePipeline(Pipeline): def __init__(self, file_list, batch_size, num_threads, device_id): super(ImagePipeline, self).__init__(batch_size, num_threads, device_id) self.input = ops.readers.File(file_root="", file_list=file_list) self.decode = ops.decoders.Image(device="mixed", output_type=types.RGB) self.resize = ops.Resize(device="gpu", resize_x=800, resize_y=800) self.normalize = ops.CropMirrorNormalize( device="gpu", dtype=types.FLOAT, output_layout=types.NCHW, crop=(800, 800), mean=[0.485 * 255, 0.456 * 255, 0.406 * 255], std=[0.229 * 255, 0.224 * 255, 0.225 * 255], ) def define_graph(self): inputs, labels = self.input() images = self.decode(inputs) images = self.resize(images) images = self.normalize(images) return images, labels
初始推理代码
def run_with_torch_tensors_on_device(x: torch.Tensor, CURR_SIZE: int, torch_type: torch.dtype = torch.float) -> torch.Tensor: binding = session.io_binding() x_tensor = x.contiguous() z_tensor = torch.zeros(CURR_SIZE, 4, dtype=torch_type, device=DEVICE).contiguous() binding.bind_input( name=session.get_inputs()[0].name, device_type=DEVICE_NAME, device_id=DEVICE_INDEX, element_type=np.float32, shape=tuple(x_tensor.shape), buffer_ptr=x_tensor.data_ptr()) binding.bind_output( name=session.get_outputs()[0].name, device_type=DEVICE_NAME, device_id=DEVICE_INDEX, element_type=np.int64, shape=tuple(x_tensor.shape), buffer_ptr=z_tensor.data_ptr()) session.run_with_iobinding(binding) return z_tensor.squeeze(0)
临时方案代码
def run_with_data_on_device(x): x_ortvalue = ort.OrtValue.ortvalue_from_numpy(x) io_binding = session.io_binding() io_binding.bind_input(name=session.get_inputs()[0].name, device_type=x_ortvalue.device_name(), device_id=0, element_type=x.dtype, shape=x_ortvalue.shape(), buffer_ptr=x_ortvalue.data_ptr()) io_binding.bind_output(name=session.get_outputs()[-1].name, device_type=DEVICE_NAME, device_id=DEVICE_INDEX, element_type=x.dtype, shape=x_ortvalue.shape()) session.run_with_iobinding(io_binding) z = io_binding.get_outputs() return z[0]
更新后代码
def run_with_torch_tensors_on_device(x: torch.Tensor, CURR_SIZE: int, torch_type: torch.dtype = torch.float) -> torch.Tensor: binding = session.io_binding() x_tensor = x.contiguous() z_tensor = torch.zeros((CURR_SIZE,91), dtype=torch_type, device=DEVICE).contiguous() binding.bind_input( name=session.get_inputs()[0].name, device_type=DEVICE_NAME, device_id=DEVICE_INDEX, element_type=np.float32, buffer_ptr=x_tensor.data_ptr(), shape=x_tensor.shape) binding.bind_output(session.get_outputs()[-1].name, "cuda") session.run_with_iobinding(binding) ort_output = binding.get_outputs() return ort_output[0]
解决方案
要实现动态形状输出的设备端内存零拷贝,核心是利用ONNX Runtime的动态IO绑定+OrtValue设备端内存复用,结合PyTorch GPU张量管理内存,避免手动预分配的形状不匹配问题:
- 放弃手动固定输出形状:不再提前创建固定shape的torch张量,让ONNX Runtime根据实际输出动态分配设备端内存
- 绑定输出仅指定设备:调用
bind_output时只传设备类型(如"cuda"),不指定固定shape和buffer_ptr - 零拷贝转换为PyTorch张量:通过
torch.as_tensor直接从OrtValue的设备指针创建张量,无数据回传
优化后的推理代码
def run_with_dynamic_output_on_device(x: torch.Tensor) -> torch.Tensor: binding = session.io_binding() x_tensor = x.contiguous() # 绑定输入:直接复用PyTorch GPU张量的设备指针 binding.bind_input( name=session.get_inputs()[0].name, device_type=DEVICE_NAME, device_id=DEVICE_INDEX, element_type=np.float32, shape=x_tensor.shape, buffer_ptr=x_tensor.data_ptr() ) # 绑定输出:仅指定GPU设备,让ONNX Runtime动态分配对应大小的内存 binding.bind_output(session.get_outputs()[-1].name, device_type="cuda", device_id=DEVICE_INDEX) session.run_with_iobinding(binding) # 获取设备端OrtValue,零拷贝转换为PyTorch张量 ort_output = binding.get_outputs()[0] output_tensor = torch.as_tensor(ort_output.numpy(), device=DEVICE) return output_tensor
关键细节
- 动态内存分配:
bind_output不指定shape时,ONNX Runtime会根据模型实际输出的动态形状,在指定GPU上分配对应内存,彻底解决形状不匹配问题 - 零拷贝转换:
torch.as_tensor(ort_output.numpy(), device=DEVICE)不会触发主机拷贝,因为OrtValue本身在GPU上,numpy()仅返回设备内存视图,PyTorch直接复用该内存 - DALI对接优化:DALI的GPU输出可直接用
torch.as_tensor转换为PyTorch张量(零拷贝),直接传入上述推理函数,实现预处理到推理全程GPU内存闭环
内容的提问来源于stack exchange,提问作者JOKKINATOR
相关产品推荐
相关产品推荐

