mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-28 20:11:22 +00:00
Copy forward signature from PyTorch model. (#6777)
This commit is contained in:
parent
c1b0cf6d0b
commit
059ed1c241
2 changed files with 134 additions and 119 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Reference in a new issue