From 8b0ade0e834b7cc70d3784a6e562bc5d2d637469 Mon Sep 17 00:00:00 2001 From: Thiago Crepaldi Date: Tue, 27 Oct 2020 11:45:47 -0700 Subject: [PATCH] Integrate automatic graph split into ORTModule --- .../orttraining/python/training/ortmodule.py | 605 ++++++++++++++---- .../orttraining_test_ortmodule_basic.py | 18 +- ...ng_test_ortmodule_basic_transform_model.py | 148 ----- run_ortmodule_mvp.sh | 8 +- 4 files changed, 498 insertions(+), 281 deletions(-) delete mode 100644 orttraining/orttraining/test/python/orttraining_test_ortmodule_basic_transform_model.py diff --git a/orttraining/orttraining/python/training/ortmodule.py b/orttraining/orttraining/python/training/ortmodule.py index 493c413777..11e004c270 100644 --- a/orttraining/orttraining/python/training/ortmodule.py +++ b/orttraining/orttraining/python/training/ortmodule.py @@ -1,10 +1,12 @@ import copy import io +import logging import onnx import onnxruntime import os import torch import warnings +from onnxruntime.capi import _pybind_state as C from . import _utils @@ -20,6 +22,13 @@ class ORTModule(torch.nn.Module): # User module is wrapped to use its initializers and save computed gradients self._original_module = module + self._original_module_grad_output_len = -1 + self._original_module_forward_input_grads = [] + self._onnx_training = None + self._onnx_training_inputs_desc = [] + self._onnx_training_outputs_desc = [] + self._onnx_gradient = None + self._grad_builder_config = C.ModuleGradientGraphBuilderConfiguration() # Forward pass self._onnx_forward = None @@ -27,154 +36,316 @@ class ORTModule(torch.nn.Module): self._onnx_forward_initializers_desc = [] self._onnx_forward_inputs_desc = [] self._onnx_forward_outputs_desc = [] + self._onnx_forward_intermediate_outputs_desc = [] # Backward pass self._onnx_backward = None self._backward_session = None self._onnx_backward_initializers_desc = [] self._onnx_backward_inputs_desc = [] + self._onnx_backward_gradient_inputs_desc = [] self._onnx_backward_outputs_desc = [] - def forward(self, *input, **kwargs): + # Log level + self._loglevel = getattr(logging, 'WARNING') + + + def forward(self, *inputs, **kwargs): + '''Forward pass starts here and continues at `_ORTModuleFunction.forward` + + ONNX model is exported the first time this method is executed. + Next, a full training graph is splitted in forward and backward graph which are used + to instantiate ONNX Runtime InferenceSession`s + + TODO: #ImproveGraphSplitting + Additionally to that, several descriptor lists are generated to help identify + model input, output, initializer, intermediate and gradient tensors. + ''' if not self._onnx_forward: - original_forward_graph = ORTModule._get_forward_graph(self._original_module, *input, **kwargs) - gradient_graph = ORTModule._build_gradient_graph(original_forward_graph) - # TODO: Remove manual split after MVP - # self.forward_graph, self.backward_graph = ORTModule._split_forward_and_backward(gradient_graph) - self._onnx_forward = original_forward_graph # TODO: hard-coding for MVP - self._onnx_backward = gradient_graph # TODO: hard-coding for MVP + self._onnx_training = ORTModule._get_forward_graph(self._original_module, *inputs, **kwargs) + self._onnx_gradient = ORTModule._build_gradient_graph(self._onnx_training, self._grad_builder_config) + self._onnx_forward, self._onnx_backward = ORTModule._split_forward_and_backward(self._onnx_gradient, self._grad_builder_config.weight_names_to_train) self._forward_session = onnxruntime.InferenceSession(self._onnx_forward.SerializeToString()) self._backward_session = onnxruntime.InferenceSession(self._onnx_backward.SerializeToString()) # Forward I/O description + if not self._onnx_training_inputs_desc: + self._onnx_training_inputs_desc = self._get_input_from_graph(self._onnx_training) + logging.debug(f'Training inputs:\n\t {self._onnx_training_inputs_desc}') + if not self._onnx_training_outputs_desc: + self._onnx_training_outputs_desc = self._get_output_from_graph(self._onnx_training) + logging.debug(f'Training outputs:\n\t {self._onnx_training_outputs_desc}') if not self._onnx_forward_initializers_desc: self._onnx_forward_initializers_desc = self._get_initializer_from_graph(self._onnx_forward) - print(f'Forward initializers: {self._onnx_forward_initializers_desc}') + logging.debug(f'Forward initializers:\n\t {self._onnx_forward_initializers_desc}') if not self._onnx_forward_inputs_desc: self._onnx_forward_inputs_desc = self._get_input_from_graph(self._onnx_forward) - print(f'Forward inputs: {self._onnx_forward_inputs_desc}') + logging.debug(f'Forward inputs:\n\t {self._onnx_forward_inputs_desc}') if not self._onnx_forward_outputs_desc: self._onnx_forward_outputs_desc = self._get_output_from_graph(self._onnx_forward) - print(f'Forward outputs: {self._onnx_forward_outputs_desc}') + logging.debug(f'Forward outputs:\n\t {self._onnx_forward_outputs_desc}') + if not self._onnx_forward_intermediate_outputs_desc: + self._onnx_forward_intermediate_outputs_desc = self._get_intermediate_from_forward_graph(self._onnx_forward) + logging.debug(f'Forward intermediate outputs:\n\t {self._onnx_forward_intermediate_outputs_desc}') # Backward I/O description if not self._onnx_backward_initializers_desc: - self._onnx_backward_initializers_desc = self._get_initializer_from_graph(self._onnx_backward) - print(f'Backward initializers: {self._onnx_backward_initializers_desc}') + self._onnx_backward_initializers_desc = self._get_input_from_graph(self._onnx_backward, True) + logging.debug(f'Backward initializers: {self._onnx_backward_initializers_desc}') if not self._onnx_backward_inputs_desc: - self._onnx_backward_inputs_desc = self._get_input_from_graph(self._onnx_backward) - print(f'Backward inputs: {self._onnx_forward_inputs_desc}') + self._onnx_backward_inputs_desc = self._get_input_from_graph(self._onnx_backward, False, self._onnx_backward_initializers_desc) + logging.debug(f'Backward inputs: {self._onnx_backward_inputs_desc}') + if not self._onnx_backward_gradient_inputs_desc: + self._onnx_backward_gradient_inputs_desc = self._get_gradient_input_from_graph(self._onnx_backward, self._onnx_forward_inputs_desc, self._onnx_forward_initializers_desc, self._onnx_forward_intermediate_outputs_desc) + logging.debug(f'Backward gradient inputs: {self._onnx_backward_gradient_inputs_desc}') if not self._onnx_backward_outputs_desc: self._onnx_backward_outputs_desc = self._get_output_from_graph(self._onnx_backward) - print(f'Backward outputs: {self._onnx_backward_outputs_desc}') + logging.debug(f'Backward outputs: {self._onnx_backward_outputs_desc}') # Use a custom torch.autograd.Function to associate self.backward_graph as the # gradient implementation for self.forward_graph. class _ORTModuleFunction(torch.autograd.Function): @staticmethod - def forward(ctx, *input, **kwargs): - # TODO: Potential optimization is to detect which inputs and weights require gradients - input_with_initializer = self._prepare_forward_input_ort(*input, **kwargs) - outputs = self._run_forward_graph(input_with_initializer) - outputs = tuple(torch.from_numpy(out) for out in outputs) + def forward(ctx, *inputs, **kwargs): + '''Performs forward pass based on user input and PyTorch initializer - # TODO: Properly save dynamic number of intermediate tensors and remove them from model output - # Tensors that need to have gradients tracked can't be saved by `save_for_backward` - # saved_tensors ==> input1, fc2.weight, 7 - ctx.save_for_backward(*[input[0], input[3], outputs[1]]) - outputs = [outputs[0]] + TODO: **kwargs are not supported - # TODO: Properly support original module output format + Model outputs are returned to the user + The following tensors are stashed (in order) for backward pass + * (Partial) user input + * (Partial) Initializers + * Intermediate tensors + + TODO: #ImproveGraphSplitting + String matching to separate user input from initializer + ''' + + # Convert input to dict of torch tensors + data_dict = self._convert_forward_input_list_to_dict(*inputs) + + # Convert dict of torch tensors to dict of numpy arrays (ORT BE requirement) + data_dict_numpy = self._convert_dict_torch_to_numpy(data_dict) + + # Feed forward + outputs, intermediate = self._run_forward_graph(data_dict_numpy) + outputs = tuple(torch.from_numpy(item) for item in outputs) + + # Save input, initializers and intermediate tensors to be used during backward + initializer_names = [item['name'] for item in self._onnx_backward_initializers_desc] + input_names = [item['name'] for item in self._onnx_backward_inputs_desc if item['name'] not in initializer_names] + ctx_input = tuple(v for k,v in data_dict.items() if k in input_names) + ctx_initializer = tuple(v for k,v in data_dict.items() if k in initializer_names) + intermediate = tuple(torch.from_numpy(item) for item in intermediate) + ctx.save_for_backward(*[*ctx_input, *ctx_initializer, *intermediate]) + + # TODO: Support original module output (currently dict is not supported) if len(outputs) == 1: return outputs[0] return outputs @staticmethod def backward(ctx, *grad_output): - # TODO: Properly restore dynamic number of intermediate tensors - # saved_tensors ==> input1, fc2.weight, 7 + '''Performs backward pass based on grad wrt output and internal state + + Internal state is composed of: + * Tensor stashed (in a particular order) during forward: + * (partial) user input, (partial) initializers and intermediate tensors + + TODO: #ImproveGraphSplitting + Length of `*grad_output` is needed to detect intermediate tensors during backward pass + + TODO: Input gradient is hard-coded to torch.tensor([1.]) + ''' saved_tensors = ctx.saved_tensors + # Used to create backward input + if self._original_module_grad_output_len == -1: + self._original_module_grad_output_len = len(grad_output) + grad_weights = self._run_backward_graph(*[*saved_tensors, *grad_output]) - grad_weights = [torch.from_numpy(grad) for grad in grad_weights] - # TODO: backward must return grad tensors in the same order forward does - # [input1_grad, fc1.weight_grad, fc1.bias_grad, fc2.weight_grad, fc2.bias_grad] - return tuple([torch.tensor([1.]), grad_weights[1], grad_weights[0], grad_weights[2], grad_weights[3]]) - return _ORTModuleFunction.apply(*self._prepare_forward_input_autograd(*input, **kwargs)) + result = [torch.tensor([1])]* len(self._onnx_training_inputs_desc) + result += [torch.from_numpy(grad) for grad in grad_weights] + return tuple(result) + + return _ORTModuleFunction.apply(*self._convert_forward_input_to_list(*inputs, **kwargs)) + + def _convert_forward_input_to_list(self, *inputs, **kwargs): + '''Creates forward `*inputs` list from user input and PyTorch initializers + + TODO: **kwargs is not supported + + ONNX Runtime forward requires an order list of: + * User input: computed from forward InferenceSession + * Initializers: computed from original PyTorch model parameters + ''' - def _prepare_forward_input_autograd(self, *input, **kwargs): # List containing both user inputs and initializers, in this order - input_with_initializer = [] + result = [] # Inputs for idx, input_data in enumerate(self._forward_session.get_inputs()): - input_with_initializer.append(input[idx]) + result.append(inputs[idx]) # Initializers for idx, param in enumerate(self._original_module.named_parameters()): - input_with_initializer.append(param[1]) + result.append(param[1]) - # TODO: [input1, fc1.weight, fc1.bias, fc2.weight, fc2.bias] - return input_with_initializer + return result + + def _convert_dict_torch_to_numpy(self, tensor_dict): + '''Convert `tensor_dict` PyTorch tensors to numpy tensors + + This is a ONNX Runtime requirement + + TODO: #UseIOBinding + ''' + result = {} + for k,v in tensor_dict.items(): + result.update({k : v.detach().cpu().numpy()}) + return result + + def _convert_forward_input_list_to_dict(self, *inputs): + '''Convert forward `*inputs` list to dict + + TODO: #ImproveGraphSplitting + Additionally, a list of gradient names of initializers are created to be used by backprop + + TODO: Input gradient is being ignored for MVP + ''' - def _prepare_forward_input_ort(self, *inputs): # Dictionary containing both inputs and initializers - input_with_initializer = {} + result = {} - # TODO: [input1, fc1.weight, fc1.bias, fc2.weight, fc2.bias] # Inputs - inputs_len = 0 + result_len = 0 for idx, input_data in enumerate(self._forward_session.get_inputs()): - inputs_len += 1 - input_with_initializer.update({input_data.name: inputs[idx].cpu().numpy()}) + result_len += 1 + result.update({input_data.name: inputs[idx]}) # Initializers for param in self._original_module.named_parameters(): - input_with_initializer.update({param[0]: inputs[inputs_len].detach().numpy()}) - inputs_len += 1 + result.update({param[0]: inputs[result_len]}) + # TODO: Create order list of input grads to use during backward. + # (for scenarios where gradients of input is required - not covered on MVP) + # if len(self._original_module_forward_input_grads) < len(self._onnx_training_inputs_desc): + # self._original_module_forward_input_grads.append(param[0]+'_grad') - return input_with_initializer + # TODO: Create order list of initializer grads to use during backward. + # if len(self._original_module_forward_input_grads) < len(self._onnx_backward_outputs_desc) + len(self._onnx_training_inputs_desc): + if len(self._original_module_forward_input_grads) < len(self._onnx_backward_outputs_desc): + self._original_module_forward_input_grads.append(param[0]+'_grad') + result_len += 1 - def _prepare_backward_input(self, *inputs, **kwargs): - # Dictionary containing initializers - input_with_initializer = {} + return result - # User input - # TODO: How to determine which user input to feed to backward - # for idx, input_data in enumerate(self._forward_session.get_inputs()): - # input_with_initializer.update({input_data.name: inputs[idx].cpu().numpy()}) - input_with_initializer.update({'input1' : inputs[0].detach().numpy()}) + def _convert_backward_input_list_to_dict(self, *inputs): + '''Convert backward `*inputs` list to dict + + ONNX Runtime backend requires dict as input, which is composed of: + * User input + Although not necessary, all user inputs are used for simplicity + * (Partial) Initializers + init_begin = len(user_input) + init_count = len(Pre-computed list of initializer) + * Intermediate tensors TODO: #ImproveGraphSplitting + Intermediate tensors are inferred from input position: + interm_begin = len(user_input) + len(initializer) + interm_count = len(all_inputs) - len(user_input) - len(initializer) - len(grad_output) + * Gradient wrt outputs TODO: #ImproveGraphSplitting + Gradient tensors are inferred from input position: + grads_begin = len(user_input) + len(initializer) + len(intermediate) + grads_count = len(all_inputs) - len(user_input) - len(initializer) - len(intermediate) + ''' + + # Dictionary containing both inputs and initializers + result = {} + + # Inputs + result_len = 0 + for idx, input_data in enumerate(self._forward_session.get_inputs()): + result.update({ input_data.name : inputs[idx]}) + result_len += 1 # Initializers - # TODO: How to determine which initializer (subset) to be used - # for idx, param in enumerate(self._original_module.named_parameters()): - # input_with_initializer.update({param[0]: param[1].detach().numpy()}) - input_with_initializer.update({'fc2.weight' : inputs[1].detach().numpy()}) + for initializer in self._onnx_backward_initializers_desc: + result.update({initializer['name']: inputs[result_len]}) + result_len += 1 - # Intermediates - # TODO: How to determine intermediates name? - input_with_initializer.update({'7': inputs[2].detach().numpy()}) + # Intermediate + intermediate_len = len(inputs) - result_len - self._original_module_grad_output_len + for idx in range(intermediate_len): + result.update({self._onnx_forward_intermediate_outputs_desc[idx]['name']: inputs[result_len]}) + result_len += 1 - # Grad output - # TODO: How to determine grad_output name? - input_with_initializer.update({'probability_grad': inputs[3].detach().numpy()}) - return input_with_initializer + # Grad outputs + for idx in range(len(inputs)-result_len): + result.update({self._onnx_backward_gradient_inputs_desc[idx]['name']: inputs[result_len]}) + result_len += 1 - def _run_forward_graph(self, data_with_initializer): # input, weights): - return self._forward_session.run(None, data_with_initializer) + return result + + def _run_forward_graph(self, inputs): + '''Execute forward pass on ONNX Runtime + + Output order has to be specified to ONNX Runtime backend + to distinguish intermediate from output tensors + ''' + + output_names = [out['name'] for out in self._onnx_forward_outputs_desc] + forward_output = self._forward_session.run(output_names, inputs) + output = forward_output[:len(self._onnx_training_outputs_desc)] + intermediates = forward_output[len(self._onnx_training_outputs_desc):] + return output, intermediates def _run_backward_graph(self, *inputs, **kwargs): - data = self._prepare_backward_input(*inputs, **kwargs) - # TODO: Hack to guarantee output order from InferenceSession.run() - return self._backward_session.run(['fc1.bias_grad', 'fc1.weight_grad', 'fc2.weight_grad', 'fc2.bias_grad'], data) + '''Execute backward pass on ONNX Runtime + + `*inputs` is converted from list to a list of detached numpy tensors before + being fed to an ONNX Runtime InferenceSession + + TODO: **kwargs are not supported + ''' + + # Convert input to dict of torch tensors + data = self._convert_backward_input_list_to_dict(*inputs) + + # Convert dict of torch tensors to dict of numpy arrays (ORT BE requirement) + data = self._convert_dict_torch_to_numpy(data) + return self._backward_session.run(self._original_module_forward_input_grads, data) @staticmethod - def _get_forward_graph(module, module_input): - # TODO: Pytorch module must be exported to ONNX and splitted - # Hard-coding with MNIST stub for MVP - return onnx.load('./model_with_training_forward_sliced.onnx') + def _get_forward_graph(module, *inputs, **kwargs): + '''Exports PyTorch `module` to ONNX with training flag, using `*inputs` as input + + TODO: Support contrib OPs support? user model has no hint + TODO: How to support dynamic axes? Dimensions are determined by samples + TODO: How to ingest **kwargs in proper order during export? + ''' + # Export the model to memory + f = io.BytesIO() + + # Deepcopy inputs, since input values may change after model run. + sample_inputs_copy = copy.deepcopy(inputs) + + # Export torch.nn.Module to ONNX + torch.onnx.export(module, + tuple(sample_inputs_copy), + f, + opset_version=ONNX_OPSET_VERSION, + do_constant_folding=False, + training=torch.onnx.TrainingMode.TRAINING) + return onnx.load_model_from_string(f.getvalue()) def _get_initializer_from_graph(self, graph): + '''Returns a descriptor list of initializers for `graph` + + The list descriptor has the following format: + [{ 'name': name, 'shape':[int1,...,intN], 'dtype': ]}] + + For ONNX types, refer to https://github.com/onnx/onnx/blob/master/onnx/onnx.in.proto#L461 + ''' + # TODO: There is a tradeoff between memory footprint and total model export time # Ideally we want to export the model using torch.onnx.export(.., export_params=False, keep_initializers_as_inputs=True) # to obtain an ONNX model with minimal size and initializers as input. @@ -190,27 +361,92 @@ class ORTModule(torch.nn.Module): initializers = [] for initializer in graph.graph.initializer: name = initializer.name - # TODO: Dynamic shape is not being handled yet shape = initializer.dims dtype = _utils.dtype_onnx_to_torch(initializer.data_type) initializers.append({'name': name, 'shape': shape, 'dtype': dtype}) return initializers - def _get_input_from_graph(self, graph): + def _get_input_from_graph(self, graph, initializers_only=False, append_initializers=[]): + '''Returns a descriptor list of input tensors for an ONNX `graph` + + When `initializers_only=True`, only input initializers are returned. Otherwise, both + user input and initializers are considered. + This is being used to get backward initializer list TODO: #ImproveGraphSplitting + + When `append_initializers` is not empty, this list is appended to the end of the result list + This is being used to get backward input list TODO: #ImproveGraphSplitting + + The list descriptor has the following format: + [{ 'name': name, 'shape':[int1,...,intN], 'dtype': ]}] + + For ONNX types, refer to https://github.com/onnx/onnx/blob/master/onnx/onnx.in.proto#L461 + ''' + inputs = [] for elem in graph.graph.input: for initializer in self._onnx_forward_initializers_desc: if elem.name == initializer['name']: + if initializers_only: + name = elem.name + shape = [dim.dim_value for dim in elem.type.tensor_type.shape.dim] + dtype = _utils.dtype_onnx_to_torch(elem.type.tensor_type.elem_type) + inputs.append({'name': name, 'shape': shape, 'dtype': dtype}) break else: - name = elem.name - # TODO: Dynamic shape is not being handled yet - shape = [dim.dim_value for dim in elem.type.tensor_type.shape.dim] - dtype = _utils.dtype_onnx_to_torch(elem.type.tensor_type.elem_type) - inputs.append({'name': name, 'shape': shape, 'dtype': dtype}) + if not initializers_only: + name = elem.name + shape = [dim.dim_value for dim in elem.type.tensor_type.shape.dim] + dtype = _utils.dtype_onnx_to_torch(elem.type.tensor_type.elem_type) + inputs.append({'name': name, 'shape': shape, 'dtype': dtype}) + if append_initializers: + inputs.extend(append_initializers) return inputs + def _get_gradient_input_from_graph(self, backward_graph, forward_input, forward_initializer, forward_intermediate): + '''Returns a descriptor list of gradient output for `backward_graph` + + Gradient output tensors are found through an elimination process, that cross reference + inputs from the backward graph to the forward input, initializer and intermediate tensors. + + The list descriptor has the following format: + [{ 'name': name, 'shape':[int1,...,intN], 'dtype': ]}] + + For ONNX types, refer to https://github.com/onnx/onnx/blob/master/onnx/onnx.in.proto#L461 + + TODO: #ImproveGraphSplitting + ''' + grads = [] + found = False + for elem in backward_graph.graph.input: + for item in forward_input: + if elem.name == item['name']: + # skip output + break + else: + for item in forward_initializer: + if elem.name == item['name']: + # skip output + break + else: + for item in forward_intermediate: + if elem.name == item['name']: + # skip output + break + else: + name = elem.name + shape = [dim.dim_value for dim in elem.type.tensor_type.shape.dim] + dtype = _utils.dtype_onnx_to_torch(elem.type.tensor_type.elem_type) + grads.append({'name': name, 'shape': shape, 'dtype': dtype}) + return grads + def _get_output_from_graph(self, graph): + '''Returns a descriptor list of output tensors for an ONNX `graph` + + The list descriptor has the following format: + [{ 'name': name, 'shape':[int1,...,intN], 'dtype': ]}] + + For ONNX types, refer to https://github.com/onnx/onnx/blob/master/onnx/onnx.in.proto#L461 + ''' outputs = [] for elem in graph.graph.output: for initializer in self._onnx_forward_initializers_desc: @@ -219,54 +455,179 @@ class ORTModule(torch.nn.Module): break else: name = elem.name - # TODO: Dynamic shape is not being handled yet shape = [dim.dim_value for dim in elem.type.tensor_type.shape.dim] dtype = _utils.dtype_onnx_to_torch(elem.type.tensor_type.elem_type) outputs.append({'name': name, 'shape': shape, 'dtype': dtype}) return outputs - @staticmethod - def _save_onnx_graph(onnx_graph, path): - r"""Persists ONNX model into :py:attr:`path` + def _get_intermediate_from_forward_graph(self, forward_graph): + '''Returns a descriptor list with all intermediate tensors for `forward_graph` - The model will be saved as a Google Protocol Buffers (aka protobuf) file as per ONNX standard. - The graph includes full information, including inference and training metadata. + Intermediate tensors are found through an elimination process, that cross reference + outputs from the forward graph to the original model (exported to ONNX) - Args: - onnx_graph (onnx.ModelProto): Either forward or backward graph - path (str): Full path, including filename, to save the ONNX model in the filesystem + The list descriptor has the following format: + [{ 'name': name, 'shape':[int1,...,intN], 'dtype': ]}] - Raises: - ValueError: raised when `path` is not valid path - """ - - assert isinstance(path, str), "'path' must be a valid path string" - dir_name = os.path.dirname(path) - file_name = os.path.basename(path) - if (dir_name and not os.path.exists(dir_name)) or not file_name: - warnings.warn("'path' is not valid or does not exist") - return - - with open(path, "wb") as f: - f.write(onnx_graph.SerializeToString()) + TODO: #ImproveGraphSplitting + ''' + intermediates = [] + for elem in forward_graph.graph.output: + for output in self._onnx_training_outputs_desc: + if elem.name == output['name']: + # skip output + break + else: + name = elem.name + shape = [dim.dim_value for dim in elem.type.tensor_type.shape.dim] + dtype = _utils.dtype_onnx_to_torch(elem.type.tensor_type.elem_type) + intermediates.append({'name': name, 'shape': shape, 'dtype': dtype}) + return intermediates @staticmethod - def _build_gradient_graph(forward_graph): - # TODO: Invoke the C++ GradientBuilder implementation via pybind. - # Return an ONNX graph that contains the forward and backward nodes, which takes the - # following inputs: - # * Module inputs - # * Module weights - # * Gradients with respect to the module outputs - # …and produces gradients with respect to the module inputs and weights. - return onnx.load('./model_with_training_backward_sliced.onnx') + def _build_gradient_graph(forward_graph, config): + '''Adds gradient nodes on top of an existing ONNX graph (with training flag) + + TODO: #SplittingGraphAtFrontend + ''' + if not config.weight_names_to_train: + weight_names_to_train = set() + for initializer in forward_graph.graph.initializer: + weight_names_to_train.add(initializer.name) + config.weight_names_to_train = weight_names_to_train + output_names = set() + for output in forward_graph.graph.output: + output_names.add(output.name) + config.output_names = output_names + return onnx.load_model_from_string(C.ModuleGradientGraphBuilder().build(forward_graph.SerializeToString(), config)) @staticmethod - def _split_forward_and_backward(gradient_graph): - # TODO: Split the result of _build_gradient_graph into two subgraphs: - # * A forward graph that takes module inputs and weights as input, and produces module - # outputs and (“stashed”) intermediate tensors as output. - # * A backward graph that takes intermediate tensors, module weights, and gradients - # respect to the module outputs as inputs, and produces gradients with respect to the - # module inputs and weights. - return (None, None) + def _split_forward_and_backward(onnx_model, weight_names_to_train): + '''Splits the result of _build_gradient_graph into two subgraphs + + * A forward graph that takes module inputs and weights as input, and produces module + outputs and (“stashed”) intermediate tensors as output. + * A backward graph that takes input, intermediate tensors, module weights, and gradients + with respect to the module outputs as inputs, and produces gradients with respect to the + module inputs and weights. + + TODO: #SplittingGraphAtFrontend + ''' + + def remove_nodes(onnx_model, nodes_to_remove): + all_nodes = [] + for node in onnx_model.graph.node: + if node not in nodes_to_remove: + all_nodes.append(node) + + onnx_model.graph.ClearField('node') + onnx_model.graph.node.extend(all_nodes) + + def add_output(model, name, data_type = None, docstring = None): + new_output = model.graph.value_info.add() + new_output.name = name + if data_type: + new_output.type.CopyFrom(data_type) + if docstring: + new_output.doc_string = docstring + model.graph.output.append(new_output) + + def add_input_from_initializer(model, initializer, docstring=None): + new_input = onnx.helper.make_tensor_value_info(initializer.name, initializer.data_type, initializer.dims, docstring) + model.graph.input.append(new_input) + + def add_input(model, name, data_type = None, dims = None, docstring = None): + new_input = onnx.helper.make_tensor_value_info(name, data_type, dims, docstring) + model.graph.input.append(new_input) + + forward_graph_outputs = set() + backward_graph_inputs = set() + backward_graph_outputs = set() + + # Get forward graph + forward_model = copy.deepcopy(onnx_model) + nodes_to_remove_from_forward_graph = [] + initializers = {} + for initializer in forward_model.graph.initializer: + initializers[initializer.name] = initializer + forward_graph_initializer_names = set() + for node in forward_model.graph.node: + if node.doc_string == 'Backward pass': + # nodes belongs to backward graph + nodes_to_remove_from_forward_graph.append(node) + for input in node.input: + backward_graph_inputs.add(input) + for output in node.output: + backward_graph_outputs.add(output) + else: + # nodes belogs to forward graph + for input in node.input: + if input in initializers: + forward_graph_initializer_names.add(input) + for output in node.output: + forward_graph_outputs.add(output) + + forward_model.graph.ClearField('initializer') + for initializer_name in forward_graph_initializer_names: + forward_model.graph.initializer.append(initializers[initializer_name]) + # training weights need to be added to input + if initializer_name in weight_names_to_train: + add_input_from_initializer(forward_model, initializers[initializer_name]) + + # outputs from forward graph that are also inputs of backwoard graph need to be added as graph output. + for output in forward_graph_outputs: + if output in backward_graph_inputs: + add_output(forward_model, output) + + remove_nodes(forward_model, nodes_to_remove_from_forward_graph) + + # Get backward graph + tensor_elem_types = {} + infered_model = onnx.shape_inference.infer_shapes(onnx_model) + for value_info in infered_model.graph.value_info: + tensor_elem_types[value_info.name] = value_info.type.tensor_type.elem_type + + backward_model = copy.deepcopy(onnx_model) + initializers = {} + for initializer in backward_model.graph.initializer: + initializers[initializer.name] = initializer + + nodes_to_remove_from_backward_graph = [] + for node in backward_model.graph.node: + if node.doc_string != 'Backward pass': + nodes_to_remove_from_backward_graph.append(node) + + backward_graph_initializer_names = set() + for input in backward_graph_inputs: + if input in forward_graph_outputs: + # inputs of backward graph that are also outputs from forward graph need to be added to backward graph input + add_input(backward_model, input, tensor_elem_types[input] if input in tensor_elem_types else 1) + elif input in forward_graph_initializer_names: + # inputs from forward graph initializers need to be added to backward graph input + add_input_from_initializer(backward_model, initializers[input]) + elif input in initializers: + backward_graph_initializer_names.add(input) + + # gradient of forward graph output will be the input of backward graph + for output in backward_model.graph.output: + if output.name + '_grad' in backward_graph_inputs: + add_input(backward_model, output.name + '_grad', output.type.tensor_type.elem_type) + + backward_model.graph.ClearField('initializer') + for initializer_name in backward_graph_initializer_names: + backward_model.graph.initializer.append(initializers[initializer_name]) + + # add gradient output to backward graph output + # TODO: need to add gradient of graph input to backward graph output + new_backward_graph_outputs = set() + for output in backward_graph_outputs: + if output.endswith('_grad') and output[:-5] in forward_graph_initializer_names: + new_backward_graph_outputs.add(output) + + backward_model.graph.ClearField('output') + for output in new_backward_graph_outputs: + add_output(backward_model, output) + + remove_nodes(backward_model, nodes_to_remove_from_backward_graph) + + return forward_model, backward_model diff --git a/orttraining/orttraining/test/python/orttraining_test_ortmodule_basic.py b/orttraining/orttraining/test/python/orttraining_test_ortmodule_basic.py index 6af177b80a..8a93aa862a 100644 --- a/orttraining/orttraining/test/python/orttraining_test_ortmodule_basic.py +++ b/orttraining/orttraining/test/python/orttraining_test_ortmodule_basic.py @@ -1,4 +1,5 @@ import argparse +import logging import torch from torchvision import datasets, transforms @@ -81,12 +82,12 @@ def main(): parser = argparse.ArgumentParser(description='PyTorch MNIST Example') parser.add_argument('--train-steps', type=int, default=-1, metavar='N', help='number of steps to train. Set -1 to run through whole dataset (default: -1)') - parser.add_argument('--lr', type=float, default=0.001, metavar='LR', - help='learning rate (default: 0.001)') + parser.add_argument('--lr', type=float, default=0.01, metavar='LR', + help='learning rate (default: 0.01)') parser.add_argument('--batch-size', type=int, default=20, metavar='N', help='input batch size for training (default: 20)') - parser.add_argument('--test-batch-size', type=int, default=1000, metavar='N', - help='input batch size for testing (default: 1000)') + parser.add_argument('--test-batch-size', type=int, default=20, metavar='N', + help='input batch size for testing (default: 20)') parser.add_argument('--no-cuda', action='store_true', default=False, help='disables CUDA training') parser.add_argument('--seed', type=int, default=42, metavar='S', @@ -99,6 +100,9 @@ def main(): help='views forward and backward graphs') parser.add_argument('--epochs', type=int, default=10, metavar='N', help='number of epochs to train (default: 10)') + parser.add_argument('--log-level', choices=['DEBUG', 'INFO', 'WARNING', 'ERROR', 'CRITICAL'], default='WARNING', + help='Log level (default: WARNING)') + args = parser.parse_args() @@ -131,6 +135,12 @@ def main(): if not args.pytorch_only: print('Training MNIST on ORTModule....') model = ORTModule(model) + + # Set log level + numeric_level = getattr(logging, args.log_level.upper(), None) + if not isinstance(numeric_level, int): + raise ValueError('Invalid log level: %s' % loglevel) + logging.basicConfig(level=numeric_level) else: print('Training MNIST on vanilla PyTorch....') optimizer = torch.optim.SGD(model.parameters(), lr=args.lr) diff --git a/orttraining/orttraining/test/python/orttraining_test_ortmodule_basic_transform_model.py b/orttraining/orttraining/test/python/orttraining_test_ortmodule_basic_transform_model.py deleted file mode 100644 index 0965b5ad06..0000000000 --- a/orttraining/orttraining/test/python/orttraining_test_ortmodule_basic_transform_model.py +++ /dev/null @@ -1,148 +0,0 @@ -# coding=utf8 -import copy -import sys -import onnx -from onnx import helper, shape_inference -from onnx import TensorProto -import numpy as np -from onnx import numpy_helper - -if len(sys.argv) < 2: - print("Please give model path...") - exit(1) - -input_model_name = sys.argv[1] -output_forward_model_name = input_model_name[:-5] + '_forward_sliced.onnx' -output_backward_model_name = input_model_name[:-5] + '_backward_sliced.onnx' - -def add_input_from_initializer(model, initializer, docstring=None): - new_input = onnx.helper.make_tensor_value_info(initializer.name, initializer.data_type, initializer.dims, docstring) - model.graph.input.append(new_input) - -def add_input(model, name, data_type = None, dims = None, docstring = None): - new_input = onnx.helper.make_tensor_value_info(name, data_type, dims, docstring) - model.graph.input.append(new_input) - -def add_output(model, name, data_type = None, docstring = None): - new_output = model.graph.value_info.add() - new_output.name = name - if data_type: - new_output.type.CopyFrom(data_type) - if docstring: - new_output.doc_string = docstring - model.graph.output.append(new_output) - -def find_model_input(model, input_name): - for input in model.graph.input: - if input.name == input_name: - return input - return None - -def find_model_output(model, output_name): - for output in model.graph.output: - if output.name == output_name: - return output - return None - -def find_initializer(model, name): - for initializer in model.graph.initializer: - if initializer.name == name: - return initializer - return None - -def find_node(model, name): - for node in model.graph.node: - if node.name == name: - return node - return None - -############################################################################### -# FORWARD PASS GRAPH ########################################################## -############################################################################### -model = onnx.load(input_model_name) - -# Remove model inputs -nodes = ['label'] -for node in nodes: - node = find_model_input(model, node) - model.graph.input.remove(node) - -# Remove model outputs -nodes = ['loss', 'fc1.bias_grad', 'fc1.weight_grad', 'fc2.bias_grad', 'fc2.weight_grad'] -for node in nodes: - node = find_model_output(model, node) - model.graph.output.remove(node) - -# Add input with same name, type and shape as the initializers -# They are: [fc1.bias, fc1.weight, fc2.bias, fc2.weight] -forward_initializer_names = ['fc1.bias', 'fc1.weight', 'fc2.bias', 'fc2.weight'] -forward_initializer = {} -for node in forward_initializer_names: - node = find_initializer(model, node) - add_input_from_initializer(model, node, f'thiagofc: add {node.name} as model input') - forward_initializer.update({node.name : copy.deepcopy(node)}) - -# Remove initializers from model -# TODO: Do this when we are able to distinguish inputs from initializers -# for node in forward_initializer_names: -# node = find_initializer(model, init) -# model.graph.initializer.remove(node) - -# Remove backward-related initializers -nodes = ['loss_grad', 'ZeroConstant'] -for node in nodes: - node = find_initializer(model, node) - model.graph.initializer.remove(node) - -# Remove OPs -nodes = ['SoftmaxCrossEntropyLoss_3', 'SoftmaxCrossEntropyLoss_3_Grad/SoftmaxCrossEntropyLossGrad_0', - 'Gemm_2_Grad/ReduceSum_3', 'Gemm_2_Grad/Identity_4', 'Gemm_2_Grad/Gemm_2', 'Gemm_2_Grad/Gemm_1', - 'Relu_1_Grad/ReluGrad_0', 'Gemm_0_Grad/Gemm_1', 'Gemm_0_Grad/ReduceSum_2', 'Gemm_0_Grad/Identity_3'] -for node in nodes: - node = find_node(model, node) - model.graph.node.remove(node) - -# Add new outputs: -# They are: 7 -add_output(model, '7', None, 'thiagofc: add 7 as model output') - -with open(output_forward_model_name, "wb") as f: - f.write(model.SerializeToString()) - - -############################################################################### -# BACKWARD PASS GRAPH ########################################################## -############################################################################### -model = onnx.load(input_model_name) - -# Add new inputs: -# TODO: Should we specify types here? ORT graph doesn't have that info available, but ONNX API needs it -add_input_from_initializer(model, forward_initializer['fc2.weight'], 'thiagofc: add fc2.weight as model input') -add_input(model, '7', 1, None, 'thiagofc: add 7 as model input') -add_input(model, 'probability_grad', 1, None, 'thiagofc: add probability_grad as model input') - -# Remove model inputs -# They are: label -node = find_model_input(model, 'label') -model.graph.input.remove(node) - -# Remove model outputs -nodes = ['loss', 'probability'] -for node in nodes: - node = find_model_output(model, node) - model.graph.output.remove(node) - -# Remove OP nodes from forward pass -nodes = ['Gemm_0', 'Relu_1', 'Gemm_2', 'SoftmaxCrossEntropyLoss_3', 'SoftmaxCrossEntropyLoss_3_Grad/SoftmaxCrossEntropyLossGrad_0'] -for node in nodes: - node = find_node(model, node) - model.graph.node.remove(node) - -# Remove initializers -forward_initializer_names.extend(['loss_grad']) -for node in forward_initializer_names: - node = find_initializer(model, node) - model.graph.initializer.remove(node) - -with open(output_backward_model_name, "wb") as f: - f.write(model.SerializeToString()) diff --git a/run_ortmodule_mvp.sh b/run_ortmodule_mvp.sh index 5bb742eeee..8a14061d41 100755 --- a/run_ortmodule_mvp.sh +++ b/run_ortmodule_mvp.sh @@ -14,11 +14,5 @@ export PYTHONPATH=$(pwd) echo "Copying PyTorch frontend source-code to build folder" cp -Rf ../../../orttraining/orttraining/python/training/* ../../../build/Linux/RelWithDebInfo/onnxruntime/training/ -echo "Running MNIST through Optimized API (ORTTrainer) to get full training graph" -python ../../../samples/python/mnist/ort_mnist.py --train-steps 1 --test-batch-size 0 --save-path "." - -echo "Splitting full training graph in forward and backward graphs" -python ../../../orttraining/orttraining/test/python/orttraining_test_ortmodule_basic_transform_model.py model_with_training.onnx - echo "Running Flexible API (ORTModule)" -python ../../../orttraining/orttraining/test/python/orttraining_test_ortmodule_basic.py \ No newline at end of file +python ../../../orttraining/orttraining/test/python/orttraining_test_ortmodule_basic.py