codekingpro/portable-devtools
114k
1import copy
2import logging
3from collections import OrderedDict
4from collections.abc import Mapping
5from typing import Any
6
7import numpy
8import torch
9from onnx import TensorProto
10
11from onnxruntime import InferenceSession, RunOptions
12
13# Type alias
14ShapeDict = Mapping[str, tuple | list[int]]
15
16logger = logging.getLogger(__name__)
17
18
19class TypeHelper:
20 @staticmethod
21 def get_input_type(ort_session: InferenceSession, name: str) -> str:
22 for _i, input in enumerate(ort_session.get_inputs()):
23 if input.name == name:
24 return input.type
25 raise ValueError(f"input name {name} not found")
26
27 @staticmethod
28 def get_output_type(ort_session, name: str) -> str:
29 for _i, output in enumerate(ort_session.get_outputs()):
30 if output.name == name:
31 return output.type
32
33 raise ValueError(f"output name {name} not found")
34
35 @staticmethod
36 def ort_type_to_numpy_type(ort_type: str):
37 ort_type_to_numpy_type_map = {
38 "tensor(int64)": numpy.int64,
39 "tensor(int32)": numpy.int32,
40 "tensor(float)": numpy.float32,
41 "tensor(float16)": numpy.float16,
42 "tensor(bool)": bool,
43 "tensor(uint8)": numpy.uint8,
44 "tensor(int8)": numpy.int8,
45 "tensor(double)": numpy.float64,
46 "tensor(int16)": numpy.int16,
47 "tensor(uint16)": numpy.uint16,
48 "tensor(uint32)": numpy.uint32,
49 "tensor(uint64)": numpy.uint64,
50 "tensor(complex64)": numpy.complex64,
51 "tensor(complex128)": numpy.complex128,
52 }
53 if ort_type not in ort_type_to_numpy_type_map:
54 raise ValueError(f"{ort_type} not found in map")
55
56 return ort_type_to_numpy_type_map[ort_type]
57
58 @staticmethod
59 def ort_type_to_torch_type(ort_type: str):
60 ort_type_to_torch_type_map = {
61 "tensor(int64)": torch.int64,
62 "tensor(int32)": torch.int32,
63 "tensor(float)": torch.float32,
64 "tensor(float16)": torch.float16,
65 "tensor(bfloat16)": torch.bfloat16,
66 "tensor(bool)": torch.bool,
67 "tensor(uint8)": torch.uint8,
68 "tensor(int8)": torch.int8,
69 "tensor(double)": torch.float64,
70 "tensor(int16)": torch.int16,
71 "tensor(uint16)": torch.uint16,
72 "tensor(uint32)": torch.uint32,
73 "tensor(uint64)": torch.uint64,
74 "tensor(complex64)": torch.complex64,
75 "tensor(complex128)": torch.complex128,
76 "tensor(float8e4m3fn)": torch.float8_e4m3fn,
77 "tensor(float8e4m3fnuz)": torch.float8_e4m3fnuz,
78 "tensor(float8e5m2)": torch.float8_e5m2,
79 "tensor(float8e5m2fnuz)": torch.float8_e5m2fnuz,
80 "tensor(int4)": torch.int4,
81 "tensor(uint4)": torch.uint4,
82 }
83 if ort_type not in ort_type_to_torch_type_map:
84 raise ValueError(f"{ort_type} not found in map")
85
86 return ort_type_to_torch_type_map[ort_type]
87
88 @staticmethod
89 def get_io_onnx_type_map(ort_session: InferenceSession) -> dict[str, int]:
90 """Create a mapping from input/output name to onnx data type"""
91 name_to_onnx_type = {}
92 for input in ort_session.get_inputs():
93 name_to_onnx_type[input.name] = TypeHelper.ort_type_to_onnx_type(input.type)
94
95 for output in ort_session.get_outputs():
96 name_to_onnx_type[output.name] = TypeHelper.ort_type_to_onnx_type(output.type)
97 return name_to_onnx_type
98
99 @staticmethod
100 def ort_type_to_onnx_type(ort_type: str):
101 ort_type_to_onnx_type_map = {
102 "tensor(int64)": TensorProto.INT64,
103 "tensor(int32)": TensorProto.INT32,
104 "tensor(float)": TensorProto.FLOAT,
105 "tensor(float16)": TensorProto.FLOAT16,
106 "tensor(bfloat16)": TensorProto.BFLOAT16,
107 "tensor(bool)": TensorProto.BOOL,
108 "tensor(uint8)": TensorProto.UINT8,
109 "tensor(int8)": TensorProto.INT8,
110 "tensor(double)": TensorProto.DOUBLE,
111 "tensor(int16)": TensorProto.INT16,
112 "tensor(uint16)": TensorProto.UINT16,
113 "tensor(uint32)": TensorProto.UINT32,
114 "tensor(uint64)": TensorProto.UINT64,
115 "tensor(complex64)": TensorProto.COMPLEX64,
116 "tensor(complex128)": TensorProto.COMPLEX128,
117 "tensor(float8e4m3fn)": TensorProto.FLOAT8E4M3FN,
118 "tensor(float8e4m3fnuz)": TensorProto.FLOAT8E4M3FNUZ,
119 "tensor(float8e5m2)": TensorProto.FLOAT8E5M2,
120 "tensor(float8e5m2fnuz)": TensorProto.FLOAT8E5M2FNUZ,
121 "tensor(float4e2m1)": TensorProto.FLOAT4E2M1,
122 "tensor(int4)": TensorProto.INT4,
123 "tensor(uint4)": TensorProto.UINT4,
124 "tensor(string)": TensorProto.STRING,
125 }
126 if ort_type not in ort_type_to_onnx_type_map:
127 raise ValueError(f"{ort_type} not found in map")
128
129 return ort_type_to_onnx_type_map[ort_type]
130
131 @staticmethod
132 def numpy_type_to_torch_type(numpy_type: numpy.dtype):
133 numpy_type_to_torch_type_map = {
134 numpy.int64: torch.int64,
135 numpy.int32: torch.int32,
136 numpy.float32: torch.float32,
137 numpy.float16: torch.float16,
138 bool: torch.bool,
139 numpy.uint8: torch.uint8,
140 numpy.int8: torch.int8,
141 numpy.float64: torch.float64,
142 numpy.int16: torch.int16,
143 numpy.uint16: torch.uint16,
144 numpy.uint32: torch.uint32,
145 numpy.uint64: torch.uint64,
146 numpy.complex64: torch.complex64,
147 numpy.complex128: torch.complex128,
148 }
149
150 if numpy_type not in numpy_type_to_torch_type_map:
151 raise ValueError(f"{numpy_type} not found in map")
152
153 return numpy_type_to_torch_type_map[numpy_type]
154
155 @staticmethod
156 def torch_type_to_numpy_type(torch_type: torch.dtype):
157 torch_type_to_numpy_type_map = {
158 torch.int64: numpy.int64,
159 torch.int32: numpy.int32,
160 torch.float32: numpy.float32,
161 torch.float16: numpy.float16,
162 torch.bool: bool,
163 torch.uint8: numpy.uint8,
164 torch.int8: numpy.int8,
165 torch.float64: numpy.float64,
166 torch.int16: numpy.int16,
167 torch.uint16: numpy.uint16,
168 torch.uint32: numpy.uint32,
169 torch.uint64: numpy.uint64,
170 torch.complex64: numpy.complex64,
171 torch.complex128: numpy.complex128,
172 }
173
174 if torch_type not in torch_type_to_numpy_type_map:
175 raise ValueError(f"{torch_type} not found in map")
176
177 return torch_type_to_numpy_type_map[torch_type]
178
179 @staticmethod
180 def get_io_numpy_type_map(ort_session: InferenceSession) -> dict[str, numpy.dtype]:
181 """Create a mapping from input/output name to numpy data type"""
182 name_to_numpy_type = {}
183 for input in ort_session.get_inputs():
184 name_to_numpy_type[input.name] = TypeHelper.ort_type_to_numpy_type(input.type)
185
186 for output in ort_session.get_outputs():
187 name_to_numpy_type[output.name] = TypeHelper.ort_type_to_numpy_type(output.type)
188 return name_to_numpy_type
189
190 @staticmethod
191 def get_io_torch_type_map(ort_session: InferenceSession) -> dict[str, torch.dtype]:
192 """Create a mapping from input/output name to torch data type"""
193 name_to_torch_type = {}
194 for input in ort_session.get_inputs():
195 name_to_torch_type[input.name] = TypeHelper.ort_type_to_torch_type(input.type)
196
197 for output in ort_session.get_outputs():
198 name_to_torch_type[output.name] = TypeHelper.ort_type_to_torch_type(output.type)
199 return name_to_torch_type
200
201
202class IOBindingHelper:
203 @staticmethod
204 def get_output_buffers(ort_session: InferenceSession, output_shapes, device):
205 """Returns a dictionary of output name as key, and 1D tensor as value. The tensor has enough space for given shape."""
206 output_buffers = {}
207 for name, shape in output_shapes.items():
208 ort_type = TypeHelper.get_output_type(ort_session, name)
209 torch_type = TypeHelper.ort_type_to_torch_type(ort_type)
210 output_buffers[name] = torch.empty(numpy.prod(shape), dtype=torch_type, device=device)
211 return output_buffers
212
213 @staticmethod
214 def prepare_io_binding(
215 ort_session,
216 input_ids: torch.Tensor,
217 position_ids: torch.Tensor,
218 attention_mask: torch.Tensor,
219 past: list[torch.Tensor],
220 output_buffers,
221 output_shapes,
222 ):
223 """IO binding for a session: bind inputs (input_ids, position_ids, attention_mask, past_*) and outputs."""
224
225 name_to_onnx_type = TypeHelper.get_io_onnx_type_map(ort_session)
226
227 # Bind inputs and outputs to onnxruntime session
228 io_binding = ort_session.io_binding()
229
230 # Bind inputs
231 assert input_ids.is_contiguous()
232 io_binding.bind_input(
233 "input_ids",
234 input_ids.device.type,
235 0,
236 name_to_onnx_type["input_ids"],
237 list(input_ids.size()),
238 input_ids.data_ptr(),
239 )
240
241 if past is not None:
242 for i, past_i in enumerate(past):
243 assert past_i.is_contiguous()
244
245 data_ptr = past_i.data_ptr()
246 if data_ptr == 0:
247 # When past_sequence_length is 0, its data_ptr will be zero. IO Binding asserts that data_ptr shall not be zero.
248 # Here we workaround and pass data pointer of input_ids. Actual data is not used for past so it does not matter.
249 data_ptr = input_ids.data_ptr()
250
251 io_binding.bind_input(
252 f"past_{i}",
253 past_i.device.type,
254 0,
255 name_to_onnx_type[f"past_{i}"],
256 list(past_i.size()),
257 data_ptr,
258 )
259
260 if attention_mask is not None:
261 assert attention_mask.is_contiguous()
262 io_binding.bind_input(
263 "attention_mask",
264 attention_mask.device.type,
265 0,
266 name_to_onnx_type["attention_mask"],
267 list(attention_mask.size()),
268 attention_mask.data_ptr(),
269 )
270
271 if position_ids is not None:
272 assert position_ids.is_contiguous()
273 io_binding.bind_input(
274 "position_ids",
275 position_ids.device.type,
276 0,
277 name_to_onnx_type["position_ids"],
278 list(position_ids.size()),
279 position_ids.data_ptr(),
280 )
281
282 # Bind outputs
283 for output in ort_session.get_outputs():
284 output_name = output.name
285 output_buffer = output_buffers[output_name]
286 logger.debug(f"{output_name} device type={output_buffer.device.type} shape={list(output_buffer.size())}")
287 io_binding.bind_output(
288 output_name,
289 output_buffer.device.type,
290 0,
291 name_to_onnx_type[output_name],
292 output_shapes[output_name],
293 output_buffer.data_ptr(),
294 )
295
296 return io_binding
297
298 @staticmethod
299 def get_outputs_from_io_binding_buffer(ort_session, output_buffers, output_shapes, return_numpy=True):
300 """Copy results to cpu. Returns a list of numpy array."""
301 ort_outputs = []
302 for output in ort_session.get_outputs():
303 output_name = output.name
304 buffer = output_buffers[output_name]
305 shape = output_shapes[output_name]
306 copy_tensor = buffer[0 : numpy.prod(shape)].reshape(shape).clone().detach()
307 if return_numpy:
308 ort_outputs.append(copy_tensor.cpu().numpy())
309 else:
310 ort_outputs.append(copy_tensor)
311 return ort_outputs
312
313
314class CudaSession:
315 """Inference Session with IO Binding for ONNX Runtime CUDA or TensorRT provider"""
316
317 def __init__(self, ort_session: InferenceSession, device: torch.device, enable_cuda_graph=False):
318 self.ort_session = ort_session
319 self.input_names = [input.name for input in self.ort_session.get_inputs()]
320 self.output_names = [output.name for output in self.ort_session.get_outputs()]
321 self.io_name_to_onnx_type = TypeHelper.get_io_onnx_type_map(self.ort_session)
322 self.io_name_to_torch_type = TypeHelper.get_io_torch_type_map(self.ort_session)
323 self.io_binding = self.ort_session.io_binding()
324 self.enable_cuda_graph = enable_cuda_graph
325
326 self.input_tensors = OrderedDict()
327 self.output_tensors = OrderedDict()
328 self.device = device
329
330 # Pairs of input and output names that share the same buffer.
331 self.buffer_sharing: dict[str, str] = {}
332
333 def set_buffer_sharing(self, input_name: str, output_name: str):
334 assert input_name in self.input_names
335 assert output_name in self.output_names
336 self.buffer_sharing[input_name] = output_name
337 self.buffer_sharing[output_name] = input_name
338
339 def __del__(self):
340 del self.input_tensors
341 del self.output_tensors
342 del self.io_binding
343
344 def bind_input_and_buffer_sharing(self, name: str, tensor: torch.Tensor):
345 device_id = tensor.device.index if tensor.device.index is not None else 0
346 tensor_shape = [1] if len(tensor.shape) == 0 else list(tensor.shape)
347
348 self.io_binding.bind_input(
349 name,
350 tensor.device.type,
351 device_id,
352 self.io_name_to_onnx_type[name],
353 tensor_shape,
354 tensor.data_ptr(),
355 )
356
357 if name in self.buffer_sharing:
358 self.io_binding.bind_output(
359 self.buffer_sharing[name],
360 tensor.device.type,
361 device_id,
362 self.io_name_to_onnx_type[name],
363 tensor_shape,
364 tensor.data_ptr(),
365 )
366 self.output_tensors[self.buffer_sharing[name]] = tensor
367
368 def allocate_buffers(self, shape_dict: ShapeDict):
369 """Allocate tensors for I/O Binding"""
370 if self.enable_cuda_graph:
371 for name, shape in shape_dict.items():
372 if name in self.input_names:
373 # Reuse allocated buffer when the shape is same
374 if name in self.input_tensors:
375 if tuple(self.input_tensors[name].shape) == tuple(shape):
376 continue
377 raise RuntimeError("Expect static input shape for cuda graph")
378
379 torch_dtype = self.io_name_to_torch_type[name]
380 tensor = torch.empty(tuple(shape), dtype=torch_dtype).to(device=self.device)
381 self.input_tensors[name] = tensor
382 self.bind_input_and_buffer_sharing(name, tensor)
383
384 for name, shape in shape_dict.items():
385 if name in self.output_names:
386 # Reuse allocated buffer when the shape is same
387 if name in self.output_tensors and tuple(self.output_tensors[name].shape) == tuple(shape):
388 continue
389
390 if name in self.buffer_sharing:
391 continue
392
393 torch_dtype = self.io_name_to_torch_type[name]
394 tensor = torch.empty(tuple(shape), dtype=torch_dtype).to(device=self.device)
395 self.output_tensors[name] = tensor
396
397 self.io_binding.bind_output(
398 name,
399 tensor.device.type,
400 tensor.device.index if tensor.device.index is not None else 0,
401 self.io_name_to_onnx_type[name],
402 list(tensor.size()),
403 tensor.data_ptr(),
404 )
405
406 def infer(self, feed_dict: dict[str, torch.Tensor], run_options: RunOptions = None, synchronize: bool = True):
407 """Bind input tensors and run inference"""
408 for name, tensor in feed_dict.items():
409 assert isinstance(tensor, torch.Tensor) and tensor.is_contiguous()
410 if name in self.input_names:
411 if self.enable_cuda_graph:
412 assert self.input_tensors[name].nelement() == tensor.nelement()
413 assert self.input_tensors[name].dtype == tensor.dtype
414 assert tensor.device.type == "cuda"
415 self.input_tensors[name].copy_(tensor)
416 else:
417 self.bind_input_and_buffer_sharing(name, tensor)
418
419 if synchronize:
420 self.io_binding.synchronize_inputs()
421 self.ort_session.run_with_iobinding(self.io_binding, run_options)
422 self.io_binding.synchronize_outputs()
423 else:
424 self.ort_session.run_with_iobinding(self.io_binding, run_options)
425
426 return self.output_tensors
427
428 @staticmethod
429 def get_cuda_provider_options(device_id: int, enable_cuda_graph: bool, stream: int = 0) -> dict[str, Any]:
430 options = {
431 "device_id": device_id,
432 "arena_extend_strategy": "kSameAsRequested",
433 "enable_cuda_graph": enable_cuda_graph,
434 }
435
436 # Stream is address of a CUDA stream. 0 means the default stream.
437 if stream != 0:
438 options["user_compute_stream"] = str(stream)
439
440 return options
441
442
443class GpuBinding(CudaSession):
444 def __init__(
445 self,
446 ort_session: InferenceSession,
447 device: torch.device,
448 shape_dict: ShapeDict,
449 enable_gpu_graph: bool = False,
450 gpu_graph_id: int = -1,
451 stream: int = 0,
452 buffer_sharing: dict[str, str] | None = None,
453 ):
454 super().__init__(ort_session, device, enable_gpu_graph)
455 if buffer_sharing:
456 for input_name, output_name in buffer_sharing.items():
457 self.set_buffer_sharing(input_name, output_name)
458
459 self.allocate_buffers(shape_dict)
460 self.gpu_graph_id = gpu_graph_id
461 # For cuda graph, we need to keep a copy of shape_dict to check if the shape is same in inference later.
462 self.shape_dict = copy.deepcopy(shape_dict) if enable_gpu_graph else None
463 self.stream = stream
464 # The gpu graph id of last run. It will be saved to image metadata.
465 self.last_run_gpu_graph_id = None
466
467 def get_run_options(self, disable_cuda_graph_in_run: bool = False) -> RunOptions:
468 options = RunOptions()
469
470 gpu_graph_id = -1 if disable_cuda_graph_in_run else self.gpu_graph_id
471
472 options.add_run_config_entry("gpu_graph_id", str(gpu_graph_id))
473
474 self.last_run_gpu_graph_id = gpu_graph_id
475
476 return options
477
478 def infer(self, feed_dict: dict[str, torch.Tensor], disable_cuda_graph_in_run: bool = False):
479 run_options = self.get_run_options(disable_cuda_graph_in_run)
480
481 if self.stream:
482 run_options.add_run_config_entry("disable_synchronize_execution_providers", "1")
483
484 return super().infer(feed_dict, run_options)
485
486
487class GpuBindingManager:
488 """A manager for I/O bindings that support multiple CUDA Graphs.
489 One cuda graph is reused for same input shape. Automatically add a new cuda graph for new input shape.
490 """
491
492 def __init__(self, ort_session: InferenceSession, device: torch.device, stream: int = 0, max_cuda_graphs: int = 1):
493 self.ort_session = ort_session
494 self.device = device
495
496 # Binding supports cuda graphs. For a binding, it is able to disable cuda graph for a specific run.
497 self.graph_bindings = []
498
499 # Binding for not using cuda graph.
500 self.no_graph_binding = None
501
502 self.stream = stream
503
504 self.max_cuda_graphs = max_cuda_graphs
505
506 def get_binding(
507 self,
508 shape_dict: ShapeDict,
509 use_cuda_graph: bool = False,
510 buffer_sharing: dict[str, str] | None = None,
511 ) -> GpuBinding:
512 for gpu_graph_binding in self.graph_bindings:
513 # Found a cuda graph that captured with the same shape
514 if gpu_graph_binding.shape_dict == shape_dict:
515 return gpu_graph_binding
516
517 # Reached the maximum number of cuda graphs. Return a binding without cuda graph.
518 if len(self.graph_bindings) >= self.max_cuda_graphs or (not use_cuda_graph):
519 if self.no_graph_binding is None:
520 self.no_graph_binding = GpuBinding(
521 self.ort_session, self.device, shape_dict, stream=self.stream, buffer_sharing=buffer_sharing
522 )
523 else:
524 self.no_graph_binding.allocate_buffers(shape_dict)
525 return self.no_graph_binding
526
527 # This is a new input shape, create a new cuda graph
528 gpu_graph_binding = GpuBinding(
529 self.ort_session,
530 self.device,
531 shape_dict,
532 enable_gpu_graph=True,
533 gpu_graph_id=len(self.graph_bindings),
534 stream=self.stream,
535 buffer_sharing=buffer_sharing,
536 )
537 self.graph_bindings.append(gpu_graph_binding)
538 return gpu_graph_binding
539 