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

Class CheckerboardDataModule

data/checkerboarddataset.py:235–267  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

233 return lonlats, labels
234
235class CheckerboardDataModule(pl.LightningDataModule):
236 def __init__(self, num_samples=5000, batch_size=1000, num_classes = 4, num_support = 200):
237 super().__init__()
238 self.num_samples = num_samples
239 self.batch_size=batch_size
240 self.num_support = num_support
241 self.num_classes = num_classes
242
243 # mean and std distance between clusters given the number of points
244 self.mean_dist, self.std_dist = calc_avg_distances(num_support, unit="deg")
245
246 def setup(self, stage: str):
247 self.train_ds = TensorDataset(*get_data(N_samples = self.num_samples,
248 N_support = self.num_support,
249 n_classes=self.num_classes,
250 seed=0))
251 self.valid_ds = TensorDataset(*get_data(N_samples = self.num_samples,
252 N_support = self.num_support,
253 n_classes=self.num_classes,
254 seed=1))
255 self.evalu_ds = TensorDataset(*get_data(N_samples = self.num_samples,
256 N_support = self.num_support,
257 n_classes=self.num_classes,
258 grid=True))
259
260 def train_dataloader(self):
261 return DataLoader(self.train_ds, batch_size=self.batch_size, shuffle=True)
262
263 def val_dataloader(self):
264 return DataLoader(self.valid_ds, batch_size=self.batch_size, shuffle=False)
265
266 def test_dataloader(self):
267 return DataLoader(self.evalu_ds, batch_size=self.batch_size, shuffle=False)
268
269if __name__ == '__main__':
270 main()

Callers 3

fitFunction · 0.90
tuneFunction · 0.90
fitFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected