MCPcopy Create free account
hub / github.com/catalys1/mae-pytorch / _BaseDataModule

Class _BaseDataModule

datamodule.py:42–87  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

40
41
42class _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
90class _FGVCDataModule(_BaseDataModule):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected