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

Class ERA5DataModule

data/era5dataset.py:80–107  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

78 return split_era5_dataframe(era5_df, label_key=label_key, random_seed=random_seed)
79
80class 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

Callers 2

fitFunction · 0.90
fitFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected