| 304 | public ::testing::WithParamInterface<TestCase> {}; |
| 305 | |
| 306 | TEST_P(ParameterizedWindowDatasetOpTest, GetNext) { |
| 307 | int thread_num = 2, cpu_num = 2; |
| 308 | TestCase test_case = GetParam(); |
| 309 | TF_ASSERT_OK(InitThreadPool(thread_num)); |
| 310 | TF_ASSERT_OK(InitFunctionLibraryRuntime({}, cpu_num)); |
| 311 | |
| 312 | std::unique_ptr<OpKernel> window_dataset_kernel; |
| 313 | TF_ASSERT_OK(CreateWindowDatasetKernel(test_case.expected_output_dtypes, |
| 314 | test_case.expected_output_shapes, |
| 315 | &window_dataset_kernel)); |
| 316 | |
| 317 | DatasetBase* range_dataset; |
| 318 | TF_ASSERT_OK(CreateRangeDataset<int64>( |
| 319 | test_case.range_data_param.start, test_case.range_data_param.end, |
| 320 | test_case.range_data_param.step, "range", &range_dataset)); |
| 321 | Tensor range_dataset_tensor(DT_VARIANT, TensorShape({})); |
| 322 | TF_ASSERT_OK( |
| 323 | StoreDatasetInVariantTensor(range_dataset, &range_dataset_tensor)); |
| 324 | Tensor size = test_case.size; |
| 325 | Tensor shift = test_case.shift; |
| 326 | Tensor stride = test_case.stride; |
| 327 | Tensor drop_remainder = test_case.drop_remainder; |
| 328 | gtl::InlinedVector<TensorValue, 4> inputs( |
| 329 | {TensorValue(&range_dataset_tensor), TensorValue(&size), |
| 330 | TensorValue(&shift), TensorValue(&stride), |
| 331 | TensorValue(&drop_remainder)}); |
| 332 | |
| 333 | std::unique_ptr<OpKernelContext> window_dataset_op_ctx; |
| 334 | TF_ASSERT_OK(CreateWindowDatasetContext(window_dataset_kernel.get(), &inputs, |
| 335 | &window_dataset_op_ctx)); |
| 336 | DatasetBase* dataset; |
| 337 | TF_ASSERT_OK(CreateDataset(window_dataset_kernel.get(), |
| 338 | window_dataset_op_ctx.get(), &dataset)); |
| 339 | core::ScopedUnref scoped_unref_dataset(dataset); |
| 340 | |
| 341 | std::unique_ptr<IteratorContext> iterator_ctx; |
| 342 | TF_ASSERT_OK( |
| 343 | CreateIteratorContext(window_dataset_op_ctx.get(), &iterator_ctx)); |
| 344 | std::unique_ptr<IteratorBase> iterator; |
| 345 | TF_ASSERT_OK( |
| 346 | dataset->MakeIterator(iterator_ctx.get(), "Iterator", &iterator)); |
| 347 | |
| 348 | bool end_of_sequence = false; |
| 349 | auto expected_outputs_it = test_case.expected_outputs.begin(); |
| 350 | while (!end_of_sequence) { |
| 351 | // Owns the window_datasets, which are stored as the variant tensors in the |
| 352 | // vector. |
| 353 | std::vector<Tensor> out_tensors; |
| 354 | TF_EXPECT_OK( |
| 355 | iterator->GetNext(iterator_ctx.get(), &out_tensors, &end_of_sequence)); |
| 356 | if (!end_of_sequence) { |
| 357 | for (const auto& window_dataset_tensor : out_tensors) { |
| 358 | // Not owned. |
| 359 | DatasetBase* window_dataset; |
| 360 | TF_ASSERT_OK(GetDatasetFromVariantTensor(window_dataset_tensor, |
| 361 | &window_dataset)); |
| 362 | std::unique_ptr<IteratorBase> window_dataset_iterator; |
| 363 | TF_ASSERT_OK(window_dataset->MakeIterator( |
nothing calls this directly
no test coverage detected