| 40 | |
| 41 | |
| 42 | class _BaseDataModule(LightningDataModule): |
| 43 | def __init__( |
| 44 | self, |
| 45 | data_dir: str, |
| 46 | batch_size: int = 64, |
| 47 | num_workers: int = 4, |
| 48 | pin_memory: bool = True, |
| 49 | size: int = 224, |
| 50 | augment: bool = True, |
| 51 | num_samples: Optional[int] = None, |
| 52 | ): |
| 53 | super().__init__() |
| 54 | |
| 55 | self.augment = augment |
| 56 | self.data_dir = data_dir |
| 57 | self.batch_size = batch_size |
| 58 | self.num_workers = num_workers |
| 59 | self.pin_memory = pin_memory |
| 60 | if isinstance(size, int): |
| 61 | self.size = (size, size) |
| 62 | else: |
| 63 | self.size = size |
| 64 | self.num_samples = num_samples |
| 65 | |
| 66 | def setup(self, stage=None): |
| 67 | pass |
| 68 | |
| 69 | def train_dataloader(self): |
| 70 | return DataLoader( |
| 71 | dataset = self.data_train, |
| 72 | batch_size = self.batch_size, |
| 73 | num_workers = self.num_workers, |
| 74 | pin_memory = self.pin_memory, |
| 75 | shuffle = True, |
| 76 | drop_last = True |
| 77 | ) |
| 78 | |
| 79 | def val_dataloader(self): |
| 80 | return DataLoader( |
| 81 | dataset = self.data_val, |
| 82 | batch_size = self.batch_size, |
| 83 | num_workers = self.num_workers, |
| 84 | pin_memory = self.pin_memory, |
| 85 | shuffle = False, |
| 86 | drop_last = False |
| 87 | ) |
| 88 | |
| 89 | |
| 90 | class _FGVCDataModule(_BaseDataModule): |
nothing calls this directly
no outgoing calls
no test coverage detected