| 186 | public ::testing::WithParamInterface<TestCase> {}; |
| 187 | |
| 188 | TEST_P(ParameterizedMapDefunOpTest, NormalTests) { |
| 189 | int thread_num = 2, cpu_num = 2; |
| 190 | TestCase test_case = GetParam(); |
| 191 | TF_ASSERT_OK(InitThreadPool(thread_num)); |
| 192 | TF_ASSERT_OK(InitFunctionLibraryRuntime(test_case.func_lib, cpu_num)); |
| 193 | |
| 194 | std::unique_ptr<OpKernel> map_defun_kernel; |
| 195 | TF_ASSERT_OK(CreateMapDefunOpKernel( |
| 196 | test_case.t_arguments, test_case.t_captured, test_case.output_dtypes, |
| 197 | test_case.output_shapes, test_case.func, |
| 198 | test_case.max_intra_op_parallelism, &map_defun_kernel)); |
| 199 | gtl::InlinedVector<TensorValue, 4> inputs; |
| 200 | for (auto& arg : test_case.arguments) { |
| 201 | inputs.emplace_back(&arg); |
| 202 | } |
| 203 | for (auto& captured_input : test_case.captured_inputs) { |
| 204 | inputs.emplace_back(&captured_input); |
| 205 | } |
| 206 | std::unique_ptr<OpKernelContext> context; |
| 207 | TF_ASSERT_OK( |
| 208 | CreateMapDefunContext(map_defun_kernel.get(), &inputs, &context)); |
| 209 | TF_ASSERT_OK(RunOpKernel(map_defun_kernel.get(), context.get())); |
| 210 | |
| 211 | EXPECT_EQ(context->num_outputs(), test_case.expected_outputs.size()); |
| 212 | for (int i = 0; i < context->num_outputs(); ++i) { |
| 213 | TF_EXPECT_OK(ExpectEqual(*context->mutable_output(i), |
| 214 | test_case.expected_outputs[i])); |
| 215 | } |
| 216 | } |
| 217 | |
| 218 | INSTANTIATE_TEST_SUITE_P(MapDefunOpTest, ParameterizedMapDefunOpTest, |
| 219 | ::testing::ValuesIn(std::vector<TestCase>( |
nothing calls this directly
no test coverage detected