| 37 | |
| 38 | |
| 39 | class DataModulePretrain(LightningDataModule): |
| 40 | def __init__( |
| 41 | self, |
| 42 | dataset_cfgs: list[DatasetCfg], |
| 43 | data_module_cfg: DataModulePretrainCfg, |
| 44 | frame_sampler_cfg: FrameSamplerCfg, |
| 45 | global_rank: int, |
| 46 | ) -> None: |
| 47 | super().__init__() |
| 48 | self.dataset_cfgs = dataset_cfgs |
| 49 | self.data_module_cfg = data_module_cfg |
| 50 | self.frame_sampler_cfg = frame_sampler_cfg |
| 51 | self.global_rank = global_rank |
| 52 | |
| 53 | def get_persistent(self, loader_cfg: DataLoaderStageCfg) -> bool | None: |
| 54 | return None if loader_cfg.num_workers == 0 else loader_cfg.persistent_workers |
| 55 | |
| 56 | def get_generator(self, loader_cfg: DataLoaderStageCfg) -> torch.Generator | None: |
| 57 | if loader_cfg.seed is None: |
| 58 | return None |
| 59 | generator = Generator() |
| 60 | generator.manual_seed(loader_cfg.seed + self.global_rank) |
| 61 | return generator |
| 62 | |
| 63 | def train_dataloader(self): |
| 64 | dataset = get_dataset(self.dataset_cfgs, "train", self.frame_sampler_cfg) |
| 65 | return DataLoader( |
| 66 | dataset, |
| 67 | self.data_module_cfg.train.batch_size, |
| 68 | shuffle=not isinstance(dataset, IterableDataset), |
| 69 | num_workers=self.data_module_cfg.train.num_workers, |
| 70 | generator=self.get_generator(self.data_module_cfg.train), |
| 71 | worker_init_fn=worker_init_fn, |
| 72 | persistent_workers=self.get_persistent(self.data_module_cfg.train), |
| 73 | ) |
| 74 | |
| 75 | def val_dataloader(self): |
| 76 | dataset = get_dataset(self.dataset_cfgs, "val", self.frame_sampler_cfg) |
| 77 | return DataLoader( |
| 78 | ValidationWrapper(dataset, 1), |
| 79 | self.data_module_cfg.val.batch_size, |
| 80 | num_workers=self.data_module_cfg.val.num_workers, |
| 81 | generator=self.get_generator(self.data_module_cfg.val), |
| 82 | worker_init_fn=worker_init_fn, |
| 83 | persistent_workers=self.get_persistent(self.data_module_cfg.val), |
| 84 | ) |