| 78 | return split_era5_dataframe(era5_df, label_key=label_key, random_seed=random_seed) |
| 79 | |
| 80 | class ERA5DataModule(pl.LightningDataModule): |
| 81 | def __init__(self, num_workers=0, batch_size=1000,data_root='/home/kklemmer/sphericalharmonics/data/era5', label_key='t2m'): |
| 82 | super().__init__() |
| 83 | self.batch_size=batch_size |
| 84 | self.data_root=data_root |
| 85 | self.num_workers=num_workers |
| 86 | self.label_key = label_key |
| 87 | |
| 88 | |
| 89 | def setup(self, stage: str): |
| 90 | data_by_split = get_era5_data_by_split(data_root=self.data_root, label_key=self.label_key, random_seed=0) |
| 91 | self.train_ds = TensorDataset(*data_by_split['train']) |
| 92 | self.valid_ds = TensorDataset(*data_by_split['val']) |
| 93 | self.evalu_ds = TensorDataset(*data_by_split['test']) |
| 94 | |
| 95 | self.test_locs = data_by_split['test'][0].detach() |
| 96 | |
| 97 | def train_dataloader(self): |
| 98 | return DataLoader(self.train_ds, batch_size=self.batch_size, num_workers=self.num_workers, shuffle=True) |
| 99 | |
| 100 | def val_dataloader(self): |
| 101 | return DataLoader(self.valid_ds, batch_size=self.batch_size, num_workers=self.num_workers, shuffle=False) |
| 102 | |
| 103 | def test_dataloader(self): |
| 104 | return DataLoader(self.evalu_ds, batch_size=self.batch_size, num_workers=self.num_workers, shuffle=False) |
| 105 | |
| 106 | def get_test_locs(self): |
| 107 | return self.test_locs |
| 108 | |
| 109 | # if __name__ == '__main__': |
| 110 | # import matplotlib.pyplot as plt |