File size: 1,098 Bytes
ffdf763
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
from pathlib import Path
import sys,numpy as np,torch
ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
from model.precip_extremes_gan import *
c=load_config(ROOT);d=np.load(ROOT/c["data"]["path"]);ck=torch.load(ROOT/c["paths"]["checkpoint"],map_location="cpu",weights_only=True);net=PrecipExtremesGAN(**ck["model_config"]);net.load_state_dict(ck["model"]);net.eval();mask=d["split"]=="test";x=torch.from_numpy(d["predictors"][mask]);members=[]
with torch.no_grad():
 base=net.baseline(x,c["data"]["high_grid"])
 for _ in range(c["evaluation"]["ensemble_members"]):members.append(base+net.generator(torch.cat((x,torch.randn(len(x),1,*x.shape[-2:])),1),c["data"]["high_grid"]))
z=torch.stack(members);p=(torch.exp(z)-c["data"]["log_epsilon"]).clamp_min(0).numpy();b=(torch.exp(base)-c["data"]["log_epsilon"]).clamp_min(0).numpy();path=ROOT/c["paths"]["predictions"];path.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(path,predictors=x.numpy(),target=d["precipitation"][mask],baseline=b,members=p,mean=p.mean(0),period=d["period"][mask],unit=d["unit"]);print(path)