| 6 | |
| 7 | |
| 8 | class DataModule(pl.LightningDataModule): |
| 9 | def __init__( |
| 10 | self, |
| 11 | train_dataset: object, |
| 12 | batch_size: int, |
| 13 | num_workers: int |
| 14 | ): |
| 15 | r"""Data module. To get one batch of data: |
| 16 | |
| 17 | code-block:: python |
| 18 | |
| 19 | data_module.setup() |
| 20 | |
| 21 | for batch_data_dict in data_module.train_dataloader(): |
| 22 | print(batch_data_dict.keys()) |
| 23 | break |
| 24 | |
| 25 | Args: |
| 26 | train_sampler: Sampler object |
| 27 | train_dataset: Dataset object |
| 28 | num_workers: int |
| 29 | distributed: bool |
| 30 | """ |
| 31 | super().__init__() |
| 32 | self._train_dataset = train_dataset |
| 33 | self.num_workers = num_workers |
| 34 | self.batch_size = batch_size |
| 35 | self.collate_fn = collate_fn |
| 36 | |
| 37 | |
| 38 | def prepare_data(self): |
| 39 | # download, split, etc... |
| 40 | # only called on 1 GPU/TPU in distributed |
| 41 | pass |
| 42 | |
| 43 | def setup(self, stage: Optional[str] = None) -> NoReturn: |
| 44 | r"""called on every device.""" |
| 45 | |
| 46 | # make assignments here (val/train/test split) |
| 47 | # called on every process in DDP |
| 48 | |
| 49 | # SegmentSampler is used for selecting segments for training. |
| 50 | # On multiple devices, each SegmentSampler samples a part of mini-batch |
| 51 | # data. |
| 52 | self.train_dataset = self._train_dataset |
| 53 | |
| 54 | |
| 55 | def train_dataloader(self) -> torch.utils.data.DataLoader: |
| 56 | r"""Get train loader.""" |
| 57 | train_loader = DataLoader( |
| 58 | dataset=self.train_dataset, |
| 59 | batch_size=self.batch_size, |
| 60 | collate_fn=self.collate_fn, |
| 61 | num_workers=self.num_workers, |
| 62 | pin_memory=True, |
| 63 | persistent_workers=False, |
| 64 | shuffle=True |
| 65 | ) |