MCPcopy Create free account
hub / github.com/FreeformRobotics/OTS / DataLoaderIter

Class DataLoaderIter

lib/utils/data/dataloader.py:188–338  ·  view source on GitHub ↗

Iterates once over the DataLoader's dataset, as specified by the sampler

Source from the content-addressed store, hash-verified

186
187
188class DataLoaderIter(object):
189 "Iterates once over the DataLoader's dataset, as specified by the sampler"
190
191 def __init__(self, loader):
192 self.dataset = loader.dataset
193 self.collate_fn = loader.collate_fn
194 self.batch_sampler = loader.batch_sampler
195 self.num_workers = loader.num_workers
196 self.pin_memory = loader.pin_memory and torch.cuda.is_available()
197 self.timeout = loader.timeout
198 self.done_event = threading.Event()
199
200 self.sample_iter = iter(self.batch_sampler)
201
202 if self.num_workers > 0:
203 self.worker_init_fn = loader.worker_init_fn
204 self.index_queue = multiprocessing.SimpleQueue()
205 self.worker_result_queue = multiprocessing.SimpleQueue()
206 self.batches_outstanding = 0
207 self.worker_pids_set = False
208 self.shutdown = False
209 self.send_idx = 0
210 self.rcvd_idx = 0
211 self.reorder_dict = {}
212
213 base_seed = torch.LongTensor(1).random_(0, 2**31-1)[0]
214 self.workers = [
215 multiprocessing.Process(
216 target=_worker_loop,
217 args=(self.dataset, self.index_queue, self.worker_result_queue, self.collate_fn,
218 base_seed + i, self.worker_init_fn, i))
219 for i in range(self.num_workers)]
220
221 if self.pin_memory or self.timeout > 0:
222 self.data_queue = queue.Queue()
223 if self.pin_memory:
224 maybe_device_id = torch.cuda.current_device()
225 else:
226 # do not initialize cuda context if not necessary
227 maybe_device_id = None
228 self.worker_manager_thread = threading.Thread(
229 target=_worker_manager_loop,
230 args=(self.worker_result_queue, self.data_queue, self.done_event, self.pin_memory,
231 maybe_device_id))
232 self.worker_manager_thread.daemon = True
233 self.worker_manager_thread.start()
234 else:
235 self.data_queue = self.worker_result_queue
236
237 for w in self.workers:
238 w.daemon = True # ensure that the worker exits on process exit
239 w.start()
240
241 _set_worker_pids(id(self), tuple(w.pid for w in self.workers))
242 _set_SIGCHLD_handler()
243 self.worker_pids_set = True
244
245 # prime the prefetch loop

Callers 1

__iter__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected