From 19988370c231ebba66d2252d75865e8c03374161 Mon Sep 17 00:00:00 2001 From: Sergii Dymchenko Date: Thu, 17 Dec 2020 12:21:02 -0800 Subject: [PATCH 01/11] Add frontend test for unused input. --- .../python/orttraining_test_orttrainer_frontend.py | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py index c6c17482a5c08..fc23f6f735ce1 100644 --- a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py +++ b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py @@ -1428,3 +1428,17 @@ def testORTTrainerOptionsDisabledAdasumFlag(test_input): actual_values = orttrainer_options.ORTTrainerOptions(test_input) assert actual_values.distributed.enable_adasum == False + +def testORTTrainerUnusedInput(): + class UnusedInputModel(torch.nn.Module): + def __init__(self): + super(Net, self).__init__() + def forward(self, x, y): + return torch.mean(x) + + model = UnusedInputModel() + model_desc = {'inputs': [('x', [1]), ('y', [1])], 'outputs': [('loss', [], True)]} + optim_config = optim.LambConfig(lr=0.001) + trainer = orttrainer.ORTTrainer(model, model_desc, optim_config) + # Run just one step to make sure there are no iobinding errors for the unused input. + trainer.train_step(torch.FloatTensor([1.0]), torch.FloatTensor([1.0])) From dc61a5c2be829fd62d03b8b690fb825af70c520e Mon Sep 17 00:00:00 2001 From: Sergii Dymchenko Date: Thu, 17 Dec 2020 12:45:49 -0800 Subject: [PATCH 02/11] Fix frontend test for unused input. --- .../test/python/orttraining_test_orttrainer_frontend.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py index fc23f6f735ce1..0adcfdfcc93c9 100644 --- a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py +++ b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py @@ -1432,7 +1432,7 @@ def testORTTrainerOptionsDisabledAdasumFlag(test_input): def testORTTrainerUnusedInput(): class UnusedInputModel(torch.nn.Module): def __init__(self): - super(Net, self).__init__() + super(UnusedInputModel, self).__init__() def forward(self, x, y): return torch.mean(x) From 5fe2a891cf647621176d07af54f72e1f9e63f777 Mon Sep 17 00:00:00 2001 From: Sergii Dymchenko Date: Thu, 17 Dec 2020 12:51:14 -0800 Subject: [PATCH 03/11] Improve frontend test for unused input. --- .../test/python/orttraining_test_orttrainer_frontend.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py index 0adcfdfcc93c9..f802255f730cf 100644 --- a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py +++ b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py @@ -1440,5 +1440,9 @@ def forward(self, x, y): model_desc = {'inputs': [('x', [1]), ('y', [1])], 'outputs': [('loss', [], True)]} optim_config = optim.LambConfig(lr=0.001) trainer = orttrainer.ORTTrainer(model, model_desc, optim_config) + # Run just one step to make sure there are no iobinding errors for the unused input. - trainer.train_step(torch.FloatTensor([1.0]), torch.FloatTensor([1.0])) + try: + trainer.train_step(torch.FloatTensor([1.0]), torch.FloatTensor([1.0])) + except RuntimeError: + self.fail("RuntimeError doing train_step with unused input.") From 0b56a617d8ce99034a8620f59370fd2817705f68 Mon Sep 17 00:00:00 2001 From: Sergii Dymchenko Date: Thu, 17 Dec 2020 12:54:42 -0800 Subject: [PATCH 04/11] Fix frontend test for unused input. --- .../test/python/orttraining_test_orttrainer_frontend.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py index f802255f730cf..e2a9fabc76e78 100644 --- a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py +++ b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py @@ -1445,4 +1445,4 @@ def forward(self, x, y): try: trainer.train_step(torch.FloatTensor([1.0]), torch.FloatTensor([1.0])) except RuntimeError: - self.fail("RuntimeError doing train_step with unused input.") + pytest.fail("RuntimeError doing train_step with unused input.") From 017b0b464d12a8587dff38498b8415a892f937a4 Mon Sep 17 00:00:00 2001 From: Sergii Dymchenko Date: Thu, 17 Dec 2020 12:57:36 -0800 Subject: [PATCH 05/11] Add temp debug prints. --- orttraining/orttraining/python/training/orttrainer.py | 1 + .../test/python/orttraining_test_orttrainer_frontend.py | 2 +- 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/orttraining/orttraining/python/training/orttrainer.py b/orttraining/orttraining/python/training/orttrainer.py index 37c8b4ed51df1..2a4a3be3cc1ab 100644 --- a/orttraining/orttraining/python/training/orttrainer.py +++ b/orttraining/orttraining/python/training/orttrainer.py @@ -804,6 +804,7 @@ def _training_session_run_helper(self, is_train, inputs, inputs_desc, outputs_de else: iobinding = self._eval_io_binding + print(self._training_session.get_inputs(self)) # Bind input tensors for input, input_desc in zip(inputs, inputs_desc): device_index = _utils.get_device_index_from_input(input) diff --git a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py index e2a9fabc76e78..23ad32f5b62fa 100644 --- a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py +++ b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py @@ -1445,4 +1445,4 @@ def forward(self, x, y): try: trainer.train_step(torch.FloatTensor([1.0]), torch.FloatTensor([1.0])) except RuntimeError: - pytest.fail("RuntimeError doing train_step with unused input.") + pytest.fail("RuntimeError doing train_step with an unused input.") From c90ee11e53a129556d5c56e165efcfd70b4992bb Mon Sep 17 00:00:00 2001 From: Sergii Dymchenko Date: Thu, 17 Dec 2020 13:11:22 -0800 Subject: [PATCH 06/11] Add temp debug prints. --- orttraining/orttraining/python/training/orttrainer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/orttraining/orttraining/python/training/orttrainer.py b/orttraining/orttraining/python/training/orttrainer.py index 2a4a3be3cc1ab..bd7bc6c08618b 100644 --- a/orttraining/orttraining/python/training/orttrainer.py +++ b/orttraining/orttraining/python/training/orttrainer.py @@ -804,7 +804,7 @@ def _training_session_run_helper(self, is_train, inputs, inputs_desc, outputs_de else: iobinding = self._eval_io_binding - print(self._training_session.get_inputs(self)) + print("x"*10, ": ", self._training_session.get_inputs()) # Bind input tensors for input, input_desc in zip(inputs, inputs_desc): device_index = _utils.get_device_index_from_input(input) From 2a4aecd1203b3c5e34090f5c472929c64bf5fb5c Mon Sep 17 00:00:00 2001 From: Sergii Dymchenko Date: Thu, 17 Dec 2020 13:27:16 -0800 Subject: [PATCH 07/11] Add temp debug prints. --- orttraining/orttraining/python/training/orttrainer.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/orttraining/orttraining/python/training/orttrainer.py b/orttraining/orttraining/python/training/orttrainer.py index bd7bc6c08618b..3843610d1b2e5 100644 --- a/orttraining/orttraining/python/training/orttrainer.py +++ b/orttraining/orttraining/python/training/orttrainer.py @@ -804,7 +804,12 @@ def _training_session_run_helper(self, is_train, inputs, inputs_desc, outputs_de else: iobinding = self._eval_io_binding - print("x"*10, ": ", self._training_session.get_inputs()) + # Get the list of session input because unused inputs can be removed. + input_nodes = self._training_session.get_inputs()) + print("*"*10) + for input_node in input_nodes: + print(node.name) + # Bind input tensors for input, input_desc in zip(inputs, inputs_desc): device_index = _utils.get_device_index_from_input(input) From ae0f970b9b9599702939cee08790e4ffd4c61494 Mon Sep 17 00:00:00 2001 From: Sergii Dymchenko Date: Thu, 17 Dec 2020 13:42:47 -0800 Subject: [PATCH 08/11] Add temp debug prints. --- orttraining/orttraining/python/training/orttrainer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/orttraining/orttraining/python/training/orttrainer.py b/orttraining/orttraining/python/training/orttrainer.py index 3843610d1b2e5..3c393f99dfc33 100644 --- a/orttraining/orttraining/python/training/orttrainer.py +++ b/orttraining/orttraining/python/training/orttrainer.py @@ -805,7 +805,7 @@ def _training_session_run_helper(self, is_train, inputs, inputs_desc, outputs_de iobinding = self._eval_io_binding # Get the list of session input because unused inputs can be removed. - input_nodes = self._training_session.get_inputs()) + input_nodes = self._training_session.get_inputs() print("*"*10) for input_node in input_nodes: print(node.name) From 527db42d2a23c88607c2794c5b3d1949424a3a05 Mon Sep 17 00:00:00 2001 From: Sergii Dymchenko Date: Thu, 17 Dec 2020 13:48:31 -0800 Subject: [PATCH 09/11] Add temp debug prints. --- orttraining/orttraining/python/training/orttrainer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/orttraining/orttraining/python/training/orttrainer.py b/orttraining/orttraining/python/training/orttrainer.py index 3c393f99dfc33..386dc1a6c55ba 100644 --- a/orttraining/orttraining/python/training/orttrainer.py +++ b/orttraining/orttraining/python/training/orttrainer.py @@ -808,7 +808,7 @@ def _training_session_run_helper(self, is_train, inputs, inputs_desc, outputs_de input_nodes = self._training_session.get_inputs() print("*"*10) for input_node in input_nodes: - print(node.name) + print(input_node.name) # Bind input tensors for input, input_desc in zip(inputs, inputs_desc): From 9553e574d4960654e443830bef618550a0329e23 Mon Sep 17 00:00:00 2001 From: Sergii Dymchenko Date: Thu, 17 Dec 2020 13:54:18 -0800 Subject: [PATCH 10/11] Don't bind unused inputs in frontend. --- .../orttraining/python/training/orttrainer.py | 21 +++++++++---------- 1 file changed, 10 insertions(+), 11 deletions(-) diff --git a/orttraining/orttraining/python/training/orttrainer.py b/orttraining/orttraining/python/training/orttrainer.py index 386dc1a6c55ba..0d0224f784f7e 100644 --- a/orttraining/orttraining/python/training/orttrainer.py +++ b/orttraining/orttraining/python/training/orttrainer.py @@ -804,21 +804,20 @@ def _training_session_run_helper(self, is_train, inputs, inputs_desc, outputs_de else: iobinding = self._eval_io_binding - # Get the list of session input because unused inputs can be removed. + # Get the list of rhe actual session inputs because unused inputs can be removed. input_nodes = self._training_session.get_inputs() - print("*"*10) - for input_node in input_nodes: - print(input_node.name) + input_node_names = [input_node.name for input_node in input_nodes] # Bind input tensors for input, input_desc in zip(inputs, inputs_desc): - device_index = _utils.get_device_index_from_input(input) - iobinding.bind_input(input_desc.name, - input.device.type, - device_index, - _utils.dtype_torch_to_numpy(input.dtype), - list(input.size()), - input.data_ptr()) + if input_desc.name in input_node_names: + device_index = _utils.get_device_index_from_input(input) + iobinding.bind_input(input_desc.name, + input.device.type, + device_index, + _utils.dtype_torch_to_numpy(input.dtype), + list(input.size()), + input.data_ptr()) # Bind output tensors outputs_desc_resolved = self._resolve_symbolic_dimensions(inputs, inputs_desc, outputs_desc) From 0a02d43051f9fb6172ed923be134ab284fcbe821 Mon Sep 17 00:00:00 2001 From: Sergii Dymchenko Date: Thu, 17 Dec 2020 14:05:33 -0800 Subject: [PATCH 11/11] Fix typo. --- orttraining/orttraining/python/training/orttrainer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/orttraining/orttraining/python/training/orttrainer.py b/orttraining/orttraining/python/training/orttrainer.py index 0d0224f784f7e..7da55e965f702 100644 --- a/orttraining/orttraining/python/training/orttrainer.py +++ b/orttraining/orttraining/python/training/orttrainer.py @@ -804,7 +804,7 @@ def _training_session_run_helper(self, is_train, inputs, inputs_desc, outputs_de else: iobinding = self._eval_io_binding - # Get the list of rhe actual session inputs because unused inputs can be removed. + # Get the list of the actual session inputs because unused inputs can be removed. input_nodes = self._training_session.get_inputs() input_node_names = [input_node.name for input_node in input_nodes]