(case_name, dtype, g)
| 532 | |
| 533 | |
| 534 | def _gpu_permuted_extents_torch_case(case_name, dtype, g): |
| 535 | shapes_perms = get_permute_extents_case(case_name) |
| 536 | shapes, perms = tuple(zip(*shapes_perms)) |
| 537 | assert len(shapes) == len(perms) == len(shapes_perms) |
| 538 | input_batch = get_random_torch_batch(g, shapes, dtype) |
| 539 | assert len(input_batch) == len(shapes) |
| 540 | |
| 541 | # returns permuted view of the input tensors |
| 542 | def permuted_tensors(batch): |
| 543 | stream = current_dali_stream() |
| 544 | torch_stream = torch.cuda.ExternalStream(stream) |
| 545 | with torch.cuda.stream(torch_stream): |
| 546 | tensors = [torch_dlpack.from_dlpack(t) for t in batch] |
| 547 | assert len(tensors) == len(perms) |
| 548 | tensor_views = [t.permute(perm) for t, perm in zip(tensors, perms)] |
| 549 | out = [torch_dlpack.to_dlpack(t) for t in tensor_views] |
| 550 | return out |
| 551 | |
| 552 | @pipeline_def(batch_size=len(input_batch), num_threads=4, device_id=0) |
| 553 | def pipeline(): |
| 554 | data = fn.external_source(lambda: input_batch) |
| 555 | data = fn.dl_tensor_python_function( |
| 556 | data.gpu(), batch_processing=True, function=permuted_tensors, synchronize_stream=False |
| 557 | ) |
| 558 | return data |
| 559 | |
| 560 | p = pipeline() |
| 561 | (out,) = p.run() |
| 562 | |
| 563 | out = [numpy.array(sample) for sample in out.as_cpu()] |
| 564 | ref = [numpy.array(sample).transpose(perm) for sample, perm in zip(input_batch, perms)] |
| 565 | |
| 566 | numpy.testing.assert_equal(out, ref) |
| 567 | |
| 568 | |
| 569 | def _gpu_permuted_extents_torch_suite(): |
nothing calls this directly
no test coverage detected