TropiCycloneNet / scripts /inference.py
zhangrenchao's picture
Publish TropiCycloneNet reproduction
01384b4 verified
Raw
History Blame Contribute Delete
808 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.tropicyclonenet import *
c=load_config(ROOT);ck=torch.load(ROOT/c["paths"]["checkpoint"],map_location="cpu",weights_only=True);m=TropiCycloneNet(**ck["model_config"]);m.load_state_dict(ck["model"]);m.eval();pred=[];truth=[];prob=[]
with torch.no_grad():
for i in range(c["data"]["samples"]):o,g,e,y,_=synthetic_sample(100+i);p,q=m(o[None],g[None],e[None]);pred.append(p[0].numpy());prob.append(q[0].numpy());truth.append(y.numpy())
path=ROOT/c["paths"]["predictions"];path.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(path,prediction=np.asarray(pred),target=np.asarray(truth),generator_probability=np.asarray(prob),lead_hours=np.arange(1,5)*6);print(path)