| 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) |
|
|