MCPcopy Create free account
hub / github.com/NVIDIA/DALI / test_dali_proxy_torch_data_loader

Function test_dali_proxy_torch_data_loader

dali/test/python/test_dali_proxy.py:91–158  ·  view source on GitHub ↗
(device, include_decoder, debug=False)

Source from the content-addressed store, hash-verified

89@attr("pytorch")
90@params(("cpu", False), ("cpu", True), ("gpu", False), ("gpu", True))
91def 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,

Callers

nothing calls this directly

Calls 6

stop_threadMethod · 0.95
image_pipeFunction · 0.70
copyMethod · 0.45
feed_inputMethod · 0.45
runMethod · 0.45
cpuMethod · 0.45

Tested by

no test coverage detected