FuXi-DA / scripts /inference.py
zhangrenchao's picture
Publish FuXi-DA reproduction
15aff58 verified
Raw
History Blame Contribute Delete
1.25 kB
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)