""" aifs.compare ============ Unifies AIFS, WeatherNext2, and an ERA5-climatology baseline behind one interface: run any of the three for N 6h steps from a chosen init date, get back a list of forecast states; plot any single one; or compare any two at a matching valid time (RMSE/MAE/bias/correlation), sampling across grids where needed. AIFS's states live on its native irregular N320 grid (per-state ``latitudes``/``longitudes``); WeatherNext2 and the climatology baseline both live on the same regular 0.25° grid (module-level ``aifs.weathernext2.GRID_LATITUDES/GRID_LONGITUDES`` — climatology reuses this grid too, since it's fetched from the same ERA5 source WeatherNext2's own accuracy is checked against). Comparisons involving AIFS sample the regular-grid model onto AIFS's points (nearest-neighbour, reusing aifs.era5_verify's sampling code); WeatherNext2-vs-climatology needs no resampling at all, since both already share the exact same grid. """ from __future__ import annotations import numpy as np MODEL_AIFS = "AIFS" MODEL_WN2 = "WeatherNext2" MODEL_CLIMATOLOGY = "Climatology (ERA5 baseline)" MODELS = [MODEL_AIFS, MODEL_WN2, MODEL_CLIMATOLOGY] STEP_HOURS = 6 #: canonical_name -> per-model field key (+ level for pressure fields) and #: ERA5 (group, var, level) — the intersection of fields available from #: AIFS's PLOTTABLE_FIELDS, WeatherNext2's target variables, and #: aifs.era5_verify.EVAL_FIELD_MAP. Confirmed against each model's own #: config/docs, not guessed — see aifs.weathernext2 and aifs.era5_verify. CANONICAL_FIELDS = { "2m_temperature": { "long_name": "2m Temperature", "units": "K", "aifs": "2t", "wn2": "2m_temperature", "wn2_level": None, "era5_group": "single", "era5_var": "t2m", "era5_level": None, }, "mean_sea_level_pressure": { "long_name": "Mean Sea-Level Pressure", "units": "Pa", "aifs": "msl", "wn2": "mean_sea_level_pressure", "wn2_level": None, "era5_group": "single", "era5_var": "msl", "era5_level": None, }, "10m_u_wind": { "long_name": "10m U Wind", "units": "m/s", "aifs": "10u", "wn2": "10m_u_component_of_wind", "wn2_level": None, "era5_group": "single", "era5_var": "u10", "era5_level": None, }, "10m_v_wind": { "long_name": "10m V Wind", "units": "m/s", "aifs": "10v", "wn2": "10m_v_component_of_wind", "wn2_level": None, "era5_group": "single", "era5_var": "v10", "era5_level": None, }, "100m_u_wind": { "long_name": "100m U Wind", "units": "m/s", "aifs": "100u", "wn2": "100m_u_component_of_wind", "wn2_level": None, "era5_group": "single", "era5_var": "u100", "era5_level": None, }, "100m_v_wind": { "long_name": "100m V Wind", "units": "m/s", "aifs": "100v", "wn2": "100m_v_component_of_wind", "wn2_level": None, "era5_group": "single", "era5_var": "v100", "era5_level": None, }, "temperature_850hpa": { "long_name": "Temperature @ 850hPa", "units": "K", "aifs": "t_850", "wn2": "temperature", "wn2_level": 850, "era5_group": "pressure", "era5_var": "t", "era5_level": 850, }, "temperature_500hpa": { "long_name": "Temperature @ 500hPa", "units": "K", "aifs": "t_500", "wn2": "temperature", "wn2_level": 500, "era5_group": "pressure", "era5_var": "t", "era5_level": 500, }, "u_wind_850hpa": { "long_name": "U Wind @ 850hPa", "units": "m/s", "aifs": "u_850", "wn2": "u_component_of_wind", "wn2_level": 850, "era5_group": "pressure", "era5_var": "u", "era5_level": 850, }, "v_wind_850hpa": { "long_name": "V Wind @ 850hPa", "units": "m/s", "aifs": "v_850", "wn2": "v_component_of_wind", "wn2_level": 850, "era5_group": "pressure", "era5_var": "v", "era5_level": 850, }, "geopotential_500hpa": { "long_name": "Geopotential @ 500hPa", "units": "m^2/s^2", "aifs": "z_500", "wn2": "geopotential", "wn2_level": 500, "era5_group": "pressure", "era5_var": "z", "era5_level": 500, }, "specific_humidity_700hpa": { "long_name": "Specific Humidity @ 700hPa", "units": "kg/kg", "aifs": "q_700", "wn2": "specific_humidity", "wn2_level": 700, "era5_group": "pressure", "era5_var": "q", "era5_level": 700, }, } def extract(model: str, states: list[dict], step_index: int, canonical_name: str): """ Pulls one canonical field out of one model's states at one step, on that model's native grid. Returns ``(lats, lons, values, is_regular)``. ``lats``/``lons`` are always 1-D. For AIFS (``is_regular=False``), ``values`` matches ``lats``/``lons`` length (an irregular point cloud). For WeatherNext2 / the climatology baseline (``is_regular=True``), ``values`` is a full ``(len(lats), len(lons))`` grid. """ if canonical_name not in CANONICAL_FIELDS: raise ValueError(f"'{canonical_name}' is not a canonical field. Available: {sorted(CANONICAL_FIELDS)}") spec = CANONICAL_FIELDS[canonical_name] state = states[step_index] if model == MODEL_AIFS: if spec["aifs"] not in state["fields"]: raise KeyError(f"'{spec['aifs']}' not in this AIFS state. Available: {sorted(state['fields'])}") values = np.asarray(state["fields"][spec["aifs"]]) return np.asarray(state["latitudes"]).ravel(), np.asarray(state["longitudes"]).ravel(), values, False from aifs.weathernext2 import GRID_LATITUDES, GRID_LONGITUDES, PRESSURE_LEVELS if model == MODEL_WN2: if spec["wn2"] not in state["fields"]: raise KeyError(f"'{spec['wn2']}' not in this WeatherNext2 state. Available: {sorted(state['fields'])}") data = np.asarray(state["fields"][spec["wn2"]]) if data.ndim == 3: data = data[PRESSURE_LEVELS.index(spec["wn2_level"])] return GRID_LATITUDES, GRID_LONGITUDES, data, True if model == MODEL_CLIMATOLOGY: if canonical_name not in state["fields"]: raise KeyError(f"'{canonical_name}' not in this climatology state. Available: {sorted(state['fields'])}") return GRID_LATITUDES, GRID_LONGITUDES, np.asarray(state["fields"][canonical_name]), True raise ValueError(f"Unknown model {model!r}. Expected one of {MODELS}.") def compare(model_a: str, states_a: list[dict], step_index_a: int, model_b: str, states_b: list[dict], step_index_b: int, canonical_name: str) -> dict: """ Compares one canonical field between two models at (possibly different) steps, reducing to whichever representation avoids interpolation: if either model is AIFS (irregular), the other is sampled onto AIFS's points; if neither is, both already share one grid and compare directly. Returns ``{"lats", "lons", "values_a", "values_b", "rmse", "mae", "bias", "corr", "n"}`` — metrics computed as ``values_a - values_b``. """ from aifs.era5_verify import _nearest_indices, metrics, sample_at_points lats_a, lons_a, values_a, regular_a = extract(model_a, states_a, step_index_a, canonical_name) lats_b, lons_b, values_b, regular_b = extract(model_b, states_b, step_index_b, canonical_name) if not regular_a and regular_b: lat_idx, lon_idx = _nearest_indices(lats_a, lons_a) pts_lats, pts_lons = lats_a, lons_a pts_a, pts_b = values_a, sample_at_points(values_b, lat_idx, lon_idx) elif regular_a and not regular_b: lat_idx, lon_idx = _nearest_indices(lats_b, lons_b) pts_lats, pts_lons = lats_b, lons_b pts_a, pts_b = sample_at_points(values_a, lat_idx, lon_idx), values_b elif regular_a and regular_b: # Same shared regular grid (WeatherNext2 and/or climatology) — no # resampling needed, but keep the map-plottable 2-D shape for the # caller (metrics() flattens internally regardless). pts_lats, pts_lons = lats_a, lons_a pts_a, pts_b = values_a, values_b else: # Both irregular — only possible comparing AIFS against AIFS. if values_a.shape != values_b.shape: raise ValueError("Both fields are on an irregular grid but have different shapes — can't compare directly.") pts_lats, pts_lons = lats_a, lons_a pts_a, pts_b = values_a, values_b m = metrics(pts_a, pts_b) return {"lats": pts_lats, "lons": pts_lons, "values_a": pts_a, "values_b": pts_b, **m} # ── Plotting ──────────────────────────────────────────────────────────────── def _map_figure(lats, lons, data, title: str, cmap: str = "RdBu_r", is_regular: bool = False): """One map, in whichever style suits the grid: pcolormesh for a regular grid (WeatherNext2 / climatology), tricontourf for an irregular one (AIFS) — matching each model's own existing plotting style.""" import cartopy.crs as ccrs import cartopy.feature as cfeature import matplotlib.pyplot as plt fig, ax = plt.subplots(figsize=(9, 5), subplot_kw={"projection": ccrs.PlateCarree()}) ax.coastlines() ax.add_feature(cfeature.BORDERS, linestyle=":") if is_regular: lons_plot = np.where(lons > 180, lons - 360, lons) order = np.argsort(lons_plot) mesh = ax.pcolormesh( lons_plot[order], lats, data[:, order], transform=ccrs.PlateCarree(), cmap=cmap, shading="auto", ) fig.colorbar(mesh, ax=ax, orientation="vertical", shrink=0.7) else: import matplotlib.tri as tri lons_plot = np.where(lons > 180, lons - 360, lons) triangulation = tri.Triangulation(lons_plot, lats) contour = ax.tricontourf(triangulation, data, levels=20, transform=ccrs.PlateCarree(), cmap=cmap) fig.colorbar(contour, ax=ax, orientation="vertical", shrink=0.7) ax.set_title(title, fontsize=11) fig.tight_layout() return fig def plot_model_field(model: str, states: list[dict], step_index: int, canonical_name: str): """ Single map for one model/step/field. For AIFS, defers to aifs.plot's own plot_field — the app's original, actively-maintained AIFS renderer — rather than duplicating it. WeatherNext2 / climatology use this module's own pcolormesh renderer, since they have no pre-existing one to defer to. """ spec = CANONICAL_FIELDS[canonical_name] if model == MODEL_AIFS: from aifs.plot import plot_field return plot_field(state=states[step_index], variable=spec["aifs"]) lats, lons, values, is_regular = extract(model, states, step_index, canonical_name) date = states[step_index]["date"] title = f"{model} — {spec['long_name']} @ {date}" return _map_figure(lats, lons, values, title, is_regular=is_regular) def plot_compare_maps(model_a: str, states_a: list[dict], step_index_a: int, model_b: str, states_b: list[dict], step_index_b: int, canonical_name: str): """Model A / Model B / (A - B) maps for one canonical field.""" spec = CANONICAL_FIELDS[canonical_name] result = compare(model_a, states_a, step_index_a, model_b, states_b, step_index_b, canonical_name) lats, lons = result["lats"], result["lons"] # If either side came from a regular grid, values_a/values_b were kept # 2-D for that side; render with pcolormesh whenever the shape says so. is_regular = np.asarray(result["values_a"]).ndim == 2 date_a = states_a[step_index_a]["date"] date_b = states_b[step_index_b]["date"] fig_a = _map_figure(lats, lons, result["values_a"], f"{model_a} — {spec['long_name']} @ {date_a}", is_regular=is_regular) fig_b = _map_figure(lats, lons, result["values_b"], f"{model_b} — {spec['long_name']} @ {date_b}", is_regular=is_regular) diff = np.asarray(result["values_a"]) - np.asarray(result["values_b"]) fig_diff = _map_figure(lats, lons, diff, f"{model_a} − {model_b} — {spec['long_name']}", is_regular=is_regular) return fig_a, fig_b, fig_diff, result