"""Fetch real ECMWF Open Data initial conditions and run AIFS-Single, extracting a time series at a single (lat, lon) point. CLI-callable so a GUI (Streamlit etc.) can invoke it as a subprocess. Usage: python aifs_fetch_and_run.py --lat 34.3853 --lon 132.4553 --lead-hours 360 --out out.csv Windows fixes applied (see ecmwf-aifs-windows-setup-guide.md for details): - pathlib.PosixPath -> WindowsPath (checkpoint pickled on Linux) - flash_attn shim already installed in the aifs conda env's site-packages - earthkit-regrid URL join patched (backslash bug on Windows) - KMP_DUPLICATE_LIB_OK=TRUE must be set by the caller (env var, not here) """ import argparse import csv import datetime import math import pathlib import sys from collections import defaultdict pathlib.PosixPath = pathlib.WindowsPath import earthkit.data as ekd import earthkit.regrid as ekr import numpy as np from ecmwf.opendata import Client as OpendataClient import earthkit.regrid.db as _ekr_db def _fixed_index_matrix_path(item): return "/".join([_ekr_db.MatrixIndex.matrix_dir_name(item), item["_name"] + ".npz"]) _ekr_db.MatrixIndex.matrix_path = staticmethod(_fixed_index_matrix_path) def _fixed_url_accessor_matrix_path(self, name): url = "/".join([self._url, name]) return _ekr_db.download_and_cache( url, owner="url", verify=True, force=None, chunk_size=1024 * 1024, http_headers=None, update_if_out_of_date=False, maximum_retries=5, retry_after=10, ) _ekr_db.UrlAccessor.matrix_path = _fixed_url_accessor_matrix_path from anemoi.inference.runners.simple import SimpleRunner PARAM_SFC = ["10u", "10v", "2d", "2t", "msl", "skt", "sp", "tcw", "lsm", "z", "slor", "sdor"] PARAM_SOIL = ["vsw", "sot"] PARAM_PL = ["gh", "t", "u", "v", "w", "q"] LEVELS = [1000, 925, 850, 700, 600, 500, 400, 300, 250, 200, 150, 100, 50] SOIL_LEVELS = [1, 2] CHECKPOINT = r"\aifs_models\aifs-single-mse-1.0.ckpt" def get_open_data(date, param, levelist=[]): fields = defaultdict(list) for d in [date - datetime.timedelta(hours=6), date]: data = ekd.from_source("ecmwf-open-data", date=d, param=param, levelist=levelist) seen = set() for f in data: name = f"{f.metadata('param')}_{f.metadata('levelist')}" if levelist else f.metadata("param") if name in seen: continue seen.add(name) assert f.to_numpy().shape == (721, 1440) values = np.roll(f.to_numpy(), -f.shape[1] // 2, axis=1) values = ekr.interpolate(values, {"grid": (0.25, 0.25)}, {"grid": "N320"}) fields[name].append(values) for p, values in fields.items(): fields[p] = np.stack(values) return fields def rh_percent(t_c, td_c): a, b = 17.625, 243.04 es = math.exp((a * t_c) / (b + t_c)) e = math.exp((a * td_c) / (b + td_c)) return round(100.0 * e / es, 1) def main(): ap = argparse.ArgumentParser() ap.add_argument("--lat", type=float, required=True) ap.add_argument("--lon", type=float, required=True) ap.add_argument("--lead-hours", type=int, default=360, help="Forecast length in hours, multiple of 6") ap.add_argument("--out", type=str, required=True) args = ap.parse_args() print(f"Fetching latest ECMWF Open Data...", flush=True) date = OpendataClient().latest() print(f"Initial date is {date}", flush=True) fields = {} fields.update(get_open_data(date, param=PARAM_SFC)) soil = get_open_data(date, param=PARAM_SOIL, levelist=SOIL_LEVELS) mapping = {"sot_1": "stl1", "sot_2": "stl2", "vsw_1": "swvl1", "vsw_2": "swvl2"} for k, v in soil.items(): fields[mapping[k]] = v fields.update(get_open_data(date, param=PARAM_PL, levelist=LEVELS)) for level in LEVELS: gh = fields.pop(f"gh_{level}") fields[f"z_{level}"] = gh * 9.80665 input_state = dict(date=date, fields=fields) print("Loading model and starting forecast (this is the slow part)...", flush=True) runner = SimpleRunner(CHECKPOINT, device="cuda") idx = None rows = [] prev_tp = None for state in runner.run(input_state=input_state, lead_time=args.lead_hours): lat_arr = state["latitudes"] lon_arr = state["longitudes"] if idx is None: lon_signed = np.where(lon_arr > 180, lon_arr - 360, lon_arr) dist2 = (lat_arr - args.lat) ** 2 + (lon_signed - args.lon) ** 2 idx = int(np.argmin(dist2)) print(f"Nearest grid point: lat={lat_arr[idx]:.3f} lon={lon_arr[idx]:.3f}", flush=True) f = state["fields"] t2m_c = float(f["2t"][idx]) - 273.15 d2m_c = float(f["2d"][idx]) - 273.15 u10 = float(f["10u"][idx]) v10 = float(f["10v"][idx]) wind_speed = float(np.hypot(u10, v10)) msl_hpa = float(f["msl"][idx]) / 100.0 tcc = float(f["tcc"][idx]) tp_m = float(f["tp"][idx]) if "tp" in f else None precip_6h_mm = tp_m * 1000.0 if tp_m is not None else None # raw model output is already period-based (SimpleRunner default: not accumulated) row = dict( date=state["date"].isoformat(), date_jst=(state["date"] + datetime.timedelta(hours=9)).strftime("%Y-%m-%d %H:%M") + " JST", t2m_c=round(t2m_c, 1), humidity_pct=rh_percent(t2m_c, d2m_c), wind_speed_ms=round(wind_speed, 1), msl_hpa=round(msl_hpa, 1), cloud_cover=round(tcc, 2), precip_6h_mm=round(precip_6h_mm, 2) if precip_6h_mm is not None else "", ) rows.append(row) print(row, flush=True) out_path = pathlib.Path(args.out) out_path.parent.mkdir(parents=True, exist_ok=True) with open(out_path, "w", newline="", encoding="utf-8") as fp: writer = csv.DictWriter(fp, fieldnames=list(rows[0].keys())) writer.writeheader() writer.writerows(rows) print(f"DONE: wrote {len(rows)} rows to {out_path}", flush=True) if __name__ == "__main__": main()