| From 33f12e101e41507c8f8c9fb8afb0000f88452c24 Mon Sep 17 00:00:00 2001 |
| From: Tommy Chiang <ototot@google.com> |
| Date: Wed, 9 Oct 2024 07:34:39 -0700 |
| Subject: [PATCH] Use ArrayFloatNear to check values in StridedSliceOpTests |
| |
| Since StridedSliceOpTests use TYPED_TEST to test StridedSliceOp with |
| different types, this CL introduce a helper function |
| ElementsAreTypedArray to help testing float values. |
| |
| PiperOrigin-RevId: 684031126 |
| |
| PATCH_NAME=dts-stridedsliceoptests |
| --- |
| tensorflow/lite/kernels/strided_slice_test.cc | 155 ++++++++++-------- |
| 1 file changed, 89 insertions(+), 66 deletions(-) |
| |
| diff --git a/tensorflow/lite/kernels/strided_slice_test.cc b/tensorflow/lite/kernels/strided_slice_test.cc |
| index 9e63abeb..e769831c 100644 |
| --- a/tensorflow/lite/kernels/strided_slice_test.cc |
| +++ b/tensorflow/lite/kernels/strided_slice_test.cc |
| @@ -156,6 +156,15 @@ using DataTypes = |
| ::testing::Types<float, uint8_t, uint32_t, int8_t, int16_t, int32_t>; |
| TYPED_TEST_SUITE(StridedSliceOpTest, DataTypes); |
| |
| +template <typename TypeParam, typename T = TypeParam> |
| +auto ElementsAreTypedArray(std::vector<T> x) { |
| + if constexpr (std::is_floating_point_v<TypeParam>) { |
| + return ElementsAreArray(ArrayFloatNear(std::move(x))); |
| + } else { |
| + return ElementsAreArray(std::move(x)); |
| + } |
| +} |
| + |
| #if GTEST_HAS_DEATH_TEST |
| TYPED_TEST(StridedSliceOpTest, UnsupportedInputSize) { |
| EXPECT_DEATH(StridedSliceOpModel<TypeParam>({2, 2, 2, 2, 2, 2}, {5}, {5}, {5}, |
| @@ -191,7 +200,7 @@ TYPED_TEST(StridedSliceOpTest, Offset) { |
| 0, 0, 0, 0, constant_tensors, /*offset=*/true); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2, 3})); |
| if (constant_tensors) { |
| EXPECT_THAT(m.GetOutputTensor(0)->allocation_type, kTfLitePersistentRo); |
| } else { |
| @@ -212,7 +221,7 @@ TYPED_TEST(StridedSliceOpTest, OffsetArray) { |
| {2, 2}, {1, 1}, 0, 0, 0, 0, 0, constant_tensors, /*offset=*/true); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 5, 6})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2, 5, 6})); |
| if (constant_tensors) { |
| EXPECT_THAT(m.GetOutputTensor(0)->allocation_type, kTfLitePersistentRo); |
| } else { |
| @@ -229,7 +238,7 @@ TYPED_TEST(StridedSliceOpTest, OffsetConstant) { |
| /*offset=*/true); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 5, 6})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2, 5, 6})); |
| EXPECT_THAT(m.GetOutputTensor(0)->allocation_type, kTfLiteArenaRw); |
| } |
| |
| @@ -245,7 +254,7 @@ TYPED_TEST(StridedSliceOpTest, OffsetConstantStride) { |
| /*offset=*/true); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 3, 13, 15})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 3, 13, 15})); |
| EXPECT_THAT(m.GetOutputTensor(0)->allocation_type, kTfLiteArenaRw); |
| } |
| |
| @@ -261,7 +270,8 @@ TYPED_TEST(StridedSliceOpTest, OffsetConstantNegativeStride) { |
| /*offset=*/true); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({28, 26, 16, 14})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({28, 26, 16, 14})); |
| EXPECT_THAT(m.GetOutputTensor(0)->allocation_type, kTfLiteArenaRw); |
| } |
| |
| @@ -275,7 +285,7 @@ TYPED_TEST(StridedSliceOpTest, In1D) { |
| {1}, 0, 0, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({2, 3})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({2, 3})); |
| } |
| } |
| |
| @@ -289,7 +299,7 @@ TYPED_TEST(StridedSliceOpTest, In1DConst) { |
| {1}, 0, 0, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({2, 3})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({2, 3})); |
| } |
| } |
| |
| @@ -308,7 +318,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_Int32End) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({32768})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray(values)); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>(values)); |
| } |
| } |
| |
| @@ -335,7 +345,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_NegativeBegin) { |
| {3}, {1}, 0, 0, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({2, 3})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({2, 3})); |
| } |
| } |
| |
| @@ -349,7 +359,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_OutOfRangeBegin) { |
| {3}, {1}, 0, 0, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2, 3})); |
| } |
| } |
| |
| @@ -364,7 +374,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_NegativeEnd) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({2})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({2})); |
| } |
| } |
| |
| @@ -378,7 +388,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_OutOfRangeEnd) { |
| {5}, {1}, 0, 0, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({2, 3, 4})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({2, 3, 4})); |
| } |
| } |
| |
| @@ -392,7 +402,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_BeginMask) { |
| {1}, 1, 0, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2, 3})); |
| } |
| } |
| |
| @@ -408,7 +418,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_NegativeBeginNegativeStride) { |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({3})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({3})); |
| } |
| } |
| |
| @@ -422,7 +432,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_OutOfRangeBeginNegativeStride) { |
| {-1}, 0, 0, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({4})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({4})); |
| } |
| } |
| |
| @@ -437,7 +447,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_NegativeEndNegativeStride) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({3, 2})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({3, 2})); |
| } |
| } |
| |
| @@ -452,7 +462,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_OutOfRangeEndNegativeStride) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({2, 1})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({2, 1})); |
| } |
| } |
| |
| @@ -466,7 +476,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_EndMask) { |
| {1}, 0, 1, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({2, 3, 4})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({2, 3, 4})); |
| } |
| } |
| |
| @@ -480,7 +490,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_NegStride) { |
| {-1}, 0, 0, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({3, 2, 1})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({3, 2, 1})); |
| } |
| } |
| |
| @@ -494,7 +504,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_EvenLenStride2) { |
| 0, 0, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1})); |
| } |
| } |
| |
| @@ -508,7 +518,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_OddLenStride2) { |
| {2}, 0, 0, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 3})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 3})); |
| } |
| } |
| |
| @@ -523,7 +533,8 @@ TYPED_TEST(StridedSliceOpTest, In2D_Identity) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2, 3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3, 4, 5, 6})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({1, 2, 3, 4, 5, 6})); |
| } |
| } |
| |
| @@ -538,7 +549,7 @@ TYPED_TEST(StridedSliceOpTest, In2D) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({4, 5})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({4, 5})); |
| } |
| } |
| |
| @@ -553,7 +564,7 @@ TYPED_TEST(StridedSliceOpTest, In2D_Stride2) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 3})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 3})); |
| } |
| } |
| |
| @@ -568,7 +579,7 @@ TYPED_TEST(StridedSliceOpTest, In2D_NegStride) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({6, 5, 4})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({6, 5, 4})); |
| } |
| } |
| |
| @@ -583,7 +594,7 @@ TYPED_TEST(StridedSliceOpTest, In2D_BeginMask) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 4, 5})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2, 4, 5})); |
| } |
| } |
| |
| @@ -598,7 +609,7 @@ TYPED_TEST(StridedSliceOpTest, In2D_EndMask) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({4, 5, 6})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({4, 5, 6})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In2D_NegStrideBeginMask) { |
| @@ -612,7 +623,7 @@ TYPED_TEST(StridedSliceOpTest, In2D_NegStrideBeginMask) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({6, 5, 4})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({6, 5, 4})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In2D_NegStrideEndMask) { |
| @@ -626,7 +637,7 @@ TYPED_TEST(StridedSliceOpTest, In2D_NegStrideEndMask) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({5, 4})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({5, 4})); |
| } |
| } |
| |
| @@ -672,7 +683,7 @@ TYPED_TEST(StridedSliceOpTest, In3D_Strided2) { |
| {0, 0, 0}, {2, 3, 2}, {2, 2, 2}, 0, 0, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 2, 1})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 5})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 5})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In1D_ShrinkAxisMask1) { |
| @@ -685,7 +696,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_ShrinkAxisMask1) { |
| {1}, 0, 0, 0, 0, 1, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_TRUE(m.GetOutputShape().empty()); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({2})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({2})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In1D_ShrinkAxisMask1_NegativeSlice) { |
| @@ -700,7 +711,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_ShrinkAxisMask1_NegativeSlice) { |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_TRUE(m.GetOutputShape().empty()); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({3})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({3})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In2D_ShrinkAxis3_NegativeSlice) { |
| @@ -716,7 +727,7 @@ TYPED_TEST(StridedSliceOpTest, In2D_ShrinkAxis3_NegativeSlice) { |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_TRUE(m.GetOutputShape().empty()); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({2})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({2})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In2D_ShrinkAxis2_BeginEndAxis1_NegativeSlice) { |
| @@ -732,7 +743,7 @@ TYPED_TEST(StridedSliceOpTest, In2D_ShrinkAxis2_BeginEndAxis1_NegativeSlice) { |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({4})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({0, 1, 2, 3})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({0, 1, 2, 3})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In1D_BeginMaskShrinkAxisMask1) { |
| @@ -745,7 +756,7 @@ TYPED_TEST(StridedSliceOpTest, In1D_BeginMaskShrinkAxisMask1) { |
| {1}, 1, 0, 0, 0, 1, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_TRUE(m.GetOutputShape().empty()); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In2D_ShrinkAxisMask1) { |
| @@ -759,7 +770,7 @@ TYPED_TEST(StridedSliceOpTest, In2D_ShrinkAxisMask1) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2, 3})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In2D_ShrinkAxisMask2) { |
| @@ -773,7 +784,7 @@ TYPED_TEST(StridedSliceOpTest, In2D_ShrinkAxisMask2) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 4})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 4})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In2D_ShrinkAxisMask3) { |
| @@ -787,7 +798,7 @@ TYPED_TEST(StridedSliceOpTest, In2D_ShrinkAxisMask3) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_TRUE(m.GetOutputShape().empty()); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis1) { |
| @@ -801,7 +812,8 @@ TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis1) { |
| {0, 0, 0}, {1, 3, 2}, {1, 1, 1}, 0, 0, 0, 0, 1, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({3, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3, 4, 5, 6})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({1, 2, 3, 4, 5, 6})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis2) { |
| @@ -815,7 +827,7 @@ TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis2) { |
| {0, 0, 0}, {2, 1, 2}, {1, 1, 1}, 0, 0, 0, 0, 2, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 7, 8})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2, 7, 8})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis3) { |
| @@ -829,7 +841,7 @@ TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis3) { |
| {0, 0, 0}, {1, 1, 2}, {1, 1, 1}, 0, 0, 0, 0, 3, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis4) { |
| @@ -843,7 +855,8 @@ TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis4) { |
| {0, 0, 0}, {2, 3, 1}, {1, 1, 1}, 0, 0, 0, 0, 4, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2, 3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 3, 5, 7, 9, 11})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({1, 3, 5, 7, 9, 11})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis5) { |
| @@ -857,7 +870,7 @@ TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis5) { |
| {0, 0, 0}, {1, 3, 1}, {1, 1, 1}, 0, 0, 0, 0, 5, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 3, 5})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 3, 5})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis6) { |
| @@ -871,7 +884,7 @@ TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis6) { |
| {0, 0, 0}, {2, 1, 1}, {1, 1, 1}, 0, 0, 0, 0, 6, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 7})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 7})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis7) { |
| @@ -881,7 +894,7 @@ TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis7) { |
| {0, 0, 0}, {1, 1, 1}, {1, 1, 1}, 0, 0, 0, 0, 7, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_TRUE(m.GetOutputShape().empty()); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1})); |
| } |
| |
| // This tests catches a very subtle bug that was fixed by cl/188403234. |
| @@ -892,7 +905,7 @@ TYPED_TEST(StridedSliceOpTest, RunTwice) { |
| false); |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 4, 5})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2, 4, 5})); |
| |
| auto setup_inputs = [&m]() { |
| m.template SetInput<TypeParam>({1, 2, 3, 4, 5, 6}, |
| @@ -905,7 +918,7 @@ TYPED_TEST(StridedSliceOpTest, RunTwice) { |
| setup_inputs(); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| // Prior to cl/188403234 this was {4, 5}. |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 4, 5})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2, 4, 5})); |
| } |
| TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis1Uint8) { |
| for (bool constant_tensors : {true, false}) { |
| @@ -918,7 +931,8 @@ TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis1Uint8) { |
| {0, 0, 0}, {1, 3, 2}, {1, 1, 1}, 0, 0, 0, 0, 1, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({3, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3, 4, 5, 6})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({1, 2, 3, 4, 5, 6})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis1int8) { |
| @@ -932,7 +946,8 @@ TYPED_TEST(StridedSliceOpTest, In3D_IdentityShrinkAxis1int8) { |
| {0, 0, 0}, {1, 3, 2}, {1, 1, 1}, 0, 0, 0, 0, 1, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({3, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3, 4, 5, 6})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({1, 2, 3, 4, 5, 6})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In5D_Identity) { |
| @@ -948,7 +963,8 @@ TYPED_TEST(StridedSliceOpTest, In5D_Identity) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2, 1, 2, 1, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3, 4, 9, 10, 11, 12})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({1, 2, 3, 4, 9, 10, 11, 12})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In5D_IdentityShrinkAxis1) { |
| @@ -964,7 +980,7 @@ TYPED_TEST(StridedSliceOpTest, In5D_IdentityShrinkAxis1) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 2, 1, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3, 4})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2, 3, 4})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In3D_SmallBegin) { |
| @@ -978,7 +994,8 @@ TYPED_TEST(StridedSliceOpTest, In3D_SmallBegin) { |
| {1}, {1}, 0, 0, 0, 0, 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 3, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3, 4, 5, 6})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({1, 2, 3, 4, 5, 6})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In3D_SmallBeginWithhrinkAxis1) { |
| @@ -992,7 +1009,8 @@ TYPED_TEST(StridedSliceOpTest, In3D_SmallBeginWithhrinkAxis1) { |
| {1}, {1}, 0, 0, 0, 0, 1, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({3, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3, 4, 5, 6})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({1, 2, 3, 4, 5, 6})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, In3D_BackwardSmallBeginEndMask) { |
| @@ -1082,7 +1100,7 @@ TYPED_TEST(StridedSliceOpTest, In2D_ShrinkAxis_Endmask_AtSameAxis) { |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, EllipsisMask1_NewAxisMask2) { |
| @@ -1097,7 +1115,8 @@ TYPED_TEST(StridedSliceOpTest, EllipsisMask1_NewAxisMask2) { |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2, 3, 1, 1})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 3, 5, 7, 9, 11})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({1, 3, 5, 7, 9, 11})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, EllipsisMask2_NewAxisMask1) { |
| @@ -1112,7 +1131,8 @@ TYPED_TEST(StridedSliceOpTest, EllipsisMask2_NewAxisMask1) { |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 2, 3, 1})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 3, 5, 7, 9, 11})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({1, 3, 5, 7, 9, 11})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, EllipsisMask2_NewAxisMask5) { |
| @@ -1143,7 +1163,7 @@ TYPED_TEST(StridedSliceOpTest, EllipsisMask2_NewAxisMask2) { |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 3, 1})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 3, 5})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 3, 5})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, EllipsisMask4_NewAxisMask2) { |
| @@ -1158,7 +1178,8 @@ TYPED_TEST(StridedSliceOpTest, EllipsisMask4_NewAxisMask2) { |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 1, 3, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 3, 4, 5, 6})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({1, 2, 3, 4, 5, 6})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, EllipsisMask2) { |
| @@ -1173,7 +1194,7 @@ TYPED_TEST(StridedSliceOpTest, EllipsisMask2) { |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 3, 1})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 3, 5})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 3, 5})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, NewAxisMask2) { |
| @@ -1188,7 +1209,7 @@ TYPED_TEST(StridedSliceOpTest, NewAxisMask2) { |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 1, 1, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, NewAxisMask1) { |
| @@ -1203,7 +1224,7 @@ TYPED_TEST(StridedSliceOpTest, NewAxisMask1) { |
| |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1, 2, 1, 2})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1, 2, 7, 8})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1, 2, 7, 8})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, NoInfiniteLoop) { |
| @@ -1229,7 +1250,7 @@ TYPED_TEST(StridedSliceOpTest, MinusThreeMinusFourMinusOne) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({2})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({2})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, MinusFourMinusThreeOne) { |
| @@ -1243,7 +1264,7 @@ TYPED_TEST(StridedSliceOpTest, MinusFourMinusThreeOne) { |
| constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({1})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({1})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({1})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, OneOneOne) { |
| @@ -1268,7 +1289,7 @@ TYPED_TEST(StridedSliceOpTest, OneOneOneShrinkAxis) { |
| {1}, 0, 0, 0, 0, 1, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), IsEmpty()); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({2})); |
| + EXPECT_THAT(m.GetOutput(), ElementsAreTypedArray<TypeParam>({2})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, OneOneOneShrinkAxisOOB) { |
| @@ -1318,7 +1339,8 @@ TYPED_TEST(StridedSliceOpTest, NegEndMask) { |
| 0, constant_tensors); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2, 3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({3, 2, 1, 6, 5, 4})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({3, 2, 1, 6, 5, 4})); |
| } |
| } |
| TYPED_TEST(StridedSliceOpTest, NoopOffset) { |
| @@ -1326,7 +1348,8 @@ TYPED_TEST(StridedSliceOpTest, NoopOffset) { |
| {0, -1}, {2, -3}, {1, -1}, 0, 0b10, 0, 0, 0); |
| ASSERT_EQ(m.Invoke(), kTfLiteOk); |
| EXPECT_THAT(m.GetOutputShape(), ElementsAreArray({2, 3})); |
| - EXPECT_THAT(m.GetOutput(), ElementsAreArray({3, 2, 1, 6, 5, 4})); |
| + EXPECT_THAT(m.GetOutput(), |
| + ElementsAreTypedArray<TypeParam>({3, 2, 1, 6, 5, 4})); |
| } |
| } // namespace |
| } // namespace tflite |