From 368e4a324f5235823cd8a948c505536b13e31ada Mon Sep 17 00:00:00 2001 From: Vincent Wang Date: Mon, 26 Apr 2021 09:12:03 +0800 Subject: [PATCH] SqueezeGrad Bugfix (#7412) * squeezegrad bugfix * fix ut Co-authored-by: Vincent Wang --- .../core/graph/gradient_builder.cc | 38 ++++++++----------- .../test/gradient/gradient_ops_test.cc | 8 ++++ 2 files changed, 24 insertions(+), 22 deletions(-) diff --git a/orttraining/orttraining/core/graph/gradient_builder.cc b/orttraining/orttraining/core/graph/gradient_builder.cc index 9630f45125..c75e717b77 100755 --- a/orttraining/orttraining/core/graph/gradient_builder.cc +++ b/orttraining/orttraining/core/graph/gradient_builder.cc @@ -788,36 +788,30 @@ IMPLEMENT_GRADIENT_BUILDER(GetReluGradient) { } IMPLEMENT_GRADIENT_BUILDER(GetSqueezeGradient) { - std::vector result; size_t numInputs = GetSrcNodeInputSize(); - if (SrcNodeOpsetVersion() < 13) { //axes attribute + if (SrcNodeOpsetVersion() < 13) { // Axes attribute exists. auto attributes = SrcNodeAttributes(); std::vector axes_values; if (attributes.find("axes") != attributes.end()) { axes_values = RetrieveValues(attributes.at("axes")); - result.push_back( - NodeDef("Unsqueeze", - {GO(0)}, - {GI(0)}, - {MakeAttribute("axes", axes_values)})); + return std::vector{NodeDef("Unsqueeze", + {GO(0)}, + {GI(0)}, + {MakeAttribute("axes", axes_values)})}; } - } else if (numInputs == 2) { //optional input 'axes' is provided - result.push_back( - NodeDef(OpDef{"Unsqueeze", kOnnxDomain, 13}, - {GO(0), I(1)}, - {GI(0)})); - } else { // if axes attribute/input not provided for squeeze - result.push_back( - NodeDef("Shape", - {I(0)}, - {IA("I0_shape")})); - result.push_back( - NodeDef("Reshape", - {GO(0), IA("I0_shape")}, - {GI(0)})); + } else if (numInputs == 2) { // Optional input 'axes' is provided + return std::vector{NodeDef(OpDef{"Unsqueeze", kOnnxDomain, 13}, + {GO(0), I(1)}, + {GI(0)})}; } - return result; + // If axes attribute/input is not provided for squeeze, no matter which OpSet version. + return std::vector{NodeDef("Shape", + {I(0)}, + {IA("I0_shape")}), + NodeDef("Reshape", + {GO(0), IA("I0_shape")}, + {GI(0)})}; } IMPLEMENT_GRADIENT_BUILDER(GetAddSubGradient) { diff --git a/orttraining/orttraining/test/gradient/gradient_ops_test.cc b/orttraining/orttraining/test/gradient/gradient_ops_test.cc index c0cc462d2f..63c651d4ff 100755 --- a/orttraining/orttraining/test/gradient/gradient_ops_test.cc +++ b/orttraining/orttraining/test/gradient/gradient_ops_test.cc @@ -1308,6 +1308,14 @@ static void RunSqueezeUnsqueezeTests(const OpDef& op_def, x_datas.push_back(random.Gaussian(x_shapes[i], 0.f, 5.f)); std::vector input = {x_shape}; std::vector attributes = {}; + + // Test case w/o axes attribute/input, only valid for Squeeze Op. + if (op_def.type == "Squeeze") { + gradient_checker.ComputeGradientError(op_def, input, {y_shape}, &max_error, x_datas, attributes); + EXPECT_IS_TINIER_THAN(max_error, error_tolerance); + } + + // test case w/ axes attribute/input. if (axes_input) { std::vector axes_float; std::transform(begin(axes), end(axes), std::back_inserter(axes_float), [](int64_t i) { return static_cast(i); });