diff --git a/orttraining/orttraining/core/graph/gradient_builder.cc b/orttraining/orttraining/core/graph/gradient_builder.cc index 55b521e13a..f807fe6af4 100755 --- a/orttraining/orttraining/core/graph/gradient_builder.cc +++ b/orttraining/orttraining/core/graph/gradient_builder.cc @@ -1838,5 +1838,24 @@ IMPLEMENT_GRADIENT_BUILDER(GetScatterNDGradient) { return result; } +IMPLEMENT_GRADIENT_BUILDER(GetScatterElementsGradient) { + auto attributes = SrcNodeAttributes(); + auto axis = utils::HasInt(attributes.at("axis")) ? attributes.at("axis").i() : 0; + std::vector result; + if (IsGradientRequiredForSrcNodeInput(0)) { + result.emplace_back(NodeDef("Shape", {I(2)}, {IA("Shape_updates")})); + result.emplace_back(NodeDef("ConstantOfShape", {IA("Shape_updates")}, {IA("Zero_Shape_updates")}, + {MakeAttribute("value", ScalarTensorProtoByElemType(0.0f, IElemType(0)))})); + result.emplace_back(NodeDef("ScatterElements", {GO(0), I(1), IA("Zero_Shape_updates")}, {GI(0)}, + {MakeAttribute("axis", axis)})); + } + + if (IsGradientRequiredForSrcNodeInput(2)) { + result.emplace_back(NodeDef("GatherElements", {GO(0), I(1)}, {GI(2)}, + {MakeAttribute("axis", axis)})); + } + return result; +} + } // namespace training } // namespace onnxruntime diff --git a/orttraining/orttraining/core/graph/gradient_builder.h b/orttraining/orttraining/core/graph/gradient_builder.h index 8947f40329..9edccb02cb 100755 --- a/orttraining/orttraining/core/graph/gradient_builder.h +++ b/orttraining/orttraining/core/graph/gradient_builder.h @@ -77,6 +77,7 @@ DECLARE_GRADIENT_BUILDER(GetPadGradient) DECLARE_GRADIENT_BUILDER(GetIdentityGradient) DECLARE_GRADIENT_BUILDER(GetPythonOpGradient) DECLARE_GRADIENT_BUILDER(GetScatterNDGradient) +DECLARE_GRADIENT_BUILDER(GetScatterElementsGradient) DECLARE_GRADIENT_BUILDER(GetTriluGradient) DECLARE_GRADIENT_BUILDER(GetExternalGradient) diff --git a/orttraining/orttraining/core/graph/gradient_builder_registry.cc b/orttraining/orttraining/core/graph/gradient_builder_registry.cc index 6fc1fda644..df728a715e 100755 --- a/orttraining/orttraining/core/graph/gradient_builder_registry.cc +++ b/orttraining/orttraining/core/graph/gradient_builder_registry.cc @@ -108,6 +108,7 @@ void GradientBuilderRegistry::RegisterGradientBuilders() { REGISTER_GRADIENT_BUILDER("Identity", GetIdentityGradient); REGISTER_GRADIENT_BUILDER("PythonOp", GetPythonOpGradient); REGISTER_GRADIENT_BUILDER("ScatterND", GetScatterNDGradient); + REGISTER_GRADIENT_BUILDER("ScatterElements", GetScatterElementsGradient); REGISTER_GRADIENT_BUILDER("Trilu", GetTriluGradient); 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 3b7818a252..05c27198df 100644 --- a/orttraining/orttraining/test/gradient/gradient_ops_test.cc +++ b/orttraining/orttraining/test/gradient/gradient_ops_test.cc @@ -2760,6 +2760,60 @@ TEST(GradientCheckerTest, ScatterNDGrad) { } } +TEST(GradientCheckerTest, ScatterElementsGrad) { + float max_error; + GradientChecker gradient_checker; + OpDef op_def{"ScatterElements", kOnnxDomain, 13}; + + { // without axis + TensorInfo data_info({3, 3}, true); + TensorInfo indices_info({2, 3}, false, nullptr, DataTypeImpl::GetTensorType()); + TensorInfo updates_info({2, 3}, true); + std::vector> input_datas = {{ 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, + 0.0f, 0.0f, 0.0f, 0.0f}, + {1, 0, 2, 0, 2, 1}, + {1.0f, 1.1f, 1.2f, 2.0f, 2.1f, 2.2f}}; + + TensorInfo output_info({3, 3}, true); + + ASSERT_STATUS_OK(gradient_checker.ComputeGradientError(op_def, {data_info, indices_info, updates_info}, + {output_info}, &max_error, input_datas)); + EXPECT_IS_TINY(max_error); + } + + { // with axis + TensorInfo data_info({1, 5}, true); + TensorInfo indices_info({1, 2}, false, nullptr, DataTypeImpl::GetTensorType()); + TensorInfo updates_info({1, 2}, true); + std::vector> input_datas = {{1.0f, 2.0f, 3.0f, 4.0f, 5.0f}, + {1, 3}, + {1.1f, 2.1f}}; + + TensorInfo output_info({1, 5}, true); + + ASSERT_STATUS_OK(gradient_checker.ComputeGradientError(op_def, {data_info, indices_info, updates_info}, + {output_info}, &max_error, input_datas, + {MakeAttribute("axis", static_cast(1))})); + EXPECT_IS_TINY(max_error); + } + + { // with -ve axis + TensorInfo data_info({1, 5}, true); + TensorInfo indices_info({1, 2}, false, nullptr, DataTypeImpl::GetTensorType()); + TensorInfo updates_info({1, 2}, true); + std::vector> input_datas = {{1.0f, 2.0f, 3.0f, 4.0f, 5.0f}, + {1, 3}, + {1.1f, 2.1f}}; + + TensorInfo output_info({1, 5}, true); + + ASSERT_STATUS_OK(gradient_checker.ComputeGradientError(op_def, {data_info, indices_info, updates_info}, + {output_info}, &max_error, input_datas, + {MakeAttribute("axis", static_cast(-1))})); + EXPECT_IS_TINY(max_error); + } +} + TEST(GradientCheckerTest, TriluGrad) { float max_error; GradientChecker gradient_checker;