From 9ba9da0c956265a245b0753e6ca8804bd2fe7e8c Mon Sep 17 00:00:00 2001 From: Thiago Crepaldi Date: Fri, 30 Apr 2021 13:50:23 -0700 Subject: [PATCH] Fix unused registered buffers issue on ORTModule (#7525) --- .../python/training/ortmodule/_io.py | 11 ++++--- .../python/orttraining_test_ortmodule_api.py | 32 +++++++++++++++++++ 2 files changed, 38 insertions(+), 5 deletions(-) diff --git a/orttraining/orttraining/python/training/ortmodule/_io.py b/orttraining/orttraining/python/training/ortmodule/_io.py index b4e1ce866b..d47beae74d 100644 --- a/orttraining/orttraining/python/training/ortmodule/_io.py +++ b/orttraining/orttraining/python/training/ortmodule/_io.py @@ -65,7 +65,7 @@ def _combine_input_buffers_initializers(param_names, onnx_input_names, input_inf # User inputs non_none_inputs = [inp for inp in inputs if inp is not None] - named_buffers_iter = iter(buffer_names) + buffer_names_dict = {buffer_name: inp for buffer_name, inp in buffer_names} result = [] for input_idx, name in enumerate(onnx_input_names): @@ -82,16 +82,17 @@ def _combine_input_buffers_initializers(param_names, onnx_input_names, input_inf elif input_idx >= len(non_none_inputs): # Registered buffers are translated to user_input+initializer in ONNX - buffer_name, inp = next(named_buffers_iter) - assert buffer_name == name, f'Input name {name} expected, but {buffer_name} found!' + try: + inp = buffer_names_dict[name] + except KeyError: + raise KeyError(f'Registered buffer name {name} not found.') if inp is not None: result.append(inp) else: raise RuntimeError(f'Input is present in ONNX graph but not provided: {name}.') # Initializers - for param in param_names: - result.append(param[1]) + result.extend([param[1] for param in param_names]) return result diff --git a/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py b/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py index f3ca5b8216..9681b4c6f0 100644 --- a/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py +++ b/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py @@ -1585,6 +1585,38 @@ def test_model_with_registered_buffers(): output = ort_model(x) assert output is not None + +def test_model_with_unused_registered_buffers(): + class UnusedBufferNet(torch.nn.Module): + def __init__(self, input_size, hidden_size, num_classes): + super(UnusedBufferNet, self).__init__() + + self.fc1 = torch.nn.Linear(input_size, hidden_size) + self.relu = torch.nn.ReLU() + self.fc2 = torch.nn.Linear(hidden_size, num_classes) + self.register_buffer("buffer1s", torch.ones(num_classes)) + self.register_buffer("buffer2s", 1+torch.ones(num_classes)) + self.register_buffer("buffer3s", 2+torch.ones(num_classes)) + + def forward(self, input1): + out = self.fc1(input1) + out = self.relu(out) + out = self.fc2(out) + out += self.buffer3s + return out + device = 'cuda' + + N, D_in, H, D_out = 64, 784, 500, 10 + model = UnusedBufferNet(D_in, H, D_out).to(device) + ort_model = ORTModule(model) + # Check that the original forward signature is preserved. + assert signature(model.forward) == signature(ort_model.forward) + x = torch.randn(N, D_in, device=device) + # Make sure model runs without any exception + output = ort_model(x) + assert output is not None + + def test_model_with_constant_and_registered_parameters(): class NeuralNetWithRegisteredParamsWithConstant(torch.nn.Module): def __init__(self, input_size, hidden_size, num_classes):