mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-20 19:12:24 +00:00
gradient and test (#10455)
Co-authored-by: Aishwarya Bhandare <aibhanda@microsoft.com@orttrainingdev8.d32nl1ml4oruzj4qz3bqlggovf.px.internal.cloudapp.net>
This commit is contained in:
parent
435e14d60a
commit
7e5d68eea6
4 changed files with 75 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Reference in a new issue