Copy forward signature from PyTorch model. (#6777)

This commit is contained in:
Sergii Dymchenko 2021-02-26 13:02:13 -08:00 committed by GitHub
parent c1b0cf6d0b
commit 059ed1c241
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
2 changed files with 134 additions and 119 deletions

View file

@ -1,3 +1,4 @@
import functools
import io
import logging
import onnx
@ -75,6 +76,130 @@ class ORTModule(torch.nn.Module):
def __init__(self, module):
assert isinstance(module, torch.nn.Module), "'module' must be a torch.nn.Module"
# Create forward dynamically, so each ORTModule instance will have its own copy.
# This is needed to be able to copy the forward signatures from the original PyTorch models
# and possibly have different signatures for different instances.
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, we build a full training graph with module_gradient_graph_builder.
Finally, we instantiate the ONNX Runtime InferenceSession.
'''
# TODO: using pytorch for evaluation for now. We will use ORT for evaluation latter.
if not self._is_training:
return self._original_module(*inputs, **kwargs)
# Exporting module to ONNX for the first time
if not self._onnx_training:
device_from_module = _utils.get_device_from_module(self._original_module)
if not self._device or self._device != device_from_module:
self._device = device_from_module
if not self._device:
raise RuntimeError('A device must be specified in the model or data!')
self._get_inference_graph_and_init_gradient_graph_builder(*inputs, **kwargs)
_, _, input_names_require_grad, new_input_shape = \
_ortmodule_output_transformation.parse_inputs_for_onnx_export(
self._original_module_input_names, self._onnx_inference, *inputs, **kwargs)
# If inputs requiring gradient change from one call to forward to the next, the module_gradient_graph_builder
# needs to be reinitialized so it can compute the backward output for the new inputs that require_grad
if input_names_require_grad != self._input_names_require_grad:
self._input_names_require_grad = input_names_require_grad
self._initialize_module_gradient_graph_builder()
if self._current_input_shape is None or self._current_input_shape != new_input_shape:
self._current_input_shape = new_input_shape
self._build_training_graph()
self._create_training_session()
module_device = _utils.get_device_from_module(self._original_module)
if self._device != module_device:
self._device = module_device
self._create_training_session()
# 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, *inputs, **kwargs):
'''Performs forward pass based on user input and PyTorch initializer
Autograd Function's apply() doesn't support keyword arguments,
so `*inputs` has all the arguments - keyword arguments converted
to positional by the caller.
Module outputs are returned to the user
'''
# Use IO binding
_create_iobinding(self._training_io_binding, inputs, self._onnx_training, self._device)
# Run and return module outputs.
forward_outputs, run_id = self._training_session.run_forward(self._training_io_binding, self._run_options)
user_outputs = tuple(_ort_output_to_torch_tensor(forward_output) for forward_output in forward_outputs)
ctx.run_id = run_id
return user_outputs
@staticmethod
def backward(ctx, *grad_outputs):
'''Performs backward pass based on grad wrt module output
'''
# Use IO binding
# Push user output grads to ONNX backend.
backward_grad_output_ortvalue = []
# backward_output_grad_names_map only contains the subset of module outputs that need a gradient,
# we filter out the invalid entries in grad_outputs, accessing using the mapped index.
for _, i in self._onnx_graphs_info.backward_output_grad_names_map.items():
grad_output = grad_outputs[i]
if not grad_output.is_contiguous():
grad_output = grad_output.contiguous()
backward_grad_output_ortvalue.append(onnxruntime.OrtValue.ortvalue_from_data_ptr(list(grad_output.size()), _utils.dtype_torch_to_numpy(
grad_output.dtype), grad_output.device.type, _utils.get_device_index(grad_output.device), grad_output.data_ptr()))
# Run and get results
run_id = ctx.run_id
self._training_session.run_backward(backward_grad_output_ortvalue, run_id)
backward_outputs = self._training_io_binding.get_outputs()
# Return input and initializer gradients
num_user_input_grads = len(self._input_names_require_grad)
results = []
for input_name in self._onnx_graphs_info.user_input_names:
try:
# Append to the results the backward output for each input that required grad
results.append(_ort_output_to_torch_tensor(
backward_outputs[self._input_names_require_grad.index(input_name)]))
except ValueError:
# input_name is not found in the self._input_names_require_grad list
# Append None to results for each input that did not require grad
results.append(None)
# Append gradients of initializer to results
results += [_ort_output_to_torch_tensor(backward_output)
for backward_output in backward_outputs[num_user_input_grads:]]
# The OrtValue has a shared_ptr to the data. At this point there are two shared_ptrs to the data, one through the
# OrtValue in the output iobinding, and the other through the copy in OrtDLManagedTensor.
# The following call clears the iobinding output, reducing the use_count to 1, so that once torch finishes computation
# on the DLpack tensors, the memory can be freed.
self._training_io_binding.clear_binding_outputs()
return tuple(results)
return _ortmodule_output_transformation.populate_user_output_from_schema_and_outputs(self._original_module_output_schema,
self._onnx_graphs_info.user_output_names,
_ORTModuleFunction.apply(*self._convert_training_graph_input_to_list(*inputs, **kwargs)))
# Bind the forward method.
self.forward = _forward.__get__(self)
# Copy the forward signature from the PyTorch module.
functools.update_wrapper(self.forward.__func__, module.forward.__func__)
super(ORTModule, self).__init__()
# Support contrib OPs
@ -186,121 +311,6 @@ class ORTModule(torch.nn.Module):
self._is_training = mode
self._flattened_output_module.train(mode)
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, we build a full training graph with module_gradient_graph_builder.
Finally, we instantiate the ONNX Runtime InferenceSession.
'''
# TODO: using pytorch for evaluation for now. We will use ORT for evaluation latter.
if not self._is_training:
return self._original_module(*inputs, **kwargs)
# Exporting module to ONNX for the first time
if not self._onnx_training:
device_from_module = _utils.get_device_from_module(self._original_module)
if not self._device or self._device != device_from_module:
self._device = device_from_module
if not self._device:
raise RuntimeError('A device must be specified in the model or data!')
self._get_inference_graph_and_init_gradient_graph_builder(*inputs, **kwargs)
_, _, input_names_require_grad, new_input_shape = \
_ortmodule_output_transformation.parse_inputs_for_onnx_export(
self._original_module_input_names, self._onnx_inference, *inputs, **kwargs)
# If inputs requiring gradient change from one call to forward to the next, the module_gradient_graph_builder
# needs to be reinitialized so it can compute the backward output for the new inputs that require_grad
if input_names_require_grad != self._input_names_require_grad:
self._input_names_require_grad = input_names_require_grad
self._initialize_module_gradient_graph_builder()
if self._current_input_shape is None or self._current_input_shape != new_input_shape:
self._current_input_shape = new_input_shape
self._build_training_graph()
self._create_training_session()
module_device = _utils.get_device_from_module(self._original_module)
if self._device != module_device:
self._device = module_device
self._create_training_session()
# 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, *inputs, **kwargs):
'''Performs forward pass based on user input and PyTorch initializer
Autograd Function's apply() doesn't support keyword arguments,
so `*inputs` has all the arguments - keyword arguments converted
to positional by the caller.
Module outputs are returned to the user
'''
# Use IO binding
_create_iobinding(self._training_io_binding, inputs, self._onnx_training, self._device)
# Run and return module outputs.
forward_outputs, run_id = self._training_session.run_forward(self._training_io_binding, self._run_options)
user_outputs = tuple(_ort_output_to_torch_tensor(forward_output) for forward_output in forward_outputs)
ctx.run_id = run_id
return user_outputs
@staticmethod
def backward(ctx, *grad_outputs):
'''Performs backward pass based on grad wrt module output
'''
# Use IO binding
# Push user output grads to ONNX backend.
backward_grad_output_ortvalue = []
# backward_output_grad_names_map only contains the subset of module outputs that need a gradient,
# we filter out the invalid entries in grad_outputs, accessing using the mapped index.
for _, i in self._onnx_graphs_info.backward_output_grad_names_map.items():
grad_output = grad_outputs[i]
if not grad_output.is_contiguous():
grad_output = grad_output.contiguous()
backward_grad_output_ortvalue.append(onnxruntime.OrtValue.ortvalue_from_data_ptr(list(grad_output.size()), _utils.dtype_torch_to_numpy(
grad_output.dtype), grad_output.device.type, _utils.get_device_index(grad_output.device), grad_output.data_ptr()))
# Run and get results
run_id = ctx.run_id
self._training_session.run_backward(backward_grad_output_ortvalue, run_id)
backward_outputs = self._training_io_binding.get_outputs()
# Return input and initializer gradients
num_user_input_grads = len(self._input_names_require_grad)
results = []
for input_name in self._onnx_graphs_info.user_input_names:
try:
# Append to the results the backward output for each input that required grad
results.append(_ort_output_to_torch_tensor(
backward_outputs[self._input_names_require_grad.index(input_name)]))
except ValueError:
# input_name is not found in the self._input_names_require_grad list
# Append None to results for each input that did not require grad
results.append(None)
# Append gradients of initializer to results
results += [_ort_output_to_torch_tensor(backward_output)
for backward_output in backward_outputs[num_user_input_grads:]]
# The OrtValue has a shared_ptr to the data. At this point there are two shared_ptrs to the data, one through the
# OrtValue in the output iobinding, and the other through the copy in OrtDLManagedTensor.
# The following call clears the iobinding output, reducing the use_count to 1, so that once torch finishes computation
# on the DLpack tensors, the memory can be freed.
self._training_io_binding.clear_binding_outputs()
return tuple(results)
return _ortmodule_output_transformation.populate_user_output_from_schema_and_outputs(self._original_module_output_schema,
self._onnx_graphs_info.user_output_names,
_ORTModuleFunction.apply(*self._convert_training_graph_input_to_list(*inputs, **kwargs)))
@_utils.timeit(enabled=__TEMP_ENABLE_METHOD_TIMING__)
def _convert_training_graph_input_to_list(self, *inputs, **kwargs):
'''Creates forward `*inputs` list from user input and PyTorch initializers

View file

@ -13,6 +13,7 @@ import warnings
from unittest.mock import patch
from collections import OrderedDict
from collections import namedtuple
from inspect import signature
from onnxruntime.training import _utils, ORTModule
import _test_helpers
@ -180,10 +181,12 @@ def test_forward_call_single_positional_argument():
N, D_in, H, D_out = 64, 784, 500, 10
model = NeuralNetSinglePositionalArgument(D_in, H, D_out).to(device)
model = ORTModule(model)
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 = model(x)
output = ort_model(x)
assert output is not None
def test_forward_call_multiple_positional_arguments():
@ -191,12 +194,14 @@ def test_forward_call_multiple_positional_arguments():
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)
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)
y = torch.randn(N, D_in, device=device)
# Make sure model runs without any exception
output = model(x, y)
output = ort_model(x, y)
assert output is not None
# TODO: Re-enable after "Support models with dynamically defined inputs" done.