From 6c271c63ac98c05ddf5c017a419fab1ef29cd7a7 Mon Sep 17 00:00:00 2001 From: pengwa Date: Sun, 4 Aug 2019 17:18:45 +0800 Subject: [PATCH] add test cases for commit c019bb9355a511f471e55e7302b26e1d370ed46a (#1556) --- .../core/providers/cuda/tensor/transpose.cc | 4 +- .../providers/cpu/tensor/transpose_test.cc | 116 ++++++++++-------- 2 files changed, 64 insertions(+), 56 deletions(-) diff --git a/onnxruntime/core/providers/cuda/tensor/transpose.cc b/onnxruntime/core/providers/cuda/tensor/transpose.cc index 53f6d32715..49dd7a0801 100644 --- a/onnxruntime/core/providers/cuda/tensor/transpose.cc +++ b/onnxruntime/core/providers/cuda/tensor/transpose.cc @@ -20,13 +20,13 @@ namespace cuda { .TypeConstraint("T", DataTypeImpl::GetTensorType()), \ Transpose); -// special case acceleration using cublas matrix tranpose +// special case acceleration using cublas matrix transpose std::tuple TryTransposeWithCublas(const std::vector& perm, const TensorShape& input_shape) { int M = 0; int N = 0; if (perm.size() == 4 && input_shape[0] == 1 && perm[0] == 0) { - // NCHW < ->NHWC when N == 1 + // NCHW <-> NHWC when N == 1 if ((perm[1] == 2 && perm[2] == 3 && perm[3] == 1) || (perm[1] == 3 && perm[2] == 1 && perm[3] == 2)) { if (perm[1] == 2) { diff --git a/onnxruntime/test/providers/cpu/tensor/transpose_test.cc b/onnxruntime/test/providers/cpu/tensor/transpose_test.cc index 0aa6d665e4..d09d711105 100644 --- a/onnxruntime/test/providers/cpu/tensor/transpose_test.cc +++ b/onnxruntime/test/providers/cpu/tensor/transpose_test.cc @@ -43,7 +43,7 @@ TEST(TransposeOpTest, TwoDimNoAttr) { 2.0f, 5.0f, 3.0f, 6.0f}; - TransposeTest(input_shape, input_vals, nullptr, expected_shape, expected_vals, false);//TensorRT: SegFault error + TransposeTest(input_shape, input_vals, nullptr, expected_shape, expected_vals, false); //TensorRT: SegFault error } TEST(TransposeOpTest, TwoDimNoAttrStr) { @@ -113,37 +113,23 @@ TEST(TransposeOpTest, ThreeDim) { std::vector perm = {0, 2, 1}; std::vector expected_shape({4, 3, 2}); auto expected_vals = { - 1.0f, - 4.0f, - 2.0f, - 5.0f, - 3.0f, - 6.0f, + 1.0f, 4.0f, + 2.0f, 5.0f, + 3.0f, 6.0f, - 1.1f, - 4.1f, - 2.1f, - 5.1f, - 3.1f, - 6.1f, + 1.1f, 4.1f, + 2.1f, 5.1f, + 3.1f, 6.1f, - 1.2f, - 4.2f, - 2.2f, - 5.2f, - 3.2f, - 6.2f, + 1.2f, 4.2f, + 2.2f, 5.2f, + 3.2f, 6.2f, - 1.3f, - 4.3f, - 2.3f, - 5.3f, - 3.3f, - 6.3f, + 1.3f, 4.3f, + 2.3f, 5.3f, + 3.3f, 6.3f}; - }; - - TransposeTest(input_shape, input_vals, &perm, expected_shape, expected_vals, false); //TensorRT: illegal error + TransposeTest(input_shape, input_vals, &perm, expected_shape, expected_vals, false); //TensorRT: illegal error } TEST(TransposeOpTest, ThreeDimStr) { @@ -164,38 +150,60 @@ TEST(TransposeOpTest, ThreeDimStr) { std::vector perm = {0, 2, 1}; std::vector expected_shape({4, 3, 2}); std::initializer_list expected_vals = { - "1", - "4", - "2", - "5", - "3", - "6", + "1", "4", + "2", "5", + "3", "6", - "1", - "4", - "2", - "5", - "3", - "6", + "1", "4", + "2", "5", + "3", "6", - "1", - "4", - "2", - "5", - "3", - "6", + "1", "4", + "2", "5", + "3", "6", - "1", - "4", - "2", - "5", - "3", - "6" - - }; + "1", "4", + "2", "5", + "3", "6"}; TransposeTest(input_shape, input_vals, &perm, expected_shape, expected_vals); } +TEST(TransposeOpTest, NCHW2NHWC) { + std::vector input_shape({1, 3, 2, 2}); + std::vector input_vals = { + "1", "2", "3", "4", + "5", "6", "7", "8", + "9", "10", "11", "12"}; + + std::vector perm = {0, 2, 3, 1}; + std::vector expected_shape({1, 2, 2, 3}); + std::initializer_list expected_vals = { + "1", "5", "9", + "2", "6", "10", + "3", "7", "11", + "4", "8", "12"}; + + TransposeTest(input_shape, input_vals, &perm, expected_shape, expected_vals, false); +} + +TEST(TransposeOpTest, NHWC2NCHW) { + std::vector input_shape({1, 2, 2, 3}); + std::vector input_vals = { + "1", "2", "3", + "4", "5", "6", + "7", "8", "9", + "10", "11", "12"}; + + std::vector perm = {0, 3, 1, 2}; + std::vector expected_shape({1, 3, 2, 2}); + std::initializer_list expected_vals = { + "1", "4", "7", "10", + "2", "5", "8", "11", + "3", "6", "9", "12"}; + + TransposeTest(input_shape, input_vals, &perm, expected_shape, expected_vals, false); +} + } // namespace test } // namespace onnxruntime