(device, include_decoder, debug=False)
| 89 | @attr("pytorch") |
| 90 | @params(("cpu", False), ("cpu", True), ("gpu", False), ("gpu", True)) |
| 91 | def test_dali_proxy_torch_data_loader(device, include_decoder, debug=False): |
| 92 | # Shows how DALI proxy is used in practice with a PyTorch data loader |
| 93 | |
| 94 | from nvidia.dali.plugin.pytorch.experimental import proxy as dali_proxy |
| 95 | import torchvision.datasets as datasets |
| 96 | from torch.utils import data as torchdata |
| 97 | |
| 98 | batch_size = 4 |
| 99 | num_threads = 3 |
| 100 | device_id = 0 |
| 101 | nworkers = 4 |
| 102 | pipe = image_pipe( |
| 103 | dali_device=device, |
| 104 | include_decoder=include_decoder, |
| 105 | random_pipe=False, |
| 106 | batch_size=batch_size, |
| 107 | num_threads=num_threads, |
| 108 | device_id=device_id, |
| 109 | prefetch_queue_depth=2 + nworkers, |
| 110 | ) |
| 111 | |
| 112 | pipe_ref = image_pipe( |
| 113 | dali_device=device, |
| 114 | include_decoder=include_decoder, |
| 115 | random_pipe=False, |
| 116 | batch_size=batch_size, |
| 117 | num_threads=num_threads, |
| 118 | device_id=device_id, |
| 119 | prefetch_queue_depth=1, |
| 120 | ) |
| 121 | |
| 122 | dali_server = dali_proxy.DALIServer(pipe) |
| 123 | if include_decoder: |
| 124 | dataset = datasets.ImageFolder(jpeg, transform=dali_server.proxy, loader=read_filepath) |
| 125 | dataset_ref = datasets.ImageFolder(jpeg, transform=lambda x: x.copy(), loader=read_filepath) |
| 126 | else: |
| 127 | dataset = datasets.ImageFolder(jpeg, transform=dali_server.proxy) |
| 128 | dataset_ref = datasets.ImageFolder(jpeg, transform=lambda x: x.copy()) |
| 129 | |
| 130 | loader = dali_proxy.DataLoader( |
| 131 | dali_server, |
| 132 | dataset, |
| 133 | batch_size=batch_size, |
| 134 | num_workers=nworkers, |
| 135 | drop_last=True, |
| 136 | ) |
| 137 | |
| 138 | def ref_collate_fn(batch): |
| 139 | filepaths, labels = zip(*batch) # Separate the inputs and labels |
| 140 | # Just return the batch as they are, a list of individual tensors |
| 141 | return filepaths, labels |
| 142 | |
| 143 | loader_ref = torchdata.dataloader.DataLoader( |
| 144 | dataset_ref, |
| 145 | batch_size=batch_size, |
| 146 | num_workers=1, |
| 147 | collate_fn=ref_collate_fn, |
| 148 | shuffle=False, |
nothing calls this directly
no test coverage detected