diff --git a/orttraining/orttraining/python/training/ortmodule/_graph_execution_manager.py b/orttraining/orttraining/python/training/ortmodule/_graph_execution_manager.py index 6d5ba7cdb0..1afa9e8c51 100644 --- a/orttraining/orttraining/python/training/ortmodule/_graph_execution_manager.py +++ b/orttraining/orttraining/python/training/ortmodule/_graph_execution_manager.py @@ -266,7 +266,7 @@ class GraphExecutionManager(GraphExecutionInterface): self._set_device_from_module(inputs, kwargs) self._onnx_models.exported_model = self._get_exported_model( - *inputs, **kwargs) + schema, *inputs, **kwargs) _cpp_ext._load_aten_op_executor_cpp_extension_if_needed( self._onnx_models.exported_model) if self._debug_options.save_onnx_models.save: @@ -280,8 +280,8 @@ class GraphExecutionManager(GraphExecutionInterface): return True - def _get_exported_model(self, *inputs, **kwargs): - '''Exports PyTorch `self._flattened_module` to ONNX for inferencing or training, using `*inputs` as input + def _get_exported_model(self, input_schema, *inputs, **kwargs): + '''Exports PyTorch `self._flattened_module` to ONNX for inferencing or training, using `*inputs` and `**kwargs` as input TODO: How to support dynamic axes? Dimensions are determined by samples ''' @@ -289,6 +289,7 @@ class GraphExecutionManager(GraphExecutionInterface): # Setup dynamic axes for onnx model self._input_info = _io.parse_inputs_for_onnx_export(self._module_parameters, None, + input_schema, inputs, kwargs) output_names, output_dynamic_axes, self._module_output_schema = \ diff --git a/orttraining/orttraining/python/training/ortmodule/_io.py b/orttraining/orttraining/python/training/ortmodule/_io.py index 9a1e89acc5..dbf59bc654 100644 --- a/orttraining/orttraining/python/training/ortmodule/_io.py +++ b/orttraining/orttraining/python/training/ortmodule/_io.py @@ -85,7 +85,7 @@ class _InputInfo(object): dynamic_axes=None, schema=None, num_positionals=0, - num_positionals_non_none=0, + num_expanded_positionals_non_none=0, keyword_names=None): self.names = names self.shape = shape @@ -93,19 +93,19 @@ class _InputInfo(object): self.dynamic_axes = dynamic_axes if dynamic_axes else {} self.schema = schema if schema else [] self.num_positionals = num_positionals - self.num_positionals_non_none = num_positionals_non_none + self.num_expanded_positionals_non_none = num_expanded_positionals_non_none self.keyword_names = keyword_names def __repr__(self) -> str: return f'''_InputInfo class: - \tNames: {self.names} - \tShape: {self.shape} - \tRequire gradient: {self.require_grad_names} - \tDynamic axes: {self.dynamic_axes} - \tSchema: {self.schema} - \t#Positionals (total): {self.num_positionals} - \t#Positionals (non-None): {self.num_positionals_non_none} - \tKeyword names: {self.keyword_names}''' + \tNames: {self.names} + \tShape: {self.shape} + \tRequire gradient: {self.require_grad_names} + \tDynamic axes: {self.dynamic_axes} + \tSchema: {self.schema} + \t#Positionals (total): {self.num_positionals} + \t#Expanded Positionals (non-None): {self.num_expanded_positionals_non_none} + \tKeyword names: {self.keyword_names}''' def flatten(self, args, kwargs, device): '''Flatten args and kwargs in a single tuple of tensors with strict ordering''' @@ -114,13 +114,18 @@ class _InputInfo(object): ret += [_PrimitiveType.get_tensor(kwargs[name], device) if _PrimitiveType.is_primitive_type(kwargs[name]) else kwargs[name] for name in self.names if name in kwargs] + # if kwargs is empty, append an empty dictionary at the end of the sample inputs to make exporter + # happy. This is because the exporter is confused with kwargs and dictionary inputs otherwise. + if not kwargs: + ret.append({}) + return ret def unflatten(self, flat_args): '''Unflatten tuple of tensors into args and kwargs''' args = tuple(flat_args[:self.num_positionals]) - kwargs = {name: arg for name, arg in zip(self.names[self.num_positionals_non_none:], flat_args[self.num_positionals:]) \ + kwargs = {name: arg for name, arg in zip(self.names[self.num_expanded_positionals_non_none:], flat_args[self.num_positionals:]) \ if name in self.keyword_names} return args, kwargs @@ -141,6 +146,11 @@ def _combine_input_buffers_initializers(params, onnx_input_names, input_info, bu # each element of the list is an input by itself for inp in current_input: _expand_inputs(inp, non_none_inputs) + elif isinstance(current_input, abc.Mapping): + # If the input is a mapping (like a dict), expand the dict so that + # each element of the dict is an input by itself + for _, val in current_input.items(): + _expand_inputs(val, non_none_inputs) elif current_input is not None: # else just collect all the non none inputs within non_none_inputs non_none_inputs.append(current_input) @@ -311,24 +321,26 @@ def _extract_schema(data): elif isinstance(data, torch.Tensor): return _TensorStub(dtype=str(data.dtype), shape_dims=len(data.size())) + # Instead of replacing the tensor with a stub in the original user input, build the stubbed_schema + # from scratch from the user input. + stubbed_schema = None if isinstance(data, abc.Sequence) and not isinstance(data, str): sequence_type = type(data) - data = list(data) - for idx in range(len(data)): - data[idx] = _extract_schema(data[idx]) + stubbed_schema = [_extract_schema(val) for val in data] try: # namedtuple can be created by passing the list sequence to method _make - data = sequence_type._make(data) + stubbed_schema = sequence_type._make(stubbed_schema) except AttributeError: # If attribute error encountered, create the sequence directly - data = sequence_type(data) + stubbed_schema = sequence_type(stubbed_schema) elif isinstance(data, abc.Mapping): - for key in sorted(data): - data[key] = _extract_schema(data[key]) + dict_type = type(data) + stubbed_schema = {key: _extract_schema(data[key]) for key in data} + stubbed_schema = dict_type(**stubbed_schema) else: raise wrap_exception(ORTModuleIOError, TypeError(f'ORTModule does not support the following model data type {type(data)}')) - return data + return stubbed_schema def _parse_outputs_and_extract_names_and_dynamic_axes(module_output): @@ -408,7 +420,7 @@ class _FlattenedModule(torch.nn.Module): return _transform_output_to_flat_tuple(self._original_module(*new_args, **new_kwargs)) -def parse_inputs_for_onnx_export(all_input_parameters, onnx_graph, inputs, kwargs): +def parse_inputs_for_onnx_export(all_input_parameters, onnx_graph, schema, inputs, kwargs): def _add_dynamic_shape(name, input): dynamic_axes[name] = {} @@ -417,21 +429,35 @@ def parse_inputs_for_onnx_export(all_input_parameters, onnx_graph, inputs, kwarg return dynamic_axes def _add_input(name, input, onnx_graph, onnx_graph_input_names): - if input is None: - # Drop all None inputs. - return + """Returns number of expanded non none inputs that _add_input processed""" + if input is None: + # Drop all None inputs and return 0. + return 0 + + num_expanded_non_none_inputs = 0 if isinstance(input, abc.Sequence): # If the input is a sequence (like a list), expand the list so that # each element of the list is an input by itself. for i, val in enumerate(input): # Name each input with the index appended to the original name of the # argument. - _add_input(f"{name}_{i}", val, onnx_graph, onnx_graph_input_names) + num_expanded_non_none_inputs += \ + _add_input(f"{name}_{i}", val, onnx_graph, onnx_graph_input_names) # Return here since the list by itself is not a valid input. # All the elements of the list have already been added as inputs individually. - return + return num_expanded_non_none_inputs + elif isinstance(input, abc.Mapping): + # If the input is a mapping (like a dict), expand the dict so that + # each element of the dict is an input by itself. + for key, val in input.items(): + num_expanded_non_none_inputs += \ + _add_input(f"{name}_{key}", val, onnx_graph, onnx_graph_input_names) + + # Return here since the dict by itself is not a valid input. + # All the elements of the dict have already been added as inputs individually. + return num_expanded_non_none_inputs # InputInfo should contain all the names irrespective of whether they are # a part of the onnx graph or not. @@ -443,6 +469,9 @@ def parse_inputs_for_onnx_export(all_input_parameters, onnx_graph, inputs, kwarg dynamic_axes.update(_add_dynamic_shape(name, input)) input_shape.append(list(input.size())) + # A single input non none input was processed, return 1 + return 1 + # Ignore optional inputs explicitly specified as None # ONNX exporter may remove unused inputs onnx_graph_input_names = [] @@ -454,6 +483,7 @@ def parse_inputs_for_onnx_export(all_input_parameters, onnx_graph, inputs, kwarg input_names_require_grad = [] input_shape = [] var_positional_idx = 0 + num_expanded_non_none_positional_inputs = 0 for input_idx, input_parameter in enumerate(all_input_parameters): if input_parameter.kind == inspect.Parameter.VAR_POSITIONAL: @@ -463,7 +493,8 @@ def parse_inputs_for_onnx_export(all_input_parameters, onnx_graph, inputs, kwarg name = f'{input_parameter.name}_{var_positional_idx}' var_positional_idx += 1 inp = inputs[args_i] - _add_input(name, inp, onnx_graph, onnx_graph_input_names) + num_expanded_non_none_positional_inputs += \ + _add_input(name, inp, onnx_graph, onnx_graph_input_names) elif input_parameter.kind == inspect.Parameter.POSITIONAL_ONLY or\ input_parameter.kind == inspect.Parameter.POSITIONAL_OR_KEYWORD or\ input_parameter.kind == inspect.Parameter.KEYWORD_ONLY: @@ -471,27 +502,32 @@ def parse_inputs_for_onnx_export(all_input_parameters, onnx_graph, inputs, kwarg name = input_parameter.name inp = None input_idx += var_positional_idx + is_positional = True if input_idx < len(inputs) and inputs[input_idx] is not None: inp = inputs[input_idx] elif name in kwargs and kwargs[name] is not None: inp = kwargs[name] - _add_input(name, inp, onnx_graph, onnx_graph_input_names) + is_positional = False + num_expanded_non_none_inputs_local = \ + _add_input(name, inp, onnx_graph, onnx_graph_input_names) + if is_positional: + num_expanded_non_none_positional_inputs += num_expanded_non_none_inputs_local elif input_parameter.kind == inspect.Parameter.VAR_KEYWORD: # **kwargs is always the last argument of forward() for name,inp in kwargs.items(): if name not in input_names: _add_input(name, inp, onnx_graph, onnx_graph_input_names) - # Shallow copy is ok as we need the data structure, not the content - schema = _extract_schema({'args': copy.copy(inputs), 'kwargs': copy.copy(kwargs)}) + # input_names have been expanded so to get the correct number of non none + # positional names, we need to collect the num_expanded_non_none_positional_inputs. return _InputInfo(names=input_names, shape=input_shape, require_grad_names=input_names_require_grad, dynamic_axes=dynamic_axes, schema=schema, num_positionals=len(inputs), - num_positionals_non_none=len([i for i in inputs if i is not None]), + num_expanded_positionals_non_none=num_expanded_non_none_positional_inputs, keyword_names=kwargs.keys()) diff --git a/orttraining/orttraining/python/training/ortmodule/_training_manager.py b/orttraining/orttraining/python/training/ortmodule/_training_manager.py index e3a33be9c4..48041c2424 100644 --- a/orttraining/orttraining/python/training/ortmodule/_training_manager.py +++ b/orttraining/orttraining/python/training/ortmodule/_training_manager.py @@ -73,8 +73,12 @@ class TrainingManager(GraphExecutionManager): # If model was exported, then initialize the graph builder self._initialize_graph_builder(training=True) + # since the schema was just extracted while trying to export the model and it was either + # saved to self._input_info.schema or checked for equality with the self._input_info.schema + # it should not need to be updated again. Pass it inside parse_inputs_for_onnx_export. input_info = _io.parse_inputs_for_onnx_export(self._module_parameters, self._onnx_models.exported_model, + self._input_info.schema, inputs, kwargs) diff --git a/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py b/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py index 68d7821e71..d722c93b4b 100644 --- a/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py +++ b/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py @@ -2891,6 +2891,7 @@ def test_unused_parameters_does_not_unnecessarily_reinitialize(model): input_info = _io.parse_inputs_for_onnx_export(training_manager._module_parameters, training_manager._onnx_models.exported_model, + training_manager._input_info.schema, x, {}) @@ -3138,3 +3139,173 @@ def test_debug_options_log_level_validation_fails_on_type_mismatch(): with pytest.raises(Exception) as ex_info: _ = DebugOptions(log_level=log_level) assert f"Expected log_level of type LogLevel, got {type(log_level)}." in str(ex_info.value) + +def test_ortmodule_dict_input(): + class DictNet(torch.nn.Module): + def __init__(self): + super(DictNet, self).__init__() + self.dummy = torch.nn.Parameter(torch.FloatTensor([0])) + + def forward(self, batch): + b = batch['one_value'] + a = batch['two_value'] + return self.dummy + a + b + + device = 'cuda' + N, D_in, H, D_out = 64, 784, 500, 10 + pt_model = DictNet().to(device) + ort_model = ORTModule(copy.deepcopy(pt_model)) + x = {'one_value': torch.randn(N, D_in, device=device), 'two_value': torch.randn(N, D_in, device=device)} + x_copy = copy.deepcopy(x) + + _test_helpers.assert_values_are_close(pt_model(x), ort_model(x_copy)) + +def test_ortmodule_dict_input_with_unused_values(): + class DictNet(torch.nn.Module): + def __init__(self): + super(DictNet, self).__init__() + self.dummy = torch.nn.Parameter(torch.FloatTensor([0])) + + def forward(self, batch): + b = batch['b'] + a = batch['a'] + return self.dummy + a + + device = 'cuda' + N, D_in, H, D_out = 64, 784, 500, 10 + pt_model = DictNet().to(device) + ort_model = ORTModule(copy.deepcopy(pt_model)) + x = {'a': torch.randn(N, D_in, device=device), 'b': torch.randn(N, D_in, device=device)} + x_copy = copy.deepcopy(x) + + _test_helpers.assert_values_are_close(pt_model(x), ort_model(x_copy)) + +def test_ortmodule_dict_input_with_none_values(): + class DictNet(torch.nn.Module): + def __init__(self): + super(DictNet, self).__init__() + self.dummy = torch.nn.Parameter(torch.FloatTensor([0])) + + def forward(self, batch): + b = batch['b'] + a = batch['a'] if batch['a'] else torch.FloatTensor([2.0]).cuda() + return self.dummy + a + b + + device = 'cuda' + N, D_in, H, D_out = 64, 784, 500, 10 + pt_model = DictNet().to(device) + ort_model = ORTModule(copy.deepcopy(pt_model)) + x = {'a': None, 'b': torch.randn(N, D_in, device=device)} + x_copy = copy.deepcopy(x) + + _test_helpers.assert_values_are_close(pt_model(x), ort_model(x_copy)) + +def test_ortmodule_dict_input_with_nested_values(): + class DictNet(torch.nn.Module): + def __init__(self): + super(DictNet, self).__init__() + self.dummy = torch.nn.Parameter(torch.FloatTensor([0])) + + def forward(self, batch): + a = batch['one_value'] + b = batch['two_value']['three_value'] + c = batch['two_value']['four_value'] + d = batch['five_value']['six_value'] + e = batch['five_value']['seven_value']['eight_value'] + return self.dummy + a + b + c + d + e + + device = 'cuda' + N, D_in, H, D_out = 64, 784, 500, 10 + pt_model = DictNet().to(device) + ort_model = ORTModule(copy.deepcopy(pt_model)) + x = { + 'one_value': torch.randn(N, D_in, device=device), + 'two_value': { + 'three_value': torch.randn(N, D_in, device=device), + 'four_value': torch.randn(N, D_in, device=device) + }, + 'five_value': { + 'six_value': torch.randn(N, D_in, device=device), + 'seven_value': { + 'eight_value': torch.randn(N, D_in, device=device) + } + } + } + x_copy = copy.deepcopy(x) + + _test_helpers.assert_values_are_close(pt_model(x), ort_model(x_copy)) + +def test_ortmodule_list_dict_input_with_nested_values(): + class ListDictNet(torch.nn.Module): + def __init__(self): + super(ListDictNet, self).__init__() + self.dummy = torch.nn.Parameter(torch.FloatTensor([3])) + + def forward(self, batch): + a = batch['one_value'][0] + b = batch['two_value'][0] + c = batch['two_value'][1] + d = batch['three_value'][0] + e = batch['three_value'][1]['four_value'] + return self.dummy + a + b + c + d + e + + device = 'cuda' + N, D_in, H, D_out = 64, 784, 500, 10 + pt_model = ListDictNet().to(device) + ort_model = ORTModule(copy.deepcopy(pt_model)) + x = { + 'one_value': [torch.randn(N, D_in, device=device)], + 'two_value': [torch.randn(N, D_in, device=device), torch.randn(N, D_in, device=device)], + 'three_value': [ + torch.randn(N, D_in, device=device), + { + 'four_value': torch.randn(N, D_in, device=device) + } + ] + } + x_copy = copy.deepcopy(x) + + _test_helpers.assert_values_are_close(pt_model(x), ort_model(x_copy)) + +def test_ortmodule_list_dict_input_with_kwargs_and_registered_buffer(): + class ListDictKwargsNet(torch.nn.Module): + def __init__(self, N, D_in): + super(ListDictKwargsNet, self).__init__() + self.register_buffer("buffer", torch.ones(N, D_in, device='cuda')) + self.dummy = torch.nn.Parameter(torch.FloatTensor([3])) + + def forward(self, batch, **kwargs): + a = batch['one_value'][0] + b = batch['two_value'][0] + c = batch['two_value'][1] + d = batch['three_value'][0] + e = batch['three_value'][1]['four_value'] + out = self.buffer + self.dummy + a + b + c + d + e + if kwargs: + if 'kwargs_0' in kwargs: + out += kwargs['kwargs_0'] + if 'kwargs_1' in kwargs: + out += torch.matmul(kwargs['kwargs_0'], kwargs['kwargs_1']) + + return out + + device = 'cuda' + N, D_in, H, D_out = 64, 784, 500, 10 + pt_model = ListDictKwargsNet(N, D_in).to(device) + ort_model = ORTModule(copy.deepcopy(pt_model), DebugOptions(save_onnx=True, onnx_prefix='kwargsanddict')) + x = { + 'one_value': [torch.randn(N, D_in, device=device)], + 'two_value': [torch.randn(N, D_in, device=device), torch.randn(N, D_in, device=device)], + 'three_value': [ + torch.randn(N, D_in, device=device), + { + 'four_value': torch.randn(N, D_in, device=device) + } + ] + } + x_copy = copy.deepcopy(x) + kwargs_input = {'kwargs_0' : torch.randn(N, D_in, device=device), + 'kwargs_1' : torch.randn(D_in, D_in, device=device)} + kwargs_input_copy = copy.deepcopy(kwargs_input) + + _test_helpers.assert_values_are_close(pt_model(x, **kwargs_input), ort_model(x_copy, **kwargs_input_copy))