patdev's picture
Add zero-copy TensorRT PyTorch runtime and AniGen hybrid loader
c8091c9 verified
Raw
History Blame Contribute Delete
2.65 kB
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