blob: c295dff1c9a77457e630c74abbaed7b7cd892188 [file] [edit]
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