File size: 1,618 Bytes
aa8cca6 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 | import torch
import os
import sys
import numpy as np
from tqdm import tqdm
from onescience.datapipes.climate import ERA5Datapipe
from onescience.utils.YParams import YParams
def main():
# instantiate the training datapipe
config_file_path = os.path.join(current_path, 'conf/config.yaml')
cfg_data = YParams(config_file_path, "datapipe")
datapipe = ERA5Datapipe(
dataset_dir=cfg_data.dataset.data_dir,
used_variables=cfg_data.dataset.channels,
used_years=cfg_data.dataset.train_time,
distributed=False
)
train_dataloader, train_sampler = datapipe.get_dataloader("train")
print(f"Loaded training datapipe of length {len(train_dataloader)}")
area = torch.abs(torch.cos(torch.linspace(-90, 90, steps=cfg_data.dataset.img_size[0]) * np.pi / 180))
area /= torch.mean(area)
area = area.unsqueeze(1)
mean, mean_sqr = 0, 0
for data in tqdm(train_dataloader):
invar = data[0] # [b, N, h, w]
outvar = data[1] # [b, N, h, w]
diff = outvar - invar
weighted_diff = area * diff
weighted_diff_sqr = torch.square(weighted_diff)
mean += torch.mean(weighted_diff, dim=(2, 3)) / len(train_dataloader)
mean_sqr += torch.mean(weighted_diff_sqr, dim=(2, 3)) / len(train_dataloader)
variance = mean_sqr - mean**2 # [1,num_channel, 1,1]
std = torch.sqrt(variance)
np.save("time_diff_std.npy", std.numpy())
print(f"saving time_diff_std.npy, shapes are {std.numpy().shape}")
if __name__ == "__main__":
current_path = os.getcwd()
sys.path.append(current_path)
main() |