diff --git a/modules/dnn/src/layers/softmax_layer.cpp b/modules/dnn/src/layers/softmax_layer.cpp index a4d5f44128..63a49f6233 100644 --- a/modules/dnn/src/layers/softmax_layer.cpp +++ b/modules/dnn/src/layers/softmax_layer.cpp @@ -92,7 +92,8 @@ public: bool inplace = Layer::getMemoryShapes(inputs, requiredOutputs, outputs, internals); MatShape shape = inputs[0]; int cAxis = normalize_axis(axisRaw, shape.size()); - shape[cAxis] = 1; + if (!shape.empty()) + shape[cAxis] = 1; internals.assign(1, shape); return inplace; } diff --git a/modules/dnn/test/test_layers_1d.cpp b/modules/dnn/test/test_layers_1d.cpp index 16cb2836fe..3a0f9540d3 100644 --- a/modules/dnn/test/test_layers_1d.cpp +++ b/modules/dnn/test/test_layers_1d.cpp @@ -381,4 +381,53 @@ INSTANTIATE_TEST_CASE_P(/*nothing*/, Layer_Concat_Test, std::vector({}), std::vector({1}) )); + +typedef testing::TestWithParam, int>> Layer_Softmax_Test; +TEST_P(Layer_Softmax_Test, Accuracy_01D) { + + int axis = get<1>(GetParam()); + std::vector input_shape = get<0>(GetParam()); + if ((input_shape.size() == 0 && axis == 1) || + (!input_shape.empty() && input_shape.size() == 2 && input_shape[0] > 1 && axis == 1) || + (!input_shape.empty() && input_shape[0] > 1 && axis == 0)) // skip since not valid case + return; + + LayerParams lp; + lp.type = "Softmax"; + lp.name = "softmaxLayer"; + lp.set("axis", axis); + Ptr layer = SoftmaxLayer::create(lp); + + Mat input = Mat(input_shape.size(), input_shape.data(), CV_32F); + cv::randn(input, 0.0, 1.0); + + Mat output_ref; + cv::exp(input, output_ref); + if (axis == 1){ + cv::divide(output_ref, cv::sum(output_ref), output_ref); + } else { + cv::divide(output_ref, output_ref, output_ref); + } + + std::vector inputs{input}; + std::vector outputs; + runLayer(layer, inputs, outputs); + ASSERT_EQ(outputs.size(), 1); + ASSERT_EQ(shape(output_ref), shape(outputs[0])); + normAssert(output_ref, outputs[0]); +} + +INSTANTIATE_TEST_CASE_P(/*nothing*/, Layer_Softmax_Test, Combine( + /*input blob shape*/ + testing::Values( + std::vector({}), + std::vector({1}), + std::vector({4}), + std::vector({1, 4}), + std::vector({4, 1}) + ), + /*Axis */ + testing::Values(0, 1) +)); + }}