| 11 | |
| 12 | |
| 13 | class AstroClipDataloader(L.LightningDataModule): |
| 14 | def __init__( |
| 15 | self, |
| 16 | path: str, |
| 17 | columns: List[str] = ["image", "spectrum"], |
| 18 | batch_size: int = 512, |
| 19 | num_workers: int = 10, |
| 20 | collate_fn: Callable[[Dict[str, Tensor]], Dict[str, Tensor]] = None, |
| 21 | ) -> None: |
| 22 | super().__init__() |
| 23 | self.save_hyperparameters() |
| 24 | |
| 25 | def setup(self, stage: str) -> None: |
| 26 | self.dataset = datasets.load_from_disk(self.hparams.path) |
| 27 | self.dataset.set_format(type="torch", columns=self.hparams.columns) |
| 28 | |
| 29 | def train_dataloader(self): |
| 30 | return torch.utils.data.DataLoader( |
| 31 | self.dataset["train"], |
| 32 | batch_size=self.hparams.batch_size, |
| 33 | shuffle=True, |
| 34 | num_workers=self.hparams.num_workers, # NOTE: disable for debugging |
| 35 | drop_last=True, |
| 36 | collate_fn=self.hparams.collate_fn, |
| 37 | ) |
| 38 | |
| 39 | def val_dataloader(self): |
| 40 | return torch.utils.data.DataLoader( |
| 41 | self.dataset["test"], |
| 42 | batch_size=self.hparams.batch_size, |
| 43 | num_workers=self.hparams.num_workers, # NOTE: disable for debugging |
| 44 | drop_last=True, |
| 45 | collate_fn=self.hparams.collate_fn, |
| 46 | ) |
| 47 | |
| 48 | |
| 49 | class AstroClipCollator: |
no outgoing calls
no test coverage detected