Spaces:
Running on Zero
Running on Zero
File size: 5,824 Bytes
80f643f 8ae4133 80f643f 8ae4133 80f643f 8ae4133 80f643f 8ae4133 80f643f 8ae4133 80f643f 8ae4133 80f643f 8ae4133 80f643f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 | """
aifs.plot
=========
Cartopy-based helpers for visualising AIFS forecast output on a global map.
All functions return ``matplotlib.figure.Figure`` objects so they work in
both notebooks (``plt.show()``) and scripts (``fig.savefig(...)``).
Quickstart
----------
from aifs.plot import plot_field, plot_field_sequence
# Single map
fig = plot_field(state, "2t", title="2-m Temperature β T+6h")
fig.savefig("t2m_T+6.png", dpi=150)
# Multi-panel sequence
fig = plot_field_sequence(states, "2t", max_steps=4)
fig.savefig("t2m_sequence.png", dpi=150)
"""
from __future__ import annotations
import warnings
import numpy as np
warnings.filterwarnings("ignore", category=UserWarning)
# ββ Variable metadata βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
#: Variables that can be extracted from forecast state dicts
PLOTTABLE = [
"2t", "msl", "sp", "tcw", "10u", "10v", "swh", "mwp",
"t_850", "t_500", "u_850", "v_850", "z_500", "q_700",
]
_CMAP = {
"2t": "RdBu_r", "t_850": "RdBu_r", "t_500": "RdBu_r",
"msl": "viridis", "sp": "viridis",
"10u": "RdBu", "10v": "RdBu",
"u_850": "RdBu", "v_850": "RdBu",
"swh": "Blues", "mwp": "Blues", "tcw": "Blues",
"z_500": "plasma", "q_700": "YlGn",
}
_UNITS = {
"2t": "K", "t_850": "K", "t_500": "K",
"msl": "Pa", "sp": "Pa", "z_500": "mΒ²/sΒ²",
"10u": "m/s", "10v": "m/s", "u_850": "m/s", "v_850": "m/s",
"swh": "m", "mwp": "s", "tcw": "kg/mΒ²", "q_700": "kg/kg",
}
_LONG_NAME = {
"2t": "2-m Temperature",
"msl": "Mean Sea-Level Pressure",
"sp": "Surface Pressure",
"tcw": "Total Column Water",
"10u": "10-m U Wind",
"10v": "10-m V Wind",
"swh": "Significant Wave Height",
"mwp": "Mean Wave Period",
"t_850": "Temperature at 850 hPa",
"t_500": "Temperature at 500 hPa",
"u_850": "U Wind at 850 hPa",
"v_850": "V Wind at 850 hPa",
"z_500": "Geopotential at 500 hPa",
"q_700": "Specific Humidity at 700 hPa",
}
# ββ Grid coordinate extraction ββββββββββββββββββββββββββββββββββββββββββββββββ
def _get_latlons(state: dict) -> tuple[np.ndarray, np.ndarray]:
"""
Return (lats, lons) for the grid the forecast was run on.
The anemoi tensor handler injects ``state["latitudes"]`` and
``state["longitudes"]`` from the checkpoint metadata before the first
inference step, and these are propagated to every output state via
``new_states = input_states.copy()``. We read them directly β no
separate grid-geometry lookup needed.
Longitudes are returned in the range [0, 360) as stored by anemoi;
callers that need [-180, 180) should call ``_to_180(lons)``.
"""
lats = state.get("latitudes")
lons = state.get("longitudes")
if lats is None or lons is None:
raise KeyError(
"State dict does not contain 'latitudes'/'longitudes'. "
"Make sure you are passing a state returned by run_forecast() "
"and have not stripped those keys."
)
lats = np.asarray(lats).ravel()
lons = np.asarray(lons).ravel()
if len(lats) < 3 or len(lons) < 3:
raise ValueError(
f"Grid has only {len(lats)} points β expected ~542 080 for N320. "
"The state latitudes/longitudes may be corrupt."
)
return lats, lons
def _to_180(lons: np.ndarray) -> np.ndarray:
"""Normalise longitudes from [0, 360) to [-180, 180) for Cartopy."""
return np.where(lons > 180, lons - 360, lons)
def _extract_field(state: dict, variable: str) -> np.ndarray | None:
"""Pull ``variable`` out of ``state["fields"]``, return None if missing."""
return state.get("fields", {}).get(variable)
# ββ Public API ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
def plot_field(
state: dict,
variable: str,
) -> "matplotlib.figure.Figure":
"""
Plot a single forecast field on a global map.
Parameters
----------
state:
One element from the list returned by :func:`aifs.forecast.run_forecast`.
variable:
Short name of the field to plot (e.g. ``"2t"``, ``"msl"``).
See :data:`PLOTTABLE` for supported names.
Returns
-------
matplotlib.figure.Figure
"""
import matplotlib.pyplot as plt
import cartopy.crs as ccrs
import cartopy.feature as cfeature
import matplotlib.tri as tri
data = _extract_field(state, variable)
if data is None:
raise KeyError(
f"Variable '{variable}' not found in forecast state. "
f"Available: {sorted(state.get('fields', {}).keys())}"
)
units = _UNITS.get(variable, "")
lname = _LONG_NAME.get(variable, variable)
dt = state.get("date", "")
lats, lons = _get_latlons(state)
lons_plot = _to_180(lons)
fig, ax = plt.subplots(figsize=(11, 6), subplot_kw={"projection": ccrs.PlateCarree()})
ax.coastlines()
ax.add_feature(cfeature.BORDERS, linestyle=":")
triangulation = tri.Triangulation(lons_plot, lats)
contour = ax.tricontourf(triangulation, data, levels=20, transform=ccrs.PlateCarree(), cmap="RdBu_r")
cbar = fig.colorbar(contour, ax=ax, orientation="vertical", shrink=0.7, label=variable)
cbar.set_label(f"{lname} [{units}]", fontsize=10)
plt.title(variable .format(dt))
fig.tight_layout()
plt.show()
return fig
|