patdev's picture
Add zero-copy ORT CUDA fallback before PyTorch
9c30346 verified
Raw
History Blame Contribute Delete
2.76 kB
from __future__ import annotations
from pathlib import Path
import numpy as np
import torch
import onnxruntime as ort
_NP_DTYPES={
'tensor(float)':np.float32,'tensor(float16)':np.float16,'tensor(double)':np.float64,
'tensor(int64)':np.int64,'tensor(int32)':np.int32,'tensor(bool)':np.bool_,
}
_TORCH_DTYPES={np.float32:torch.float32,np.float16:torch.float16,np.float64:torch.float64,np.int64:torch.int64,np.int32:torch.int32,np.bool_:torch.bool}
class TorchORTCudaSession:
"""Static-shape ONNX Runtime CUDA wrapper with zero-copy PyTorch I/O binding."""
def __init__(self,path:str|Path):
so=ort.SessionOptions(); so.graph_optimization_level=ort.GraphOptimizationLevel.ORT_ENABLE_ALL
providers=['CUDAExecutionProvider','CPUExecutionProvider']
self.session=ort.InferenceSession(str(path),sess_options=so,providers=providers)
if self.session.get_providers()[0] != 'CUDAExecutionProvider':
raise RuntimeError(f'CUDAExecutionProvider unavailable: {self.session.get_providers()}')
self.inputs={x.name:x for x in self.session.get_inputs()}; self.outputs={x.name:x for x in self.session.get_outputs()}
def run(self,feeds:dict[str,torch.Tensor]):
io=self.session.io_binding(); keep=[]
for name,meta in self.inputs.items():
t=feeds[name].contiguous(); npdt=_NP_DTYPES[meta.type]; tdt=_TORCH_DTYPES[npdt]
if t.dtype!=tdt: t=t.to(dtype=tdt)
if not t.is_cuda: t=t.cuda(non_blocking=True)
feeds[name]=t; keep.append(t)
io.bind_input(name,'cuda',0,npdt,tuple(t.shape),int(t.data_ptr()))
outs={}
for name,meta in self.outputs.items():
shape=tuple(int(x) for x in meta.shape); npdt=_NP_DTYPES[meta.type]; y=torch.empty(shape,device='cuda',dtype=_TORCH_DTYPES[npdt]); keep.append(y); outs[name]=y
io.bind_output(name,'cuda',0,npdt,shape,int(y.data_ptr()))
self.session.run_with_iobinding(io); io.synchronize_outputs(); return outs
class DinoORT(torch.nn.Module):
def __init__(self,path): super().__init__(); self.engine=TorchORTCudaSession(path); self.device=torch.device('cuda')
def forward(self,pixel_values,is_training=True): return {'x_prenorm':self.engine.run({'pixel_values':pixel_values})['x_prenorm']}
def to(self,*a,**k):return self
def cpu(self):return self
def eval(self):return self
class DsineORT(torch.nn.Module):
def __init__(self,path): super().__init__(); self.engine=TorchORTCudaSession(path); self.device=torch.device('cuda')
def forward(self,image,intrins): return [self.engine.run({'image':image,'intrins':intrins})['normal']]
def to(self,*a,**k):return self
def cpu(self):return self
def eval(self):return self