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 <tuananhnguyen7198@gmail.com>
This commit is contained in:
Anh Nguyen 2022-02-15 22:38:19 -08:00 committed by GitHub
parent f436d3437e
commit 7443edb0bf
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
4 changed files with 81 additions and 3 deletions

69
orttraining/orttraining/core/graph/gradient_builder.cc Executable file → Normal file
View file

@ -1051,6 +1051,75 @@ IMPLEMENT_GRADIENT_BUILDER(GetReduceMeanGradient) {
return result;
}
IMPLEMENT_GRADIENT_BUILDER(GetReduceMaxGradient) {
std::vector<NodeDef> result;
std::vector<int64_t> 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<bool>(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<int64_t> default_reduce_axes = {};
ArgDef reduce_axes_arg_def = IA("ReduceAxes");
if (!keepdims) {
if (attributes.find("axes") != attributes.end()) {
axes_values = RetrieveValues<int64_t>(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<int64_t>(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)))

View file

@ -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

View file

@ -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);
};

View file

@ -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};