MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / TEST_P

Function TEST_P

tensorflow/core/kernels/data/window_dataset_op_test.cc:306–382  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

304 public ::testing::WithParamInterface<TestCase> {};
305
306TEST_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(

Callers

nothing calls this directly

Calls 15

GetParamFunction · 0.85
TensorValueClass · 0.85
VerifyTypesMatchFunction · 0.85
VerifyShapesCompatibleFunction · 0.85
TensorShapeClass · 0.50
ExpectEqualFunction · 0.50
getMethod · 0.45
MakeIteratorMethod · 0.45
beginMethod · 0.45
GetNextMethod · 0.45

Tested by

no test coverage detected