From 4f76c386862bfe6dcba5a427969f0071b95a9ab9 Mon Sep 17 00:00:00 2001 From: ytaous <4484531+ytaous@users.noreply.github.com> Date: Wed, 16 Feb 2022 16:02:30 -0800 Subject: [PATCH] Revert "Reduce max gradient (#9859)" (#10574) This reverts commit 7443edb0bfc0e147b660b8f43e0e73df8117a98a. --- .../core/graph/gradient_builder.cc | 69 ------------------- .../orttraining/core/graph/gradient_builder.h | 2 +- .../core/graph/gradient_builder_registry.cc | 2 +- .../test/gradient/gradient_ops_test.cc | 11 +-- 4 files changed, 3 insertions(+), 81 deletions(-) mode change 100644 => 100755 orttraining/orttraining/core/graph/gradient_builder.cc diff --git a/orttraining/orttraining/core/graph/gradient_builder.cc b/orttraining/orttraining/core/graph/gradient_builder.cc old mode 100644 new mode 100755 index ba3a1e66f7..f807fe6af4 --- a/orttraining/orttraining/core/graph/gradient_builder.cc +++ b/orttraining/orttraining/core/graph/gradient_builder.cc @@ -1051,75 +1051,6 @@ IMPLEMENT_GRADIENT_BUILDER(GetReduceMeanGradient) { return result; } -IMPLEMENT_GRADIENT_BUILDER(GetReduceMaxGradient) { - std::vector result; - std::vector axes_values; - int opset_version = SrcNodeDomain() == kOnnxDomain ? SrcNodeOpsetVersion() : OnnxOpSetVersion(); - auto attributes = SrcNodeAttributes(); - bool keepdims = true; - if (attributes.find("keepdims") != attributes.end() && - attributes.at("keepdims").has_i()) { - keepdims = static_cast(attributes.at("keepdims").i()); - } - - NodeDef zero_constant_node = ZeroConstantNode(IElemType(0)); - ArgDef ZERO = zero_constant_node.output_args[0]; - result.push_back(zero_constant_node); - NodeDef one_constant_node = OneConstantNode(IElemType(0)); - ArgDef ONE = one_constant_node.output_args[0]; - result.push_back(one_constant_node); - ArgDef grad = GO(0); - - std::vector default_reduce_axes = {}; - ArgDef reduce_axes_arg_def = IA("ReduceAxes"); - - if (!keepdims) { - if (attributes.find("axes") != attributes.end()) { - axes_values = RetrieveValues(attributes.at("axes")); - grad = IA("Unsqueezed_Grad"); - result.push_back(NodeDef("Unsqueeze", {GO(0)}, {grad}, {MakeAttribute("axes", axes_values)})); - result.push_back(NodeDef("Unsqueeze", {O(0)}, {IA("Unsqueezed_Output")}, {MakeAttribute("axes", axes_values)})); - result.push_back(NodeDef("Equal", {I(0), IA("Unsqueezed_Output")}, {IA("Mask")})); - result.push_back(NodeDef("Where", {IA("Mask"), ONE, ZERO}, {IA("Mask_float")})); - AddReduceSumNode(IA("Mask_float"), IA("ReduceSum_Mask"), axes_values, true, result); - } else { // axes is not available, O(0) will be a scalar - if (opset_version >= 13) { - result.push_back(NodeDef(OpDef{"ReduceSum", kOnnxDomain, opset_version}, - {IA("Mask_float")}, {IA("ReduceSum_Mask")}, - {{"keepdims", ONNX_NAMESPACE::MakeAttribute("keepdims", int64_t{1})}})); - } - else { - result.push_back(ConstantVectorNode(default_reduce_axes, reduce_axes_arg_def.name)); - result.push_back(NodeDef(OpDef{"ReduceSumTraining", kMSDomain, 1}, - {IA("Mask_float"), reduce_axes_arg_def}, {IA("ReduceSum_Mask")}, - {{"keepdims", ONNX_NAMESPACE::MakeAttribute("keepdims", int64_t{1})}})); - } - } - } else { - result.push_back(NodeDef("Equal", {I(0), O(0)}, {IA("Mask")})); - result.push_back(NodeDef("Where", {IA("Mask"), ONE, ZERO}, {IA("Mask_float")})); - if (attributes.find("axes") != attributes.end()) { - axes_values = RetrieveValues(attributes.at("axes")); - AddReduceSumNode(IA("Mask_float"), IA("ReduceSum_Mask"), axes_values, true, result); - } else { // axes is not available, O(0) will be a scalar - if (opset_version >= 13) { - result.push_back(NodeDef(OpDef{"ReduceSum", kOnnxDomain, opset_version}, - {IA("Mask_float")}, {IA("ReduceSum_Mask")}, - {{"keepdims", ONNX_NAMESPACE::MakeAttribute("keepdims", int64_t{1})}})); - } - else { - result.push_back(ConstantVectorNode(default_reduce_axes, reduce_axes_arg_def.name)); - result.push_back(NodeDef(OpDef{"ReduceSumTraining", kMSDomain, 1}, - {IA("Mask_float"), reduce_axes_arg_def}, {IA("ReduceSum_Mask")}, - {{"keepdims", ONNX_NAMESPACE::MakeAttribute("keepdims", int64_t{1})}})); - } - } - } - result.push_back(NodeDef("Div", {grad, IA("ReduceSum_Mask")}, {IA("Scaled_Grad")})); - result.push_back(NodeDef("Mul", {IA("Scaled_Grad"), IA("Mask_float")}, {GI(0)})); - return result; -} - // Reference computation is pytorch's logsumexp_backward // dx_i = exp(xi) / reduceSum(exp(xi)) // O(0) = log(reduceSum(exp(xi))) diff --git a/orttraining/orttraining/core/graph/gradient_builder.h b/orttraining/orttraining/core/graph/gradient_builder.h index af3a93d26e..9edccb02cb 100755 --- a/orttraining/orttraining/core/graph/gradient_builder.h +++ b/orttraining/orttraining/core/graph/gradient_builder.h @@ -79,7 +79,7 @@ DECLARE_GRADIENT_BUILDER(GetPythonOpGradient) DECLARE_GRADIENT_BUILDER(GetScatterNDGradient) DECLARE_GRADIENT_BUILDER(GetScatterElementsGradient) DECLARE_GRADIENT_BUILDER(GetTriluGradient) -DECLARE_GRADIENT_BUILDER(GetReduceMaxGradient) + DECLARE_GRADIENT_BUILDER(GetExternalGradient) } // namespace training diff --git a/orttraining/orttraining/core/graph/gradient_builder_registry.cc b/orttraining/orttraining/core/graph/gradient_builder_registry.cc index ce4fb6aca0..df728a715e 100755 --- a/orttraining/orttraining/core/graph/gradient_builder_registry.cc +++ b/orttraining/orttraining/core/graph/gradient_builder_registry.cc @@ -110,7 +110,7 @@ void GradientBuilderRegistry::RegisterGradientBuilders() { REGISTER_GRADIENT_BUILDER("ScatterND", GetScatterNDGradient); REGISTER_GRADIENT_BUILDER("ScatterElements", GetScatterElementsGradient); REGISTER_GRADIENT_BUILDER("Trilu", GetTriluGradient); - REGISTER_GRADIENT_BUILDER("ReduceMax", GetReduceMaxGradient); + REGISTER_GRADIENT_BUILDER("ExternalGradient", GetExternalGradient); }; diff --git a/orttraining/orttraining/test/gradient/gradient_ops_test.cc b/orttraining/orttraining/test/gradient/gradient_ops_test.cc index ccd3d14fbf..05c27198df 100644 --- a/orttraining/orttraining/test/gradient/gradient_ops_test.cc +++ b/orttraining/orttraining/test/gradient/gradient_ops_test.cc @@ -1,4 +1,4 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. #ifdef NDEBUG // disable for debug builds because some of these tests are slow @@ -611,15 +611,6 @@ TEST(GradientCheckerTest, ReduceMeanGrad) { RunReductionTests(op_def); } -TEST(GradientCheckerTest, ReduceMaxGrad) { - // Attribute axes supports negative values from opset 11. - OpDef op_def_11{"ReduceMax", kOnnxDomain, 11}; - RunReductionTests(op_def_11); - - OpDef op_def_13{"ReduceMax", kOnnxDomain, 11}; - RunReductionTests(op_def_13); -} - TEST(GradientCheckerTest, ReduceSumGrad) { // Attribute axes supports negative values from opset 11. OpDef op_def_11{"ReduceSum", kOnnxDomain, 11};