File size: 2,770 Bytes
f1d3656
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
import numpy as np
import pandas as pd 

def era5_fname():
    return "/gpfs/scratch/ehpc03/data/{}/ml{}/era5_{}_y{}_m{}_ml{}.grib"

def atmorep_pred():
    return "./results/id{}/results_id{}_epoch{}_pred.zarr"

def atmorep_target():
    return "./results/id{}/results_id{}_epoch{}_target.zarr"

def grib_index(field):
    grib_idxs = {"velocity_u": "u",
                 "temperature": "t", 
                 "total_precip": "tp", 
                 "velocity_v": "v", 
                 "velocity_z": "z", 
                 "vorticity" : "vo", 
                 "divergence" : "d", 
                 "specific_humidity": "q"}

    return grib_idxs[field]

##################################################################

def get_BERT(atmorep, field, sample, level):
    atmorep_sample = atmorep[f"{field}/sample={sample:05d}/ml={level:05d}"] 
    data = atmorep_sample.data[0,0] 
    datetime = pd.Timestamp(atmorep_sample.datetime[0,0])
    lats = atmorep_sample.lat[0]
    lons = atmorep_sample.lon[0]
    return data, datetime, lats, lons

def get_forecast(atmorep, field, sample,level_idx):
    atmorep_sample = atmorep[f"{field}/sample={sample:05d}"]
    data = atmorep_sample.data[level_idx, 0]
    datetime = pd.Timestamp(atmorep_sample.datetime[0])
    lats = atmorep_sample.lat
    lons = atmorep_sample.lon
    return data, datetime, lats, lons

######################################

def check_lats(lats_pred, lats_target):
    assert (lats_pred[:] == lats_target[:]).all(), "Mismatch between latitudes"
    assert (lats_pred[:] <= 90.).all(), f"latitudes are between {np.amin(lats_pred)}- {np.amax(lats_pred)}"
    assert (lats_pred[:] >= -90.).all(), f"latitudes are between {np.amin(lats_pred)}- {np.amax(lats_pred)}"

def check_lons(lons_pred, lons_target):
    assert (lons_pred[:] == lons_target[:]).all(), "Mismatch between longitudes"
    assert (lons_pred[:] >= 0.).all(), "longitudes are between {np.amin(lons_pred)}- {np.amax(lons_pred)}"
    assert (lons_pred[:] <= 360.).all(), "longitudes are between {np.amin(lons_pred)}- {np.amax(lons_pred)}"

def check_datetimes(datetimes_pred, datetimes_target):
    assert (datetimes_pred == datetimes_target), "Mismatch between datetimes"

######################################

#calculate RMSE
def compute_RMSE(pred, target):
    return np.sqrt(np.mean((pred-target)**2))


def get_max_RMSE(field):
    #TODO: optimize thresholds
    values = {"temperature" : 3, 
              "velocity_u" : 0.2, #????
              "velocity_v": 0.2,  #????
              "velocity_z": 0.2,  #????
              "vorticity" : 0.2,    #????
              "divergence": 0.2,    #????
              "specific_humidity": 0.2,  #????
              "total_precip": 1, #?????
            }
    
    return values[field]