From 7443edb0bfc0e147b660b8f43e0e73df8117a98a Mon Sep 17 00:00:00 2001 From: Anh Nguyen <94985387+anhnguyen7198@users.noreply.github.com> Date: Tue, 15 Feb 2022 22:38:19 -0800 Subject: [PATCH] Reduce max gradient (#9859) * ReduceMax gradient builder * Update gradient_builder.cc * Add CI fix * Remove whitepace * Update gradient_builder.cc * Update gradient_ops_test.cc * Fix Window CI tests Co-authored-by: root --- .../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, 81 insertions(+), 3 deletions(-) mode change 100755 => 100644 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 100755 new mode 100644 index f807fe6af4..ba3a1e66f7 --- a/orttraining/orttraining/core/graph/gradient_builder.cc +++ b/orttraining/orttraining/core/graph/gradient_builder.cc @@ -1051,6 +1051,75 @@ 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 9edccb02cb..af3a93d26e 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 df728a715e..ce4fb6aca0 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 05c27198df..ccd3d14fbf 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,6 +611,15 @@ 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};