2121import argparse
2222import importlib
2323import os
24+ import sys
2425import warnings
26+ from pathlib import Path
2527
26- import tensorrt as trt
2728import torch
2829from mmcv import Config , DictAction
2930from mmcv .utils import import_modules_from_strings
3334from projects .mmdet3d_plugin .datasets .builder import build_dataloader
3435from tqdm import tqdm
3536
36- TRT_TO_TORCH = {
37- trt .DataType .FLOAT : torch .float32 ,
38- trt .DataType .HALF : torch .float16 ,
39- trt .DataType .INT8 : torch .int8 ,
40- trt .DataType .INT32 : torch .int32 ,
41- trt .DataType .BOOL : torch .bool ,
42- trt .DataType .UINT8 : torch .uint8 ,
43- }
44- if int (trt .__version__ .split ("." )[0 ]) >= 10 :
45- TRT_TO_TORCH [trt .DataType .INT64 ] = torch .int64
46-
47- TRT_LOGGER = trt .Logger (trt .Logger .WARNING )
48- trt .init_libnvinfer_plugins (TRT_LOGGER , "" )
49-
50-
51- def aligned_tensor (shape , dtype , device , alignment = 256 ):
52- element_size = torch .empty ((), dtype = dtype ).element_size ()
53- element_count = int (torch .tensor (shape ).prod ().item ())
54- storage = torch .empty (element_count + alignment // element_size , dtype = dtype , device = device )
55- offset_bytes = (- storage .data_ptr ()) % alignment
56- offset = offset_bytes // element_size
57- return storage [offset : offset + element_count ].view (shape )
58-
59-
60- class TensorRTRunner :
61- def __init__ (self , engine_path , state_names = ()):
62- with open (engine_path , "rb" ) as engine_file :
63- engine_bytes = engine_file .read ()
64- self .engine = trt .Runtime (TRT_LOGGER ).deserialize_cuda_engine (engine_bytes )
65- if self .engine is None :
66- raise RuntimeError (f"Failed to deserialize { engine_path } " )
67- self .context = self .engine .create_execution_context ()
68- if self .context is None :
69- raise RuntimeError (f"Failed to create an execution context for { engine_path } " )
70- self .tensor_names = [
71- self .engine .get_tensor_name (index ) for index in range (self .engine .num_io_tensors )
72- ]
73- self .input_shapes = {}
74- self .output_shapes = {}
75- self .tensor_dtypes = {}
76- for name in self .tensor_names :
77- shape = tuple (self .engine .get_tensor_shape (name ))
78- dtype = TRT_TO_TORCH [self .engine .get_tensor_dtype (name )]
79- self .tensor_dtypes [name ] = dtype
80- if self .engine .get_tensor_mode (name ) == trt .TensorIOMode .INPUT :
81- self .input_shapes [name ] = shape
82- else :
83- self .output_shapes [name ] = shape
84-
85- self .state = {}
86- for base_name in state_names :
87- name = self .resolve_name (base_name )
88- if name in self .input_shapes :
89- tensor = aligned_tensor (self .input_shapes [name ], self .tensor_dtypes [name ], "cuda" )
90- tensor .zero_ ()
91- self .state [name ] = tensor
92- self .context .set_tensor_address (name , tensor .data_ptr ())
93- if self .state :
94- torch .cuda .synchronize ()
95-
96- def resolve_name (self , base_name ):
97- if base_name in self .tensor_names :
98- return base_name
99- suffixed_name = f"{ base_name } .1"
100- return suffixed_name if suffixed_name in self .tensor_names else base_name
101-
102- def reset_state (self ):
103- for tensor in self .state .values ():
104- tensor .zero_ ()
105-
106- def prepare_input (self , name , inputs ):
107- shape = self .input_shapes [name ]
108- base_name = name .rsplit (".1" , maxsplit = 1 )[0 ] if name .endswith (".1" ) else name
109- if base_name not in inputs :
110- raise KeyError (f"Missing TensorRT input { base_name } " )
111- value = inputs [base_name ].to (device = "cuda" , dtype = self .tensor_dtypes [name ])
112- if tuple (value .shape ) != shape :
113- if tuple (value .shape [1 :]) == shape :
114- value = value .squeeze (0 )
115- elif tuple (shape [1 :]) == tuple (value .shape ):
116- value = value .unsqueeze (0 )
117- else :
118- raise ValueError (
119- f"Input { base_name } has shape { tuple (value .shape )} , expected { shape } "
120- )
121- return value
122-
123- def __call__ (self , stream , ** inputs ):
124- input_buffers = {}
125- for name , shape in self .input_shapes .items ():
126- if name in self .state :
127- continue
128- value = self .prepare_input (name , inputs )
129- buffer = aligned_tensor (shape , value .dtype , value .device )
130- buffer .copy_ (value )
131- input_buffers [name ] = buffer
132- self .context .set_tensor_address (name , buffer .data_ptr ())
133-
134- outputs = {}
135- for name , shape in self .output_shapes .items ():
136- output = aligned_tensor (shape , self .tensor_dtypes [name ], "cuda" )
137- outputs [name ] = output
138- self .context .set_tensor_address (name , output .data_ptr ())
139-
140- if not self .context .execute_async_v3 (stream .cuda_stream ):
141- raise RuntimeError ("TensorRT execution failed" )
142- stream .synchronize ()
143- return outputs
37+ sys .path .insert (0 , str (Path (__file__ ).resolve ().parents [3 ]))
14438
39+ from examples .onnx_ptq .trt_runner import TensorRTRunner
14540
14641STATE_NAMES = (
14742 "memory_embedding" ,
@@ -154,8 +49,7 @@ def __call__(self, stream, **inputs):
15449
15550class Far3DDecoderRunner (TensorRTRunner ):
15651 def __init__ (self , engine_path , input_callback = None ):
157- super ().__init__ (engine_path , STATE_NAMES )
158- self .input_callback = input_callback
52+ super ().__init__ (engine_path , STATE_NAMES , input_callback )
15953 self .scene_token = None
16054 self .timestamp_offset = None
16155
@@ -175,16 +69,6 @@ def __call__(self, stream, img_metas, timestamp, **inputs):
17569 device = "cuda" ,
17670 )
17771 inputs ["timestamp" ] = (timestamp - self .timestamp_offset ).float ()
178- if self .input_callback :
179- calibration_inputs = {}
180- for name in self .input_shapes :
181- base_name = name .rsplit (".1" , maxsplit = 1 )[0 ] if name .endswith (".1" ) else name
182- if name in self .state :
183- value = self .state [name ]
184- else :
185- value = self .prepare_input (name , inputs )
186- calibration_inputs [base_name ] = value
187- self .input_callback (calibration_inputs )
18872 outputs = super ().__call__ (stream , ** inputs )
18973 for base_name in STATE_NAMES :
19074 input_name = self .resolve_name (base_name )
@@ -287,7 +171,7 @@ def main():
287171 }
288172 }
289173 )
290- if args .max_samples is not None and len (outputs ) = = args .max_samples :
174+ if args .max_samples is not None and len (outputs ) > = args .max_samples :
291175 break
292176
293177 if len (outputs ) < len (dataset ):
0 commit comments