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