File size: 2,650 Bytes
ef79ec1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c8091c9
ef79ec1
 
 
 
 
 
 
c8091c9
ef79ec1
 
 
 
 
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
50
from __future__ import annotations
from pathlib import Path
import torch

class TorchTensorRTEngine:
    """TensorRT 10 engine wrapper using PyTorch CUDA pointers (zero-copy bindings)."""
    def __init__(self, path:str|Path):
        import tensorrt as trt
        self.trt=trt; self.logger=trt.Logger(trt.Logger.ERROR)
        self.runtime=trt.Runtime(self.logger)
        self.engine=self.runtime.deserialize_cuda_engine(Path(path).read_bytes())
        if self.engine is None: raise RuntimeError(f'Could not deserialize {path}')
        self.context=self.engine.create_execution_context()
        self.inputs=[]; self.outputs=[]
        for i in range(self.engine.num_io_tensors):
            n=self.engine.get_tensor_name(i); mode=self.engine.get_tensor_mode(n)
            (self.inputs if mode==trt.TensorIOMode.INPUT else self.outputs).append(n)
    def _torch_dtype(self, dt):
        trt=self.trt
        m={trt.float32:torch.float32,trt.float16:torch.float16,trt.int32:torch.int32,trt.int64:torch.int64,trt.bool:torch.bool}
        if hasattr(trt,'bfloat16'): m[trt.bfloat16]=torch.bfloat16
        return m[dt]
    def run(self, feeds:dict[str,torch.Tensor]):
        for n,t in feeds.items():
            t=t.contiguous(); feeds[n]=t
            self.context.set_input_shape(n, tuple(t.shape)); self.context.set_tensor_address(n, int(t.data_ptr()))
        outs={}
        for n in self.outputs:
            shape=tuple(self.context.get_tensor_shape(n)); dt=self._torch_dtype(self.engine.get_tensor_dtype(n))
            y=torch.empty(shape,device='cuda',dtype=dt); outs[n]=y; self.context.set_tensor_address(n,int(y.data_ptr()))
        ok=self.context.execute_async_v3(stream_handle=torch.cuda.current_stream().cuda_stream)
        if not ok: raise RuntimeError('TensorRT execute_async_v3 failed')
        return outs

class DinoTRT(torch.nn.Module):
    def __init__(self, engine_path): super().__init__(); self.engine=TorchTensorRTEngine(engine_path); self.device=torch.device('cuda')
    def forward(self, pixel_values, is_training=True):
        out=self.engine.run({'pixel_values':pixel_values})['x_prenorm']; return {'x_prenorm':out}
    def to(self,*a,**k): return self
    def cpu(self): return self
    def eval(self): return self

class DsineTRT(torch.nn.Module):
    def __init__(self, engine_path): super().__init__(); self.engine=TorchTensorRTEngine(engine_path); self.device=torch.device('cuda')
    def forward(self, image, intrins):
        out=self.engine.run({'image':image,'intrins':intrins})['normal']; return [out]
    def to(self,*a,**k): return self
    def cpu(self): return self
    def eval(self): return self