from pathlib import Path import sys import numpy as np import torch ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from model.ace2 import build_model, hard_correct, load_config def main(): cfg = load_config(ROOT) data = np.load(ROOT / cfg["data"]["path"]) model = build_model(cfg) checkpoint = torch.load( ROOT / cfg["train"]["checkpoint"], map_location="cpu", weights_only=True ) if checkpoint["format_version"] != cfg["data"]["format_version"]: raise ValueError("checkpoint format_version mismatch") model.load_state_dict(checkpoint["model"]) model.eval() current = torch.from_numpy(data["state"][0, 0].astype(np.float32)).unsqueeze(0) forecast = [] with torch.no_grad(): for step in range(cfg["inference"]["steps"]): forcing = torch.from_numpy(data["forcing"][0, step + 1].astype(np.float32)).unsqueeze(0) current = hard_correct(current, model(current, forcing)) forecast.append(current.squeeze(0).numpy().astype(np.float16)) output = ROOT / cfg["inference"]["output"] output.parent.mkdir(parents=True, exist_ok=True) leads = np.arange(1, cfg["inference"]["steps"] + 1) * cfg["data"]["dt_hours"] np.savez_compressed(output, forecast=np.stack(forecast), lead_hours=leads) print(f"saved {output}: forecast={tuple(np.stack(forecast).shape)}, leads={leads.tolist()}h") if __name__ == "__main__": main()