(self, loader_cfg: DataLoaderStageCfg)
| 81 | return None if loader_cfg.num_workers == 0 else loader_cfg.persistent_workers |
| 82 | |
| 83 | def get_generator(self, loader_cfg: DataLoaderStageCfg) -> torch.Generator | None: |
| 84 | if loader_cfg.seed is None: |
| 85 | return None |
| 86 | generator = Generator() |
| 87 | generator.manual_seed(loader_cfg.seed + self.global_rank) |
| 88 | return generator |
| 89 | |
| 90 | def train_dataloader(self): |
| 91 | datasets = get_dataset(self.dataset_cfgs, "train", self.step_tracker) |
no outgoing calls