gradient and test (#10455)

Co-authored-by: Aishwarya Bhandare <aibhanda@microsoft.com@orttrainingdev8.d32nl1ml4oruzj4qz3bqlggovf.px.internal.cloudapp.net>
This commit is contained in:
ashbhandare 2022-02-08 10:18:22 -08:00 committed by GitHub
parent 435e14d60a
commit 7e5d68eea6
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
4 changed files with 75 additions and 0 deletions

View file

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

View file

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

View file

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

View file

@ -2760,6 +2760,60 @@ TEST(GradientCheckerTest, ScatterNDGrad) {
}
}
TEST(GradientCheckerTest, ScatterElementsGrad) {
float max_error;
GradientChecker<float, float, float> 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<int64_t>());
TensorInfo updates_info({2, 3}, true);
std::vector<std::vector<float>> 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<int64_t>());
TensorInfo updates_info({1, 2}, true);
std::vector<std::vector<float>> 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<int64_t>(1))}));
EXPECT_IS_TINY(max_error);
}
{ // with -ve axis
TensorInfo data_info({1, 5}, true);
TensorInfo indices_info({1, 2}, false, nullptr, DataTypeImpl::GetTensorType<int64_t>());
TensorInfo updates_info({1, 2}, true);
std::vector<std::vector<float>> 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<int64_t>(-1))}));
EXPECT_IS_TINY(max_error);
}
}
TEST(GradientCheckerTest, TriluGrad) {
float max_error;
GradientChecker<float, float, float> gradient_checker;