ACE2-Seasonal / scripts /inference.py
zhangrenchao's picture
Publish ACE2-Seasonal reproduction
01fa30d verified
Raw
History Blame Contribute Delete
1.04 kB
from pathlib import Path
import sys,numpy as np,torch
ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
from model.ace2_seasonal import *
c=load_config(ROOT);ck=torch.load(ROOT/c["paths"]["checkpoint"],map_location="cpu",weights_only=True);m=ACE2Seasonal(**ck["model_config"]);m.load_state_dict(ck["model"]);base=state(30)[None];members=[]
with torch.no_grad():
for j in range(c["inference"]["ensemble_members"]):
x=base+.002*j;anomaly=x[:,5:7]-state(0)[None,5:7];seq=[]
for s in range(c["inference"]["engineering_steps"]):boundary=state(0,s+1)[None,5:7]+anomaly;boundary[:,1].clamp_(0,1);x=m(x,boundary);seq.append(x[0].numpy())
members.append(seq)
target=np.stack([state(30,s+1).numpy() for s in range(c["inference"]["engineering_steps"])]);p=ROOT/c["paths"]["predictions"];p.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(p,ensemble=np.array(members),target=target,lead_hours=np.arange(1,c["inference"]["engineering_steps"]+1)*6,logical_shape=np.array([8,180,360]),is_complete_global=False);print(p)