File size: 2,755 Bytes
9c30346
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
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