From 0695a11e5fb729aed2d16ddd4a451bf38703bfa3 Mon Sep 17 00:00:00 2001 From: wangyanling10 Date: Fri, 8 Apr 2022 10:45:09 +0800 Subject: [PATCH] restore conv2d infer bug --- mindspore/core/ops/conv2d.cc | 41 +++++++++++++++++++++++++++++++----- 1 file changed, 36 insertions(+), 5 deletions(-) diff --git a/mindspore/core/ops/conv2d.cc b/mindspore/core/ops/conv2d.cc index fe5207d9f62..b6129b982ae 100644 --- a/mindspore/core/ops/conv2d.cc +++ b/mindspore/core/ops/conv2d.cc @@ -36,6 +36,11 @@ constexpr size_t stride_num = 2; constexpr size_t dilation_num = 2; constexpr size_t padding_num = 4; constexpr size_t start_index = 2; +constexpr size_t top_padding = 0; +constexpr size_t bottom_padding = 1; +constexpr size_t left_padding = 2; +constexpr size_t right_padding = 3; + void CheckShapeAnyAndPositive(const std::string &op, const ShapeVector &shape) { for (size_t i = 0; i < shape.size(); ++i) { if ((shape[i] < 0) && (shape[i] != Shape::SHP_ANY)) { @@ -153,12 +158,30 @@ void Conv2DPadFunction(std::vector *output_hw, std::vector *pa } } -abstract::ShapePtr Conv2dInferShape(const PrimitivePtr &primitive, const std::vector &input_args) { - MS_EXCEPTION_IF_NULL(primitive); - auto prim_name = primitive->name(); - for (const auto &item : input_args) { - MS_EXCEPTION_IF_NULL(item); +bool CheckConv2dShape(const std::string &prim_name, const std::vector &input_args, + const std::vector &x_shape, const std::vector &w_shape, + const std::vector &padding, int64_t pad_mode, uint64_t w_axis, uint64_t h_axis) { + auto x_shape_ptr = CheckAndConvertUtils::GetTensorInputShape(prim_name, input_args, 0); + auto w_shape_ptr = CheckAndConvertUtils::GetTensorInputShape(prim_name, input_args, 1); + if (x_shape_ptr->IsDynamic() || w_shape_ptr->IsDynamic()) { + return true; } + if (w_shape[w_axis] != Shape::SHP_ANY && pad_mode != PadMode::SAME) { + int64_t input_height = x_shape[h_axis]; + int64_t input_width = x_shape[w_axis]; + if (pad_mode == PadMode::PAD) { + input_height += padding[left_padding] + padding[right_padding]; + input_width += padding[top_padding] + padding[bottom_padding]; + } + if (input_height < w_shape[h_axis] || input_width < w_shape[w_axis]) { + return false; + } + } + return true; +} + +abstract::ShapePtr Conv2dInferShape(const PrimitivePtr &primitive, const std::vector &input_args) { + auto prim_name = primitive->name(); auto x_shape_map = CheckAndConvertUtils::ConvertShapePtrToShapeMap(input_args[0]->BuildShape()); auto w_shape_map = CheckAndConvertUtils::ConvertShapePtrToShapeMap(input_args[1]->BuildShape()); auto x_shape = x_shape_map[kShape]; @@ -219,6 +242,11 @@ abstract::ShapePtr Conv2dInferShape(const PrimitivePtr &primitive, const std::ve std::vector padding = CheckAttrIntOrTuple(primitive->GetAttr("pad"), 0, padding_num); int64_t pad_mode; CheckAndConvertUtils::GetPadModEnumValue(primitive->GetAttr("pad_mode"), &pad_mode); + if (!CheckConv2dShape(prim_name, input_args, x_shape, w_shape, padding, pad_mode, w_axis, h_axis)) { + MS_LOG(EXCEPTION) + << "Shape error for Conv2d, input shape's h and w after padding is less than kernel_size's h and w dims."; + } + std::vector output_hw; std::vector pad_list; std::vector output_hw_min; @@ -231,6 +259,7 @@ abstract::ShapePtr Conv2dInferShape(const PrimitivePtr &primitive, const std::ve dilation, pad_mode, padding, true); Conv2DPadFunction(&output_hw_max, &pad_list_max, x_max_shape[h_axis], x_max_shape[w_axis], kernel_size, stride, dilation, pad_mode, padding); + std::vector pad_list_val = {MakeValue(pad_list[0]), MakeValue(pad_list[1]), MakeValue(pad_list[2]), MakeValue(pad_list[3])}; primitive->set_attr("pad_list", MakeValue(pad_list_val)); @@ -374,9 +403,11 @@ Format Conv2D::get_format() const { AbstractBasePtr Conv2dInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { + MS_EXCEPTION_IF_NULL(primitive); for (auto item : input_args) { MS_EXCEPTION_IF_NULL(item); } + const int64_t input_num = 2; (void)CheckAndConvertUtils::CheckInteger("Conv2d infer", SizeToLong(input_args.size()), kGreaterEqual, input_num, primitive->name());