| from pathlib import Path | |
| import sys,yaml,numpy as np,torch | |
| ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT)) | |
| from model.fuxi_da import FuXiDA,make_sample | |
| c=yaml.safe_load((ROOT/"conf/config.yaml").read_text());ck=torch.load(ROOT/c["paths"]["checkpoint"],map_location="cpu",weights_only=True);assert ck["format_version"]==c["data"]["format_version"];m=FuXiDA(**ck["model_config"]);m.load_state_dict(ck["model"]);m.eval();rows=[] | |
| with torch.no_grad(): | |
| for tile in c["data"]["tile_ids"]: | |
| s=make_sample(tile,20,c["data"]["tile_size"],c["data"]["missing_probability"]);p=m(s["background"][None],s["obs"][None])[0];rows.append((s,p)) | |
| path=ROOT/c["paths"]["predictions"];path.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(path,background=np.stack([r[0]["background"].numpy() for r in rows]),observations=np.stack([r[0]["obs"].numpy() for r in rows]),target=np.stack([r[0]["target"].numpy() for r in rows]),prediction=np.stack([r[1].numpy() for r in rows]),latitude=np.stack([r[0]["latitude"].numpy() for r in rows]),origins=np.stack([r[0]["origin"].numpy() for r in rows]),tile_ids=np.array(c["data"]["tile_ids"]),format_version=np.array(c["data"]["format_version"]),is_complete_global=np.bool_(False));print(path) | |