From 1370cbe256842ee8d0b1f63f56b8b87bf48f507f Mon Sep 17 00:00:00 2001 From: Sherlock Date: Wed, 28 Jul 2021 09:55:07 -0700 Subject: [PATCH] [ORTModule] Extract output schema in module's true train/eval mode (#8516) * Extract output schema in module's true train/eval mode --- .../python/training/ortmodule/_io.py | 8 +-- .../python/orttraining_test_ortmodule_api.py | 54 +++++++++++++++++-- 2 files changed, 52 insertions(+), 10 deletions(-) diff --git a/orttraining/orttraining/python/training/ortmodule/_io.py b/orttraining/orttraining/python/training/ortmodule/_io.py index 2f2e767333..7f3fa48888 100644 --- a/orttraining/orttraining/python/training/ortmodule/_io.py +++ b/orttraining/orttraining/python/training/ortmodule/_io.py @@ -488,10 +488,7 @@ def parse_inputs_for_onnx_export(all_input_parameters, onnx_graph, inputs, kwarg def parse_outputs_for_onnx_export_and_extract_schema(module, inputs, kwargs): - - # Do an inference to grab outputs - is_train_mode = module.training - module.eval() + # Perform a forward call to grab outputs output_names = None output_dynamic_axes = None is_deepcopy = False @@ -512,8 +509,7 @@ def parse_outputs_for_onnx_export_and_extract_schema(module, inputs, kwargs): # Parse the output and extract the output_names and output_dynamic_axes to be used for onnx export output_names, output_dynamic_axes = _parse_outputs_and_extract_names_and_dynamic_axes(sample_outputs) - if is_train_mode: - module.train() + output_schema = _extract_schema(sample_outputs) if is_deepcopy: del model_copy diff --git a/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py b/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py index ca9af5265d..27d3beac6f 100644 --- a/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py +++ b/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py @@ -2196,12 +2196,58 @@ def test_unused_layer(): device = torch.device('cuda') N, D_in, H, D_out = 64, 784, 500, 10 - model = NeuralNetSinglePositionalArgument(D_in, H, D_out).to(device) - ort_model = ORTModule(model) + pt_model = Net(D_in, H, D_out).to(device) + ort_model = ORTModule(copy.deepcopy(pt_model)) x = torch.randn(N, D_in, device=device) - output = ort_model(x) - assert output is not None + pt_output = pt_model(x) + ort_output = ort_model(x) + _test_helpers.assert_values_are_close(pt_output, ort_output) + +def test_train_eval_with_various_outputs(): + class Net(torch.nn.Module): + def __init__(self, input_size, hidden_size, num_classes): + super(Net, self).__init__() + self.fc1 = torch.nn.Linear(input_size, hidden_size) + self.relu = torch.nn.ReLU() + + def forward(self, input1): + out1 = self.fc1(input1) + out2 = self.relu(out1) + # return different number of outputs for train ane eval mode + if self.training: + return out1, out2 + else: + return out2 + + def train_step(model, x): + out1, out2 = model(x) + loss = out2.sum() + loss.backward() + return out1, out2 + + device = torch.device('cuda') + N, D_in, H, D_out = 64, 784, 500, 10 + pt_model = Net(D_in, H, D_out).to(device) + ort_model = ORTModule(copy.deepcopy(pt_model)) + + # train mode + x = torch.randn(N, D_in, device=device) + pt_out1, pt_out2 = train_step(pt_model, x) + ort_out1, ort_out2 = train_step(ort_model, x) + + _test_helpers.assert_values_are_close(pt_out1, ort_out1) + _test_helpers.assert_values_are_close(pt_out2, ort_out2) + _test_helpers.assert_gradients_match_and_reset_gradient(ort_model, pt_model) + + # eval mode + pt_model.eval() + ort_model.eval() + + x = torch.randn(N, D_in, device=device) + pt_out = pt_model(x) + ort_out = ort_model(x) + _test_helpers.assert_values_are_close(pt_out, ort_out) def test_forward_dynamic_args(): device = 'cuda'