OpenMotionLab/MotionGPT
118
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 