| 630 | |
| 631 | |
| 632 | class MultiObjectDataModule(pl.LightningDataModule): |
| 633 | cfg: MultiObjectDataModuleConfig |
| 634 | |
| 635 | def __init__(self, cfg: Optional[Union[dict, DictConfig]] = None) -> None: |
| 636 | super().__init__() |
| 637 | self.cfg = parse_structured(MultiObjectDataModuleConfig, cfg) |
| 638 | |
| 639 | def setup(self, stage=None) -> None: |
| 640 | if stage in [None, "fit"]: |
| 641 | self.train_dataset = MultiObjectDataset(self.cfg, "train") |
| 642 | if stage in [None, "fit", "validate"]: |
| 643 | self.val_dataset = MultiObjectDataset(self.cfg, "val") |
| 644 | if stage in [None, "test", "predict"]: |
| 645 | self.test_dataset = MultiObjectDataset(self.cfg, "test") |
| 646 | |
| 647 | def prepare_data(self): |
| 648 | pass |
| 649 | |
| 650 | def train_dataloader(self) -> DataLoader: |
| 651 | return DataLoader( |
| 652 | self.train_dataset, |
| 653 | batch_size=self.cfg.batch_size, |
| 654 | num_workers=self.cfg.num_workers, |
| 655 | shuffle=True, |
| 656 | collate_fn=self.train_dataset.collate, |
| 657 | ) |
| 658 | |
| 659 | def val_dataloader(self) -> DataLoader: |
| 660 | return DataLoader( |
| 661 | self.val_dataset, |
| 662 | batch_size=self.cfg.eval_batch_size, |
| 663 | num_workers=self.cfg.num_workers, |
| 664 | shuffle=False, |
| 665 | collate_fn=self.val_dataset.collate, |
| 666 | ) |
| 667 | |
| 668 | def test_dataloader(self) -> DataLoader: |
| 669 | return DataLoader( |
| 670 | self.test_dataset, |
| 671 | batch_size=self.cfg.eval_batch_size, |
| 672 | num_workers=self.cfg.num_workers, |
| 673 | shuffle=False, |
| 674 | collate_fn=self.test_dataset.collate, |
| 675 | ) |
| 676 | |
| 677 | def predict_dataloader(self) -> DataLoader: |
| 678 | return self.test_dataloader() |
| 679 | |
| 680 | |
| 681 | if __name__ == "__main__": |