Iterates once over the DataLoader's dataset, as specified by the sampler
| 186 | |
| 187 | |
| 188 | class 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 |