Team Ai
Apppublic

OpenMotionLab/MotionGPT

sourceHugging Facemitupdated 1y agoView on Hugging Face
118likes
__init__.py104 linesDownload Raw Back to data
1import pytorch_lightning as pl2from torch.utils.data import DataLoader3 4 5class BASEDataModule(pl.LightningDataModule):6    def __init__(self, collate_fn):7        super().__init__()8 9        self.dataloader_options = {"collate_fn": collate_fn}10        self.persistent_workers = True11        self.is_mm = False12 13        self._train_dataset = None14        self._val_dataset = None15        self._test_dataset = None16 17    def get_sample_set(self, overrides={}):18        sample_params = self.hparams.copy()19        sample_params.update(overrides)20        return self.DatasetEval(**sample_params)21 22    @property23    def train_dataset(self):24        if self._train_dataset is None:25            self._train_dataset = self.Dataset(split=self.cfg.TRAIN.SPLIT,26                                               **self.hparams)27        return self._train_dataset28 29    @property30    def val_dataset(self):31        if self._val_dataset is None:32            params = self.hparams.copy()33            params['code_path'] = None34            params['split'] = self.cfg.EVAL.SPLIT35            self._val_dataset = self.DatasetEval(**params)36        return self._val_dataset37 38    @property39    def test_dataset(self):40        if self._test_dataset is None:41            # self._test_dataset = self.DatasetEval(split=self.cfg.TEST.SPLIT,42            #                                       **self.hparams)43            params = self.hparams.copy()44            params['code_path'] = None45            params['split'] = self.cfg.TEST.SPLIT46            self._test_dataset = self.DatasetEval( **params)47        return self._test_dataset48 49    def setup(self, stage=None):50        # Use the getter the first time to load the data51        if stage in (None, "fit"):52            _ = self.train_dataset53            _ = self.val_dataset54        if stage in (None, "test"):55            _ = self.test_dataset56 57    def train_dataloader(self):58        dataloader_options = self.dataloader_options.copy()59        dataloader_options["batch_size"] = self.cfg.TRAIN.BATCH_SIZE60        dataloader_options["num_workers"] = self.cfg.TRAIN.NUM_WORKERS61        return DataLoader(62            self.train_dataset,63            shuffle=False,64            persistent_workers=True,65            **dataloader_options,66        )67 68    def predict_dataloader(self):69        dataloader_options = self.dataloader_options.copy()70        dataloader_options[71            "batch_size"] = 1 if self.is_mm else self.cfg.TEST.BATCH_SIZE72        dataloader_options["num_workers"] = self.cfg.TEST.NUM_WORKERS73        dataloader_options["shuffle"] = False74        return DataLoader(75            self.test_dataset,76            persistent_workers=True,77            **dataloader_options,78        )79 80    def val_dataloader(self):81        # overrides batch_size and num_workers82        dataloader_options = self.dataloader_options.copy()83        dataloader_options["batch_size"] = self.cfg.EVAL.BATCH_SIZE84        dataloader_options["num_workers"] = self.cfg.EVAL.NUM_WORKERS85        dataloader_options["shuffle"] = False86        return DataLoader(87            self.val_dataset,88            persistent_workers=True,89            **dataloader_options,90        )91 92    def test_dataloader(self):93        # overrides batch_size and num_workers94        dataloader_options = self.dataloader_options.copy()95        dataloader_options[96            "batch_size"] = 1 if self.is_mm else self.cfg.TEST.BATCH_SIZE97        dataloader_options["num_workers"] = self.cfg.TEST.NUM_WORKERS98        dataloader_options["shuffle"] = False99        return DataLoader(100            self.test_dataset,101            persistent_workers=True,102            **dataloader_options,103        )104