diff --git a/orttraining/orttraining/python/training/ortmodule.py b/orttraining/orttraining/python/training/ortmodule.py index a9ca1b1fc7..23c506af4f 100644 --- a/orttraining/orttraining/python/training/ortmodule.py +++ b/orttraining/orttraining/python/training/ortmodule.py @@ -97,6 +97,7 @@ class ORTModule(torch.nn.Module): self._export_again = False # TODO: This is incorrect when different layers may be in different devices self._device = next(module.parameters()).device + self._require_export = False # User module is wrapped to use its initializers and save computed gradients self._original_module = module diff --git a/orttraining/orttraining/test/python/orttraining_ortmodule_tests.py b/orttraining/orttraining/test/python/orttraining_ortmodule_tests.py index 9544da69e4..d33e8dc881 100644 --- a/orttraining/orttraining/test/python/orttraining_ortmodule_tests.py +++ b/orttraining/orttraining/test/python/orttraining_ortmodule_tests.py @@ -21,13 +21,21 @@ def parse_arguments(): return parser.parse_args() +def run_ortmodule_api_tests(cwd, log): + log.debug('Running: ORTModule-API tests') + + command = [sys.executable, '-m', 'pytest', '-sv', 'orttraining_test_ortmodule_api.py'] + + run_subprocess(command, cwd=cwd, log=log).check_returncode() + + def main(): args = parse_arguments() cwd = args.cwd log.info("Running ortmodule tests pipeline") - # TODO: add ortmodule tests + run_ortmodule_api_tests(cwd, log) return 0 diff --git a/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py b/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py new file mode 100644 index 0000000000..d235cae0fb --- /dev/null +++ b/orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py @@ -0,0 +1,241 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# orttraining_test_ortmodule_api.py + +import torch +import pytest + +import onnxruntime +from onnxruntime.training import ORTModule + +# PyTorch model definitions for tests + +class NeuralNetSinglePositionalArgument(torch.nn.Module): + def __init__(self, input_size, hidden_size, num_classes): + super(NeuralNetSinglePositionalArgument, 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) + + def forward(self, input1): + out = self.fc1(input1) + out = self.relu(out) + out = self.fc2(out) + return out + +class NeuralNetMultiplePositionalArguments(torch.nn.Module): + def __init__(self, input_size, hidden_size, num_classes): + super(NeuralNetMultiplePositionalArguments, 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) + + def forward(self, input1, input2): + model_input = input1 + input2 + out = self.fc1(model_input) + out = self.relu(out) + out = self.fc2(out) + return out + +class NeuralNetPositionalArguments(torch.nn.Module): + def __init__(self, input_size, hidden_size, num_classes): + super(NeuralNetPositionalArguments, 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) + + def forward(self, *model_inputs): + model_input = torch.sum(torch.stack(model_inputs), dim=0) + out = self.fc1(model_input) + out = self.relu(out) + out = self.fc2(out) + return out + +class NeuralNetKeywordArguments(torch.nn.Module): + def __init__(self, input_size, hidden_size, num_classes): + super(NeuralNetKeywordArguments, 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) + + def forward(self, x=None, y=None, z=None): + model_input = torch.sum(torch.stack([x, y, z]), dim=0) + out = self.fc1(model_input) + out = self.relu(out) + out = self.fc2(out) + return out + +class NeuralNetPositionalAndKeywordArguments(torch.nn.Module): + def __init__(self, input_size, hidden_size, num_classes): + super(NeuralNetPositionalAndKeywordArguments, 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) + + def forward(self, model_input, x=None, y=None, z=None): + model_input = model_input + torch.sum(torch.stack([x, y, z]), dim=0) + out = self.fc1(model_input) + out = self.relu(out) + out = self.fc2(out) + return out + +# ORTModule-API tests + +def test_forward_call_single_positional_argument(): + device = 'cuda' + + N, D_in, H, D_out = 64, 784, 500, 10 + model = NeuralNetSinglePositionalArgument(D_in, H, D_out).to(device) + model = ORTModule(model) + x = torch.randn(N, D_in, device=device) + try: + model(x) + except Exception as exception: + raise exception + +def test_forward_call_multiple_positional_arguments(): + device = 'cuda' + + N, D_in, H, D_out = 64, 784, 500, 10 + model = NeuralNetMultiplePositionalArguments(input_size=D_in, hidden_size=H, num_classes=D_out).to(device) + model = ORTModule(model) + x = torch.randn(N, D_in, device=device) + y = torch.randn(N, D_in, device=device) + try: + model(x, y) + except Exception as exception: + raise exception + +def test_forward_call_positional_arguments(): + device = 'cuda' + + N, D_in, H, D_out = 64, 784, 500, 10 + model = NeuralNetPositionalArguments(input_size=D_in, hidden_size=H, num_classes=D_out).to(device) + model = ORTModule(model) + args = [torch.randn(N, D_in, device=device), torch.randn(N, D_in, device=device), torch.randn(N, D_in, device=device)] + try: + model(*args) + except Exception as exception: + raise exception + +def test_forward_call_keyword_arguments(): + device = 'cuda' + + N, D_in, H, D_out = 64, 784, 500, 10 + model = NeuralNetKeywordArguments(D_in, H, D_out).to(device) + model = ORTModule(model) + x = torch.randn(N, D_in, device=device) + y = torch.randn(N, D_in, device=device) + z = torch.randn(N, D_in, device=device) + try: + model(x, y, z) + except Exception as exception: + raise exception + +def test_forward_call_positional_and_keyword_arguments(): + device = 'cuda' + + N, D_in, H, D_out = 64, 784, 500, 10 + model = NeuralNetPositionalAndKeywordArguments(D_in, H, D_out).to(device) + model = ORTModule(model) + a = torch.randn(N, D_in, device=device) + x = torch.randn(N, D_in, device=device) + y = torch.randn(N, D_in, device=device) + z = torch.randn(N, D_in, device=device) + try: + model(a, x, y, z) + except Exception as exception: + raise exception + +def test_model_cuda(): + original_device = 'cpu' + to_device = 'cuda' + + N, D_in, H, D_out = 64, 784, 500, 10 + model = NeuralNetSinglePositionalArgument(D_in, H, D_out) + model = ORTModule(model) + x = torch.randn(N, D_in, device=to_device) + for _, parameter_value in model.named_parameters(): + assert parameter_value.device.type == original_device + + model = model.cuda() + model(x) + + for _, parameter_value in model.named_parameters(): + assert parameter_value.device.type == to_device + +def test_model_cpu(): + original_device = 'cuda' + to_device = 'cpu' + + N, D_in, H, D_out = 64, 784, 500, 10 + model = NeuralNetSinglePositionalArgument(D_in, H, D_out).to(original_device) + model = ORTModule(model) + x = torch.randn(N, D_in) + for _, parameter_value in model.named_parameters(): + assert parameter_value.device.type == original_device + + model = model.cpu() + model(x) + + for _, parameter_value in model.named_parameters(): + assert parameter_value.device.type == to_device + +@pytest.mark.parametrize("original_device, to_argument, requires_export, device_type, device_index", [ + ('cpu', torch.device('cuda'), True, 'cuda', 0), + ('cpu', 'cuda', True, 'cuda', 0), + ('cpu', 'cuda:0', True, 'cuda', 0), + ('cpu', 'cuda', True, 'cuda', 0), + ('cuda', 'cuda', False, 'cuda', 0), + ('cuda', 'cuda:0', False, 'cuda', 0), + ('cuda', torch.device('cuda'), False, 'cuda', 0), + ('cuda', 'cpu', True, 'cpu', 0), + ('cuda', torch.device('cpu'), True, 'cpu', 0), + ('cpu', 'cpu', False, 'cpu', None), + ('cpu', torch.device('cpu'), False, 'cpu', None), + ('cpu', torch.zeros(2, device=torch.device('cuda')), True, 'cuda', 0), + ]) +def test_model_to_device(original_device, to_argument, requires_export, device_type, device_index): + N, D_in, H, D_out = 64, 784, 500, 10 + model = NeuralNetSinglePositionalArgument(D_in, H, D_out).to(original_device) + model = ORTModule(model) + x = torch.randn(N, D_in, device=device_type) + for _, parameter_value in model.named_parameters(): + assert parameter_value.device.type == original_device + + model = model.to(to_argument) + assert model._require_export == requires_export + assert model._device == torch.device(device_type+':'+str(device_index) if device_index is not None else device_type) + model(x) + + for _, parameter_value in model.named_parameters(): + assert parameter_value.device.type == device_type + +@pytest.mark.parametrize("original_device, to_device", [ + ('cuda', 'cpu'), + ('cpu', 'cuda') + ]) +def test_model_to_device_and_back_to_original(original_device, to_device): + N, D_in, H, D_out = 64, 784, 500, 10 + model = NeuralNetSinglePositionalArgument(D_in, H, D_out).to(original_device) + model = ORTModule(model) + for _, parameter_value in model.named_parameters(): + assert parameter_value.device.type == original_device + + model = model.to(to_device) + assert model._require_export == True + assert model._device == torch.device(to_device+':0') + + for _, parameter_value in model.named_parameters(): + assert parameter_value.device.type == to_device + + model = model.to(original_device) + assert model._require_export == True + assert model._device == torch.device(original_device+':0') + for _, parameter_value in model.named_parameters(): + assert parameter_value.device.type == original_device