From c76172fab664f20f7b173b80e1a07f66bb955f7d Mon Sep 17 00:00:00 2001 From: Wei-Sheng Chin Date: Wed, 16 Jun 2021 09:19:41 -0700 Subject: [PATCH] Fix PythonOp with input which has no gradient (#8011) * Fix PythonOp with input has no gradient * Fix another bug which happens when inputs require gradient * Remove comments Co-authored-by: Peng Wang --- .../core/graph/gradient_builder.cc | 7 +- .../python/training/ortmodule/_io.py | 15 +++-- .../orttraining_test_ortmodule_autograd.py | 65 +++++++++++++++---- .../torch_custom_function_kernel_base.cc | 8 +-- 4 files changed, 72 insertions(+), 23 deletions(-) diff --git a/orttraining/orttraining/core/graph/gradient_builder.cc b/orttraining/orttraining/core/graph/gradient_builder.cc index 54619ddaa6..534d5c0a8e 100755 --- a/orttraining/orttraining/core/graph/gradient_builder.cc +++ b/orttraining/orttraining/core/graph/gradient_builder.cc @@ -1807,11 +1807,10 @@ IMPLEMENT_GRADIENT_BUILDER(GetPythonOpGradient) { std::vector output_args; for (int i = 0; i < GetSrcNodeInputSize(); ++i) { - // output_args[i] is typed to output_types[i]. - if (output_tensor_types.at(i) == 9) { - output_args.push_back(ArgDef()); - } else { + if (output_tensor_requires_grads[i]) { output_args.push_back(GI(i)); + } else { + output_args.push_back(ArgDef()); } } diff --git a/orttraining/orttraining/python/training/ortmodule/_io.py b/orttraining/orttraining/python/training/ortmodule/_io.py index 8aaa9e908c..0993476b23 100644 --- a/orttraining/orttraining/python/training/ortmodule/_io.py +++ b/orttraining/orttraining/python/training/ortmodule/_io.py @@ -175,13 +175,20 @@ def _combine_input_buffers_initializers(params, onnx_input_names, input_info, bu def deepcopy_model_input(*inputs, **kwargs): - sample_inputs_copy = [model_input.data if isinstance(model_input, torch.Tensor) else model_input - for model_input in inputs] + def extract_tensor(value): + if isinstance(value, torch.Tensor): + if value.requires_grad: + return value.data.requires_grad_() + else: + return value.data + else: + return value + sample_inputs_copy = [extract_tensor(value) for value in inputs] sample_inputs_copy = copy.deepcopy(tuple(sample_inputs_copy)) sample_kwargs_copy = {} - for name, model_input in kwargs.items(): - sample_kwargs_copy[name] = model_input.data if isinstance(model_input, torch.Tensor) else model_input + for name, value in kwargs.items(): + sample_kwargs_copy[name] = extract_tensor(value) sample_kwargs_copy = copy.deepcopy(sample_kwargs_copy) return sample_inputs_copy, sample_kwargs_copy diff --git a/orttraining/orttraining/test/python/orttraining_test_ortmodule_autograd.py b/orttraining/orttraining/test/python/orttraining_test_ortmodule_autograd.py index aeb81a1e33..7720ac1e67 100644 --- a/orttraining/orttraining/test/python/orttraining_test_ortmodule_autograd.py +++ b/orttraining/orttraining/test/python/orttraining_test_ortmodule_autograd.py @@ -50,7 +50,6 @@ def test_GeLU(): device=torch.cuda.current_device(), dtype=torch.float)) - # Always initialize bias to zero. with torch.no_grad(): self.bias.uniform_() @@ -93,7 +92,6 @@ def test_MegatronF(): device=torch.cuda.current_device(), dtype=torch.float)) - # Always initialize bias to zero. with torch.no_grad(): self.bias.uniform_() @@ -189,7 +187,6 @@ def test_InplaceUpdateInputAsOutputNotRequireGrad(): device=torch.cuda.current_device(), dtype=torch.float)) - # Always initialize bias to zero. with torch.no_grad(): self.bias.uniform_() @@ -238,7 +235,6 @@ def test_InplaceUpdateInputNotAsOutputNotRequireGrad(): device=torch.cuda.current_device(), dtype=torch.float)) - # Always initialize bias to zero. with torch.no_grad(): self.bias.uniform_() @@ -293,7 +289,6 @@ def test_InplaceUpdateInputAsOutputNotRequireGradWithMarkDirty(): device=torch.cuda.current_device(), dtype=torch.float)) - # Always initialize bias to zero. with torch.no_grad(): self.bias.uniform_() @@ -340,7 +335,6 @@ def test_InplaceUpdateInputAsOutputRequireGrad(): device=torch.cuda.current_device(), dtype=torch.float)) - # Always initialize bias to zero. with torch.no_grad(): self.bias.uniform_() @@ -396,7 +390,6 @@ def test_InplaceUpdateInputNotAsOutputRequireGrad(): device=torch.cuda.current_device(), dtype=torch.float)) - # Always initialize bias to zero. with torch.no_grad(): self.bias.uniform_() @@ -454,7 +447,6 @@ def test_InplaceUpdateInputAsOutputRequireGradWithMarkDirty(): device=torch.cuda.current_device(), dtype=torch.float)) - # Always initialize bias to zero. with torch.no_grad(): self.bias.uniform_() @@ -501,7 +493,6 @@ def test_EvalTest(): device=torch.cuda.current_device(), dtype=torch.float)) - # Always initialize bias to zero. with torch.no_grad(): self.bias.uniform_() @@ -566,7 +557,6 @@ def test_TwoOutputFunction(): device=torch.cuda.current_device(), dtype=torch.float)) - # Always initialize bias to zero. with torch.no_grad(): self.bias.uniform_() @@ -632,7 +622,6 @@ def test_InnerModuleCall(): self.use_ort = use_ort self.bias = Parameter(torch.FloatTensor([1.0] * dim).to(device)) - # Always initialize bias to zero. with torch.no_grad(): self.bias.uniform_() @@ -660,3 +649,57 @@ def test_InnerModuleCall(): result_ort = get_inner_module_call_result(x.detach(), 'cpu', True) result_pth = get_inner_module_call_result(x.detach(), 'cpu', False) compare_tensor_list(result_ort, result_pth) + + +def test_Share_Input(): + class TwoOutputFunction(torch.autograd.Function): + @staticmethod + # bias is an optional argument + def forward(ctx, x, y): + ctx.save_for_backward(x, y) + w = x + y + z = x * y + return w, z + + @staticmethod + def backward(ctx, dw, dz): + x, y = ctx.saved_tensors + dx = dw * 1.0 + dz * y + dy = dw * 1.0 + dz * x + return dx, dy + + class TwoOutputModel(torch.nn.Module): + def __init__(self, output_size): + super(TwoOutputModel, self).__init__() + self.fun = TwoOutputFunction.apply + self.bias = Parameter(torch.empty( + output_size, + device=torch.cuda.current_device(), + dtype=torch.float)) + + with torch.no_grad(): + self.bias.uniform_() + + def forward(self, x): + a, b = self.fun(x, self.bias) + c, d = self.fun(x, self.bias) + return a + b + c + d + + output_size = 2 + + def model_builder(): + return TwoOutputModel(output_size) + + def input_generator(): + return torch.randn(output_size, dtype=torch.float) + + def input_generator_with_requires_grad(): + return torch.randn(output_size, dtype=torch.float).requires_grad_() + + # generate a label that have same shape as forward output. + label_input = torch.ones([output_size]) + + # Test multi-input and multi-output custom function. + run_training_test_and_compare(model_builder, input_generator, label_input) + + run_training_test_and_compare(model_builder, input_generator_with_requires_grad, label_input) diff --git a/orttraining/orttraining/training_ops/cpu/torch/torch_custom_function_kernel_base.cc b/orttraining/orttraining/training_ops/cpu/torch/torch_custom_function_kernel_base.cc index 9e7b2f545f..8da54b52af 100644 --- a/orttraining/orttraining/training_ops/cpu/torch/torch_custom_function_kernel_base.cc +++ b/orttraining/orttraining/training_ops/cpu/torch/torch_custom_function_kernel_base.cc @@ -258,11 +258,11 @@ void PythonOpGradBase::RunBackward(OpKernelContext* context, void PythonOpGradBase::SetOutputs(OpKernelContext* context, std::vector& returned_ortvalues) const { auto* ctx_internal = reinterpret_cast(context); - auto outputs_count = static_cast(ctx_internal->OutputCount()); - // It's possible that Pytorch returns None as gradient and ORT Python side may skip them. - // In that case, returned_args may contain less arguments. - outputs_count = outputs_count > returned_ortvalues.size() ? returned_ortvalues.size() : outputs_count; + const auto outputs_count = static_cast(ctx_internal->OutputCount()); for (size_t i = 0; i < outputs_count; ++i) { + if (!output_tensor_requires_grads_[i]) { + continue; + } ORT_THROW_IF_ERROR(ctx_internal->SetOutputMLValue(static_cast(i), returned_ortvalues.at(i))); } }