MCPcopy Create free account
hub / github.com/MarcCoru/locationencoder / LandOceanDataModule

Class LandOceanDataModule

data/landoceandataset.py:78–96  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

76 return lonlats, land
77
78class LandOceanDataModule(pl.LightningDataModule):
79 def __init__(self, num_samples=5000, batch_size=1000):
80 super().__init__()
81 self.num_samples = num_samples
82 self.batch_size=batch_size
83
84 def setup(self, stage: str):
85 self.train_ds = TensorDataset(*get_data(self.num_samples, seed=0))
86 self.valid_ds = TensorDataset(*get_data(self.num_samples, seed=1))
87 self.evalu_ds = TensorDataset(*get_data(self.num_samples, grid=True))
88
89 def train_dataloader(self):
90 return DataLoader(self.train_ds, batch_size=self.batch_size, shuffle=True)
91
92 def val_dataloader(self):
93 return DataLoader(self.valid_ds, batch_size=self.batch_size, shuffle=False)
94
95 def test_dataloader(self):
96 return DataLoader(self.evalu_ds, batch_size=self.batch_size, shuffle=False)
97
98
99if __name__ == '__main__':

Callers 3

fitFunction · 0.90
tuneFunction · 0.90
fitFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected