codekingpro/portable-devtools
114k
1# Copyright (c) Microsoft Corporation. All rights reserved.
2# Licensed under the MIT License.
3
4from __future__ import annotations
5
6import inspect
7from collections import abc
8
9import torch
10
11
12def _parse_inputs_for_onnx_export(all_input_parameters, inputs, kwargs):
13 # extracted from https://github.com/microsoft/onnxruntime/blob/239c6ad3f021ff7cc2e6247eb074bd4208dc11e2/orttraining/orttraining/python/training/ortmodule/_io.py#L433
14
15 def _add_input(name, input):
16 """Returns number of expanded inputs that _add_input processed"""
17
18 if input is None:
19 # Drop all None inputs and return 0.
20 return 0
21
22 num_expanded_non_none_inputs = 0
23 if isinstance(input, abc.Sequence):
24 # If the input is a sequence (like a list), expand the list so that
25 # each element of the list is an input by itself.
26 for i, val in enumerate(input):
27 # Name each input with the index appended to the original name of the
28 # argument.
29 num_expanded_non_none_inputs += _add_input(f"{name}_{i}", val)
30
31 # Return here since the list by itself is not a valid input.
32 # All the elements of the list have already been added as inputs individually.
33 return num_expanded_non_none_inputs
34 elif isinstance(input, abc.Mapping):
35 # If the input is a mapping (like a dict), expand the dict so that
36 # each element of the dict is an input by itself.
37 for key, val in input.items():
38 num_expanded_non_none_inputs += _add_input(f"{name}_{key}", val)
39
40 # Return here since the dict by itself is not a valid input.
41 # All the elements of the dict have already been added as inputs individually.
42 return num_expanded_non_none_inputs
43
44 # InputInfo should contain all the names irrespective of whether they are
45 # a part of the onnx graph or not.
46 input_names.append(name)
47
48 # A single input non none input was processed, return 1
49 return 1
50
51 input_names = []
52 var_positional_idx = 0
53 num_expanded_non_none_positional_inputs = 0
54
55 for input_idx, input_parameter in enumerate(all_input_parameters):
56 if input_parameter.kind == inspect.Parameter.VAR_POSITIONAL:
57 # VAR_POSITIONAL parameter carries all *args parameters from original forward method
58 for args_i in range(input_idx, len(inputs)):
59 name = f"{input_parameter.name}_{var_positional_idx}"
60 var_positional_idx += 1
61 inp = inputs[args_i]
62 num_expanded_non_none_positional_inputs += _add_input(name, inp)
63 elif (
64 input_parameter.kind == inspect.Parameter.POSITIONAL_ONLY
65 or input_parameter.kind == inspect.Parameter.POSITIONAL_OR_KEYWORD
66 or input_parameter.kind == inspect.Parameter.KEYWORD_ONLY
67 ):
68 # All positional non-*args and non-**kwargs are processed here
69 name = input_parameter.name
70 inp = None
71 input_idx += var_positional_idx # noqa: PLW2901
72 is_positional = True
73 if input_idx < len(inputs) and inputs[input_idx] is not None:
74 inp = inputs[input_idx]
75 elif name in kwargs and kwargs[name] is not None:
76 inp = kwargs[name]
77 is_positional = False
78 num_expanded_non_none_inputs_local = _add_input(name, inp)
79 if is_positional:
80 num_expanded_non_none_positional_inputs += num_expanded_non_none_inputs_local
81 elif input_parameter.kind == inspect.Parameter.VAR_KEYWORD:
82 # **kwargs is always the last argument of forward()
83 for name, inp in kwargs.items():
84 if name not in input_names:
85 _add_input(name, inp)
86
87 return input_names
88
89
90def _flatten_module_input(names, args, kwargs):
91 """Flatten args and kwargs in a single tuple of tensors."""
92 # extracted from https://github.com/microsoft/onnxruntime/blob/239c6ad3f021ff7cc2e6247eb074bd4208dc11e2/orttraining/orttraining/python/training/ortmodule/_io.py#L110
93
94 def is_primitive_type(value):
95 return type(value) in {int, bool, float}
96
97 def to_tensor(value):
98 return torch.tensor(value)
99
100 ret = [to_tensor(arg) if is_primitive_type(arg) else arg for arg in args]
101 ret += [
102 to_tensor(kwargs[name]) if is_primitive_type(kwargs[name]) else kwargs[name] for name in names if name in kwargs
103 ]
104
105 # if kwargs is empty, append an empty dictionary at the end of the sample inputs to make exporter
106 # happy. This is because the exporter is confused with kwargs and dictionary inputs otherwise.
107 if not kwargs:
108 ret.append({})
109
110 return tuple(ret)
111
112
113def infer_input_info(module: torch.nn.Module, *inputs, **kwargs):
114 """
115 Infer the input names and order from the arguments used to execute a PyTorch module for usage exporting
116 the model via torch.onnx.export.
117 Assumes model is on CPU. Use `module.to(torch.device('cpu'))` if it isn't.
118
119 Example usage:
120 input_names, inputs_as_tuple = infer_input_info(module, ...)
121 torch.onnx.export(module, inputs_as_type, 'model.onnx', input_names=input_names, output_names=[...], ...)
122
123 :param module: Module
124 :param inputs: Positional inputs
125 :param kwargs: Keyword argument inputs
126 :return: Tuple of ordered input names and input values. These can be used directly with torch.onnx.export as the
127 `input_names` and `inputs` arguments.
128 """
129 module_parameters = inspect.signature(module.forward).parameters.values()
130 input_names = _parse_inputs_for_onnx_export(module_parameters, inputs, kwargs)
131 inputs_as_tuple = _flatten_module_input(input_names, inputs, kwargs)
132
133 return input_names, inputs_as_tuple
134 