| 233 | return lonlats, labels |
| 234 | |
| 235 | class 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 | |
| 269 | if __name__ == '__main__': |
| 270 | main() |