File size: 959 Bytes
838dfd3
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
from pathlib import Path
import sys,numpy as np,torch
ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
from model.global_flood_lstm import *
c=load_config(ROOT);ck=torch.load(ROOT/c["paths"]["checkpoint"],map_location="cpu",weights_only=True);members=[];target=[]
with torch.no_grad():
 for member in range(3):
  m=GlobalFloodLSTM(**ck["model_config"]);m.load_state_dict(ck["model"])
  for p in m.parameters():p.add_(torch.randn_like(p)*member*1e-4)
  seq=[]
  for i in range(c["data"]["basins"]):h,f,s,y=synthetic_sample(100+i);loc,scale,tau=m(h[None],f[None],s[None]);seq.append(loc[0].numpy());
  members.append(seq)
 for i in range(c["data"]["basins"]):target.append(synthetic_sample(100+i)[3].numpy())
p=ROOT/c["paths"]["predictions"];p.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(p,prediction=np.asarray(members),target=np.asarray(target),lead_days=np.arange(1,8),basin_ids=np.arange(c["data"]["basins"]));print(p)