Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
io_binding_helper.py539 linesDownload Raw Back to transformers
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 
codekingpro/portable-devtools · Team Ai