{ "cells": [ { "cell_type": "markdown", "id": "0", "metadata": {}, "source": [ "# June 2025 European Heatwave Ensemble Reforecast\n", "\n", "This notebook is an ensemble version of `heatwave-reforecast.ipynb` (deterministic). In order to compute the station-based ensemble Continuous Ranked Probability Score (CRPS), you need to the `ensemble` extra from the [stationbench fork](https://github.com/martibosch/stationbench) (e.g., `pip install \"stationbench[ensemble] @ git+https://github.com/martibosch/stationbench.git\"`)." ] }, { "cell_type": "code", "execution_count": null, "id": "1", "metadata": {}, "outputs": [], "source": [ "import datetime as dt\n", "import os\n", "import pathlib\n", "import time\n", "\n", "import cartopy.crs as ccrs\n", "import cartopy.feature as cfeature\n", "import icechunk\n", "import matplotlib.pyplot as plt\n", "import numpy as np\n", "import pandas as pd\n", "import seaborn as sns\n", "import stationbench\n", "import xarray as xr\n", "from stationbench.utils import regions\n", "\n", "from aifs_modal import app, run_forecast" ] }, { "cell_type": "markdown", "id": "2", "metadata": {}, "source": [ "## Set up" ] }, { "cell_type": "code", "execution_count": null, "id": "3", "metadata": {}, "outputs": [], "source": [ "# experiment label\n", "n_members = 10\n", "experiment = \"heatwave-2025-jun-ens\"\n", "\n", "# forecast window: init at heatwave onset for a clean 10-day medium/extended-range eval\n", "start_date = dt.datetime(2025, 6, 20, 0, tzinfo=dt.UTC)\n", "lead_time = 240 # in hours, i.e., until 2025-06-30 00 UTC\n", "end_date = start_date + dt.timedelta(hours=lead_time)\n", "print(f\"Experiment : {experiment}\")\n", "print(f\"Members : {n_members}\")\n", "print(f\"Start : {start_date}\")\n", "print(f\"End : {end_date} (T+{lead_time} h)\")\n", "\n", "# modal storage\n", "storage_bucket = \"aifs-modal-unibe\"\n", "outputs_prefix = \"aifs-outputs-ens\"\n", "outputs_branch = experiment\n", "\n", "# stationbench\n", "region = \"switzerland\"\n", "lat_slice = (44.5, 48.5)\n", "lon_slice = (4.5, 11.5)\n", "stationbench_dir = pathlib.Path(\"../data/stationbench\")\n", "stations_filepath = stationbench_dir / f\"{experiment}-stations.nc\"\n", "era5_filepath = stationbench_dir / f\"{experiment}-era5.nc\"\n", "\n", "# viz\n", "figwidth = plt.rcParams[\"figure.figsize\"][0]\n", "figheight = plt.rcParams[\"figure.figsize\"][1]" ] }, { "cell_type": "markdown", "id": "4", "metadata": {}, "source": [ "## 1. Running the ensemble reforecast on Modal\n", "\n", "Initial conditions are ingested automatically by `run_forecast` when they are\n", "not already present on the Modal IC Volume. No separate ingestion step is\n", "needed." ] }, { "cell_type": "code", "execution_count": null, "id": "5", "metadata": {}, "outputs": [], "source": [ "# Initial conditions are handled automatically by run_forecast.\n", "# No manual ingestion step needed." ] }, { "cell_type": "markdown", "id": "6", "metadata": {}, "source": [ "## 2. Running the ensemble reforecast on Modal\n", "\n", "Ensemble members run sequentially on a single GPU using `run_forecast` with `n_members`:\n" ] }, { "cell_type": "code", "execution_count": null, "id": "7", "metadata": {}, "outputs": [], "source": [ "# with modal.enable_output(): # uncomment for live logs\n", "with app.run():\n", " _start = time.perf_counter()\n", " run_forecast.remote(\n", " start_date,\n", " storage_bucket,\n", " n_members=n_members,\n", " lead_time=lead_time,\n", " outputs_prefix=outputs_prefix,\n", " outputs_branch=outputs_branch,\n", " include_pressure_levels=False,\n", " )\n", " print(\n", " f\"Ran {lead_time}-hour ensemble forecast with {n_members} members in \"\n", " f\"{time.perf_counter() - _start:.1f}s\"\n", " )" ] }, { "cell_type": "markdown", "id": "8", "metadata": {}, "source": [ "## 3. Loading ensemble forecast outputs\n", "\n", "All members are stored in a single zarr group with an `ensemble_member` dimension:" ] }, { "cell_type": "code", "execution_count": null, "id": "9", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ " \u001b[2m2026-03-18T10:44:43.908107Z\u001b[0m \u001b[33m WARN\u001b[0m \u001b[1;33maws_runtime::env_config::normalize\u001b[0m\u001b[33m: \u001b[33mprofile [plugins] ignored; sections in the AWS config file (other than [default]) must have a prefix i.e. [profile my-profile]\u001b[0m\n", " \u001b[2;3mat\u001b[0m /home/conda/feedstock_root/build_artifacts/icechunk_1773674412996/_build_env/.cargo/registry/src/index.crates.io-1949cf8c6b5b557f/aws-runtime-1.5.18/src/env_config/normalize.rs:121\n", "\n", "10 members, 40 time steps\n" ] }, { "data": { "text/html": [ "
<xarray.Dataset> Size: 37GB\n",
"Dimensions: (ensemble_member: 10, valid_time: 40, lat: 721, lon: 1440,\n",
" pressure: 13)\n",
"Coordinates:\n",
" * ensemble_member (ensemble_member) int64 80B 0 1 2 3 4 5 6 7 8 9\n",
" * valid_time (valid_time) datetime64[ns] 320B 2025-06-20T06:00:00 ......\n",
" * lat (lat) float64 6kB 90.0 89.75 89.5 ... -89.5 -89.75 -90.0\n",
" * lon (lon) float64 12kB 0.0 0.25 0.5 0.75 ... 359.2 359.5 359.8\n",
" * pressure (pressure) int64 104B 50 100 150 200 ... 700 850 925 1000\n",
"Data variables: (12/22)\n",
" 100u (ensemble_member, valid_time, lat, lon) float32 2GB dask.array<chunksize=(5, 5, 721, 1440), meta=np.ndarray>\n",
" 100v (ensemble_member, valid_time, lat, lon) float32 2GB dask.array<chunksize=(5, 5, 721, 1440), meta=np.ndarray>\n",
" 10v (ensemble_member, valid_time, lat, lon) float32 2GB dask.array<chunksize=(5, 5, 721, 1440), meta=np.ndarray>\n",
" hcc (ensemble_member, valid_time, lat, lon) float32 2GB dask.array<chunksize=(5, 5, 721, 1440), meta=np.ndarray>\n",
" 10u (ensemble_member, valid_time, lat, lon) float32 2GB dask.array<chunksize=(5, 5, 721, 1440), meta=np.ndarray>\n",
" cp (ensemble_member, valid_time, lat, lon) float32 2GB dask.array<chunksize=(5, 5, 721, 1440), meta=np.ndarray>\n",
" ... ...\n",
" mcc (ensemble_member, valid_time, lat, lon) float32 2GB dask.array<chunksize=(5, 5, 721, 1440), meta=np.ndarray>\n",
" sp (ensemble_member, valid_time, lat, lon) float32 2GB dask.array<chunksize=(5, 5, 721, 1440), meta=np.ndarray>\n",
" tcw (ensemble_member, valid_time, lat, lon) float32 2GB dask.array<chunksize=(5, 5, 721, 1440), meta=np.ndarray>\n",
" strd (ensemble_member, valid_time, lat, lon) float32 2GB dask.array<chunksize=(5, 5, 721, 1440), meta=np.ndarray>\n",
" tp (ensemble_member, valid_time, lat, lon) float32 2GB dask.array<chunksize=(5, 5, 721, 1440), meta=np.ndarray>\n",
" tcc (ensemble_member, valid_time, lat, lon) float32 2GB dask.array<chunksize=(5, 5, 721, 1440), meta=np.ndarray><xarray.Dataset> Size: 3GB\n",
"Dimensions: (time: 1, member: 10, valid_time: 40, latitude: 721,\n",
" longitude: 1440, prediction_timedelta: 40)\n",
"Coordinates:\n",
" * time (time) datetime64[ns] 8B 2025-06-20\n",
" * member (member) int64 80B 0 1 2 3 4 5 6 7 8 9\n",
" * valid_time (valid_time) datetime64[ns] 320B 2025-06-20T06:00:0...\n",
" * latitude (latitude) float64 6kB 90.0 89.75 ... -89.75 -90.0\n",
" * longitude (longitude) float64 12kB 0.0 0.25 0.5 ... 359.5 359.8\n",
" * prediction_timedelta (prediction_timedelta) timedelta64[ns] 320B 06:00:0...\n",
"Data variables:\n",
" 2t (time, member, valid_time, latitude, longitude) float32 2GB dask.array<chunksize=(1, 5, 5, 721, 1440), meta=np.ndarray>\n",
" 10si (time, member, valid_time, latitude, longitude) float32 2GB dask.array<chunksize=(1, 5, 5, 721, 1440), meta=np.ndarray><xarray.Dataset> Size: 4MB\n",
"Dimensions: (time: 1441, station_id: 157)\n",
"Coordinates:\n",
" * time (time) datetime64[ns] 12kB 2025-06-20 ... 2025-06-30\n",
" * station_id (station_id) <U3 2kB 'ABO' 'AEG' 'AIG' ... 'WFJ' 'WYN' 'ZER'\n",
" longitude (station_id) float64 1kB ...\n",
" latitude (station_id) float64 1kB ...\n",
"Data variables:\n",
" 2m_temperature (time, station_id) float64 2MB ...\n",
" 10m_wind_speed (time, station_id) float64 2MB ...<xarray.Dataset> Size: 80MB\n",
"Dimensions: (member: 10, station_id: 157, lead_time: 40, metric: 2,\n",
" valid_time: 40)\n",
"Coordinates:\n",
" * member (member) int64 80B 0 1 2 3 4 5 6 7 8 9\n",
" * station_id (station_id) <U3 2kB 'ABO' 'AEG' 'AIG' ... 'WFJ' 'WYN' 'ZER'\n",
" latitude (station_id) float64 1kB 46.49 47.13 46.33 ... 47.26 46.03\n",
" longitude (station_id) float64 1kB 7.561 8.608 6.924 ... 7.787 7.752\n",
" * lead_time (lead_time) timedelta64[ns] 320B 0 days 06:00:00 ... 10 d...\n",
" * metric (metric) object 16B 'crps' 'mbe'\n",
"Dimensions without coordinates: valid_time\n",
"Data variables:\n",
" 2m_temperature (metric, member, valid_time, station_id, lead_time) float64 40MB ...\n",
" 10m_wind_speed (metric, member, valid_time, station_id, lead_time) float64 40MB ...| \n", " | CRPS (K) | \n", "MBE (K) | \n", "
|---|---|---|
| bin | \n", "\n", " | \n", " |
| AIFS-ENS days 1-7 | \n", "4.616505 | \n", "-0.892011 | \n", "
| AIFS-ENS days 7-10 | \n", "4.943851 | \n", "-1.861350 | \n", "
<xarray.Dataset> Size: 40MB\n",
"Dimensions: (member: 10, station_id: 157, lead_time: 40, valid_time: 40)\n",
"Coordinates:\n",
" * member (member) int64 80B 0 1 2 3 4 5 6 7 8 9\n",
" * station_id (station_id) <U3 2kB 'ABO' 'AEG' 'AIG' ... 'WFJ' 'WYN' 'ZER'\n",
" * lead_time (lead_time) timedelta64[ns] 320B 0 days 06:00:00 ... 10 d...\n",
"Dimensions without coordinates: valid_time\n",
"Data variables:\n",
" 2m_temperature (member, valid_time, station_id, lead_time) float64 20MB ...\n",
" 10m_wind_speed (member, valid_time, station_id, lead_time) float64 20MB ...