from pathlib import Path import sys,numpy as np,torch ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT)) from model.sfno_bvmc import * c=load_config(ROOT);torch.manual_seed(c["seed"]);ck=torch.load(ROOT/c["paths"]["checkpoint"],map_location="cpu",weights_only=True);base=CompactSFNO(**ck["model_config"]);base.load_state_dict(ck["model"]);models=[] for j in range(c["ensemble"]["checkpoint_members"]): m=CompactSFNO(**ck["model_config"]);m.load_state_dict(ck["model"]) with torch.no_grad(): for p in m.parameters():p.add_(torch.randn_like(p)*(j+1)*1e-4) m.eval();models.append(m) x=state(c["data"]["tile_origins"][0],c["data"]["tile_size"],30,0)[None];members=[] with torch.no_grad(): for m in models: plus,minus=centered_bred_vectors(m,x,c["ensemble"]["bred_cycles"],c["ensemble"]["perturbation_norm"]) for perturb in (plus,minus): current=x+perturb;seq=[] for step in range(c["ensemble"]["forecast_steps"]): out=m(current);seq.append(out[0].numpy());current=torch.cat((out,current[:,-3:]),dim=1) members.append(seq) targets=np.stack([state(c["data"]["tile_origins"][0],32,30,s+1,74).numpy() for s in range(c["ensemble"]["forecast_steps"])]);path=ROOT/c["paths"]["predictions"];path.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(path,ensemble=np.asarray(members),target=targets,lead_hours=np.arange(1,c["ensemble"]["forecast_steps"]+1)*6,checkpoint_indices=np.repeat(np.arange(len(models)),2),bred_sign=np.tile(np.array([1,-1]),len(models)),logical_shape=np.array([74,721,1440]),is_complete_global=np.bool_(False));print(path)