(self)
| 88 | return generator |
| 89 | |
| 90 | def train_dataloader(self): |
| 91 | datasets = get_dataset(self.dataset_cfgs, "train", self.step_tracker) |
| 92 | data_loaders = [] |
| 93 | for dataset in datasets: |
| 94 | dataset = self.dataset_shim(dataset, "train") |
| 95 | data_loaders.append( |
| 96 | DataLoader( |
| 97 | dataset, |
| 98 | self.data_loader_cfg.train.batch_size, |
| 99 | shuffle=not isinstance(dataset, IterableDataset), |
| 100 | num_workers=self.data_loader_cfg.train.num_workers, |
| 101 | generator=self.get_generator(self.data_loader_cfg.train), |
| 102 | worker_init_fn=worker_init_fn, |
| 103 | persistent_workers=self.get_persistent(self.data_loader_cfg.train), |
| 104 | ) |
| 105 | ) |
| 106 | return data_loaders if len(data_loaders) > 1 else data_loaders[0] |
| 107 | |
| 108 | def val_dataloader(self): |
| 109 | datasets = get_dataset(self.dataset_cfgs, "val", self.step_tracker) |
nothing calls this directly
no test coverage detected