mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-24 19:43:35 +00:00
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:
parent
f436d3437e
commit
7443edb0bf
4 changed files with 81 additions and 3 deletions
69
orttraining/orttraining/core/graph/gradient_builder.cc
Executable file → Normal file
69
orttraining/orttraining/core/graph/gradient_builder.cc
Executable file → Normal 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)))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
Loading…
Reference in a new issue