"""Run AIFS-Single ONCE and save the FULL global forecast (all 542,080 grid points, all timesteps) to a single .npz file. Any city/location can then be extracted from this file later with zero additional GPU computation. Pipeline step 1/2: 全球GPU計算 -> 保存 (step 2/2 is extract_city_from_global.py: 都市名->緯度経度->Grid Point検索->Pickup->CSV/Graph) Usage: python aifs_run_global.py --lead-hours 360 --out global_forecast.npz """ import argparse import datetime import pathlib 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" # Fields kept for every grid point in the saved output (enough for the GUI's # temperature / humidity / wind / pressure / cloud / precip charts). SAVE_FIELDS = ["2t", "2d", "10u", "10v", "msl", "tcc", "tp"] 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 main(): ap = argparse.ArgumentParser() 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("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, ~15-20 min)...", flush=True) runner = SimpleRunner(CHECKPOINT, device="cuda") dates = [] lat = lon = None series = {name: [] for name in SAVE_FIELDS} for state in runner.run(input_state=input_state, lead_time=args.lead_hours): if lat is None: lat = np.array(state["latitudes"], dtype=np.float32) lon = np.array(state["longitudes"], dtype=np.float32) dates.append(state["date"].isoformat()) f = state["fields"] for name in SAVE_FIELDS: series[name].append(np.array(f[name], dtype=np.float32)) print(f"step done: {state['date'].isoformat()}", flush=True) save_kwargs = {"latitudes": lat, "longitudes": lon, "dates": np.array(dates)} for name in SAVE_FIELDS: save_kwargs[name] = np.stack(series[name]) # shape (n_steps, n_points) out_path = pathlib.Path(args.out) out_path.parent.mkdir(parents=True, exist_ok=True) np.savez_compressed(out_path, **save_kwargs) print(f"DONE: wrote global forecast ({len(dates)} steps x {len(lat)} points) to {out_path}", flush=True) print(f"File size: {out_path.stat().st_size / 1e6:.1f} MB", flush=True) if __name__ == "__main__": main()