Global-Flood-LSTM / scripts /inference.py
zhangrenchao's picture
Publish Global-Flood-LSTM reproduction
838dfd3 verified
Raw History Blame Contribute Delete
959 Bytes
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)