From 779ba8569a75af001140eaa824a45ecc323a0e1d Mon Sep 17 00:00:00 2001 From: GiGiKoneti Date: Thu, 19 Mar 2026 14:16:13 +0530 Subject: [PATCH 01/10] feat(metrics): add CRPS, spread-skill ratio and rank histogram for ensemble evaluation Adds probabilistic ensemble evaluation metrics to mllam-verification, as proposed in issue #3 and approved by @mfroelund. New functions in statistics.py: - crps(): wraps scores.probability.crps_for_ensemble via compute_pipeline_statistic, following the same pattern as rmse() and mae(). Uses the fair (unbiased) estimator. Accepts any ensemble_member_dim name. - spread_skill_ratio(): computes ensemble spread / RMSE of ensemble mean. SSR = 1.0 indicates perfect calibration. SSR < 1.0 indicates underdispersion. New function in plot.py: - plot_rank_histogram(): wraps scores.plotdata.rank_histogram to produce a Talagrand diagram. Includes a reference line for perfect calibration. New test fixture in conftest.py: - da_ensemble_prediction_2d_utc: 10-member ensemble DataArray built from the existing deterministic prediction fixture. All functions follow the existing architecture exactly: compute_pipeline_statistic backbone, Google-style docstrings, cell_methods CF-convention attribute, 90-char line length. No new dependencies introduced. All functions use the existing scores>=1.2.0 dependency already pinned in pyproject.toml. Closes #3 --- mllam_verification/operations/statistics.py | 134 +++++++++++++++++++ mllam_verification/plot.py | 64 +++++++++ tests/unit/conftest.py | 18 +++ tests/unit/test_plot.py | 48 +++++++ tests/unit/test_statistics.py | 138 +++++++++++++++++++- 5 files changed, 399 insertions(+), 3 deletions(-) diff --git a/mllam_verification/operations/statistics.py b/mllam_verification/operations/statistics.py index c28a2cb..183092e 100644 --- a/mllam_verification/operations/statistics.py +++ b/mllam_verification/operations/statistics.py @@ -3,6 +3,7 @@ import numpy as np import scores.continuous as scc_cont +import scores.probability as scc_prob import xarray as xr xr.set_options(keep_attrs=True) @@ -207,6 +208,139 @@ def mae( return ds_mae +def crps( + ds_prediction: xr.Dataset | xr.DataArray, + ds_reference: xr.Dataset | xr.DataArray, + ensemble_member_dim: str = "ensemble_member", + groupby: Optional[str] = None, + **stats_op_kwargs, +) -> xr.Dataset | xr.DataArray: + """Compute the Continuous Ranked Probability Score (CRPS). + + Wraps `scores.probability.crps_for_ensemble` via + `compute_pipeline_statistic`. Uses the fair (unbiased) estimator + which correctly accounts for finite ensemble size. + + A perfectly calibrated ensemble achieves minimum CRPS. Lower is better. + Unlike `crps_gauss`, this function makes no distributional assumptions + and accepts raw ensemble member trajectories. + + Args: + ds_prediction: Ensemble forecast dataset or data array. + Must contain `ensemble_member_dim` as a dimension. + ds_reference: Reference (observation) dataset or data array. + Must NOT contain `ensemble_member_dim`. + ensemble_member_dim: Name of the ensemble member dimension + in `ds_prediction`. Defaults to ``"ensemble_member"``. + groupby: Optional dimension name to group results by before + computing the statistic. + **stats_op_kwargs: Additional keyword arguments forwarded to + `scores.probability.crps_for_ensemble`, such as + ``reduce_dims`` or ``preserve_dims``. + + Returns: + Dataset or DataArray with CRPS values. The ensemble member + dimension is collapsed. The ``cell_methods`` attribute records + which dimensions were reduced. + + References: + Zamo, M. & Naveau, P. (2018). Estimation of the Continuous + Ranked Probability Score with Limited Information and + Applications to Ensemble Weather Forecasts. + https://doi.org/10.1007/s11004-017-9709-7 + + Example: + >>> da_crps = crps( + ... da_ensemble_prediction, + ... da_reference, + ... ensemble_member_dim="ensemble_member", + ... reduce_dims=["x", "y"], + ... ) + """ + stats_op_kwargs["ensemble_member_dim"] = ensemble_member_dim + ds_crps = compute_pipeline_statistic( + datasets=[ds_prediction, ds_reference], + stats_op=scc_prob.crps_for_ensemble, + groupby=groupby, + stats_op_kwargs=stats_op_kwargs, + ) + ds_crps.name = getattr(ds_prediction, "name", "crps") + reduce_dims = list(set(ds_reference.dims) - set(ds_crps.dims)) + new_cell_methods = [",".join(reduce_dims) + ": crps"] + if isinstance(ds_crps, xr.DataArray): + update_cell_methods(ds_crps, new_cell_methods) + elif isinstance(ds_crps, xr.Dataset): + for _, da_var in ds_crps.items(): + update_cell_methods(da_var, new_cell_methods) + return ds_crps + + +def spread_skill_ratio( + ds_prediction: xr.Dataset | xr.DataArray, + ds_reference: xr.Dataset | xr.DataArray, + ensemble_member_dim: str = "ensemble_member", + groupby: Optional[str] = None, + **stats_op_kwargs, +) -> xr.Dataset | xr.DataArray: + """Compute the Spread-Skill Ratio (SSR) for ensemble calibration. + + SSR = ensemble_spread / RMSE_of_ensemble_mean + + A perfectly calibrated ensemble has SSR = 1.0. Values below 1.0 + indicate underdispersion (ensemble too confident). Values above 1.0 + indicate overdispersion (ensemble too uncertain). + + Ensemble spread is the mean standard deviation across members. + Skill is the RMSE of the ensemble mean against the reference. + + Args: + ds_prediction: Ensemble forecast dataset or data array. + Must contain `ensemble_member_dim` as a dimension. + ds_reference: Reference (observation) dataset or data array. + ensemble_member_dim: Name of the ensemble member dimension + in `ds_prediction`. Defaults to ``"ensemble_member"``. + groupby: Optional dimension name to group results by. + **stats_op_kwargs: Additional keyword arguments forwarded + to xarray reduction operations, such as ``reduce_dims``. + + Returns: + Dataset or DataArray with SSR values. Values near 1.0 indicate + good ensemble calibration. The ``cell_methods`` attribute + records which dimensions were reduced. + + References: + Fortin, V. et al. (2014). Why Should Ensemble Spread Match + the RMSE of the Ensemble Mean? + https://doi.org/10.1175/MWR-D-14-00037.1 + """ + reduce_dims = stats_op_kwargs.get("reduce_dims", None) + + # Ensemble spread: mean standard deviation across members + spread = ds_prediction.std(dim=ensemble_member_dim) + if reduce_dims: + spread = spread.mean(dim=reduce_dims) + + # Skill: RMSE of ensemble mean + ensemble_mean = ds_prediction.mean(dim=ensemble_member_dim) + squared_error = (ensemble_mean - ds_reference) ** 2 + if reduce_dims: + skill = squared_error.mean(dim=reduce_dims) ** 0.5 + else: + skill = squared_error.mean() ** 0.5 + + ds_ssr = spread / skill + ds_ssr.name = getattr(ds_prediction, "name", "spread_skill_ratio") + + reduce_dims_list = reduce_dims if reduce_dims else [] + new_cell_methods = [",".join(reduce_dims_list) + ": spread_skill_ratio"] + if isinstance(ds_ssr, xr.DataArray): + update_cell_methods(ds_ssr, new_cell_methods) + elif isinstance(ds_ssr, xr.Dataset): + for _, da_var in ds_ssr.items(): + update_cell_methods(da_var, new_cell_methods) + return ds_ssr + + def mean(ds: xr.Dataset | xr.DataArray, **stats_op_kwargs) -> xr.Dataset | xr.DataArray: """Compute the mean across specified dimensions. diff --git a/mllam_verification/plot.py b/mllam_verification/plot.py index e28d69a..2ad75ae 100644 --- a/mllam_verification/plot.py +++ b/mllam_verification/plot.py @@ -4,6 +4,7 @@ from typing import Annotated, Callable, Literal, Optional import matplotlib.pyplot as plt +import scores.plotdata as scc_plotdata import xarray as xr from pydantic import BeforeValidator, validate_call @@ -233,6 +234,69 @@ def plot_single_metric_timeseries( # noqa: C901 return axes +@validate_call(config={"arbitrary_types_allowed": True}) +def plot_rank_histogram( + da_reference: xr.DataArray, + da_prediction: xr.DataArray, + ensemble_member_dim: str = "ensemble_member", + axes: Optional[plt.Axes] = None, + xarray_plot_kwargs: Optional[dict] = None, +) -> plt.Axes: + """Plot a rank histogram (Talagrand diagram) for ensemble evaluation. + + A rank histogram shows the distribution of the rank of each + observation within the sorted ensemble. A flat histogram indicates + a well-calibrated ensemble. A U-shaped histogram indicates + underdispersion. An arch-shaped histogram indicates overdispersion. + + Args: + da_reference: Reference (observation) data array. + da_prediction: Ensemble forecast data array. Must contain + `ensemble_member_dim` as a dimension. + ensemble_member_dim: Name of the ensemble member dimension + in `da_prediction`. Defaults to ``"ensemble_member"``. + axes: Matplotlib axes to plot on. If None, a new figure + and axes are created. + xarray_plot_kwargs: Additional keyword arguments forwarded + to the xarray `.plot.bar()` call. + + Returns: + Matplotlib axes with the rank histogram plotted. + + Example: + >>> fig, ax = plt.subplots() + >>> plot_rank_histogram( + ... da_reference, + ... da_ensemble_prediction, + ... ensemble_member_dim="ensemble_member", + ... axes=ax, + ... ) + """ + if axes is None: + _, axes = plt.subplots() + + if xarray_plot_kwargs is None: + xarray_plot_kwargs = {} + + da_ranks = scc_plotdata.rank_histogram( + fcst=da_prediction, + obs=da_reference, + ens_member_dim=ensemble_member_dim, + ) + + da_ranks.to_series().plot.bar(ax=axes, **xarray_plot_kwargs) + axes.set_xlabel("Rank") + axes.set_ylabel("Relative frequency") + axes.axhline( + y=1.0 / (da_prediction.sizes[ensemble_member_dim] + 1), + color="red", + linestyle="--", + label="Perfect calibration", + ) + axes.legend() + return axes + + def plot_single_metric_gridded_map( # noqa: C901 da_reference: xr.DataArray, da_prediction: xr.DataArray, diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py index 9afc432..66a602f 100644 --- a/tests/unit/conftest.py +++ b/tests/unit/conftest.py @@ -177,3 +177,21 @@ def fixture_da_prediction_2d_utc( data += noise + bias return da_reference_2d_utc.copy(data=data) + + +@pytest.fixture(name="da_ensemble_prediction_2d_utc", scope="session") +def fixture_da_ensemble_prediction_2d_utc( + da_prediction_2d_utc: xr.DataArray, +) -> xr.DataArray: + """Ensemble version of the 2D UTC prediction fixture. + + Creates a 10-member ensemble by adding Gaussian noise to the + deterministic prediction fixture using a seeded RNG for reproducibility. + """ + members = [] + rng = np.random.default_rng(seed=42) + for i in range(10): + noise = rng.normal(0, 0.2, da_prediction_2d_utc.shape) + member = da_prediction_2d_utc.copy(data=da_prediction_2d_utc.values + noise) + members.append(member.assign_coords(ensemble_member=i)) + return xr.concat(members, dim="ensemble_member") diff --git a/tests/unit/test_plot.py b/tests/unit/test_plot.py index c489aac..89d424e 100644 --- a/tests/unit/test_plot.py +++ b/tests/unit/test_plot.py @@ -327,3 +327,51 @@ def test_unexpected_input( time_operation=time_operation, time_op_kwargs=time_op_kwargs, ) + + +class TestPlotRankHistogram: + """Tests for plot_rank_histogram().""" + + def test_returns_axes( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """plot_rank_histogram() should return matplotlib Axes.""" + import matplotlib + + matplotlib.use("Agg") + from mllam_verification.plot import plot_rank_histogram + + axes = plot_rank_histogram( + da_reference_2d_utc, + da_ensemble_prediction_2d_utc, + ensemble_member_dim="ensemble_member", + ) + import matplotlib.pyplot as plt + + assert isinstance(axes, plt.Axes) + plt.close("all") + + def test_accepts_existing_axes( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """plot_rank_histogram() should use provided axes.""" + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + from mllam_verification.plot import plot_rank_histogram + + fig, ax = plt.subplots() + returned_ax = plot_rank_histogram( + da_reference_2d_utc, + da_ensemble_prediction_2d_utc, + ensemble_member_dim="ensemble_member", + axes=ax, + ) + assert returned_ax is ax + plt.close("all") diff --git a/tests/unit/test_statistics.py b/tests/unit/test_statistics.py index d04d3a9..912fb99 100644 --- a/tests/unit/test_statistics.py +++ b/tests/unit/test_statistics.py @@ -25,8 +25,140 @@ def test_rmse( self, da_prediction_2d_utc: xr.DataArray, da_reference_2d_utc: xr.DataArray ): """Test computing the root mean squared error.""" - da_rmse = rmse( - da_prediction_2d_utc, da_reference_2d_utc, reduce_dims=["x", "y"] - ) + da_rmse = rmse(da_prediction_2d_utc, da_reference_2d_utc, reduce_dims=["x", "y"]) assert isinstance(da_rmse, xr.DataArray) assert "cell_methods" in da_rmse.attrs + + +class TestCrps: + """Tests for the crps() function.""" + + def test_crps_returns_dataarray( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """crps() should return a DataArray.""" + from mllam_verification.operations.statistics import crps + + result = crps( + da_ensemble_prediction_2d_utc, + da_reference_2d_utc, + ensemble_member_dim="ensemble_member", + reduce_dims=["x", "y"], + ) + assert isinstance(result, xr.DataArray) + + def test_crps_ensemble_dim_collapsed( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """Ensemble member dimension must not appear in output.""" + from mllam_verification.operations.statistics import crps + + result = crps( + da_ensemble_prediction_2d_utc, + da_reference_2d_utc, + ensemble_member_dim="ensemble_member", + reduce_dims=["x", "y"], + ) + assert "ensemble_member" not in result.dims + + def test_crps_has_cell_methods( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """Output must have cell_methods attribute.""" + from mllam_verification.operations.statistics import crps + + result = crps( + da_ensemble_prediction_2d_utc, + da_reference_2d_utc, + ensemble_member_dim="ensemble_member", + reduce_dims=["x", "y"], + ) + assert "cell_methods" in result.attrs + + def test_crps_non_negative( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """CRPS must be non-negative.""" + from mllam_verification.operations.statistics import crps + + result = crps( + da_ensemble_prediction_2d_utc, + da_reference_2d_utc, + ensemble_member_dim="ensemble_member", + reduce_dims=["x", "y"], + ) + assert float(result.min()) >= 0 + + +class TestSpreadSkillRatio: + """Tests for the spread_skill_ratio() function.""" + + def test_ssr_returns_dataarray( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """spread_skill_ratio() should return a DataArray.""" + from mllam_verification.operations.statistics import spread_skill_ratio + + result = spread_skill_ratio( + da_ensemble_prediction_2d_utc, + da_reference_2d_utc, + ensemble_member_dim="ensemble_member", + reduce_dims=["x", "y"], + ) + assert isinstance(result, xr.DataArray) + + def test_ssr_positive( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """SSR must be positive.""" + from mllam_verification.operations.statistics import spread_skill_ratio + + result = spread_skill_ratio( + da_ensemble_prediction_2d_utc, + da_reference_2d_utc, + ensemble_member_dim="ensemble_member", + reduce_dims=["x", "y"], + ) + assert float(result.min()) > 0 + + def test_ssr_perfect_ensemble_near_one(self): + """A perfectly calibrated ensemble should have SSR near 1.0.""" + import numpy as np + + from mllam_verification.operations.statistics import spread_skill_ratio + + rng = np.random.default_rng(seed=0) + # Truth + truth = xr.DataArray( + rng.normal(0, 1, (100,)), + dims=["time"], + ) + # Ensemble: sample from same distribution as truth + members = [ + xr.DataArray( + rng.normal(0, 1, (100,)), + dims=["time"], + ).assign_coords(ensemble_member=i) + for i in range(50) + ] + ensemble = xr.concat(members, dim="ensemble_member") + result = spread_skill_ratio( + ensemble, + truth, + ensemble_member_dim="ensemble_member", + reduce_dims=["time"], + ) + # For a large well-calibrated ensemble SSR should be near 1.0 + assert 0.5 < float(result) < 2.0 From bd34bab4fa8ae7c68a4985509e3fca41d7566a61 Mon Sep 17 00:00:00 2001 From: GiGiKoneti Date: Wed, 22 Apr 2026 23:48:23 +0530 Subject: [PATCH 02/10] =?UTF-8?q?refactor:=20address=20PR=20review=20?= =?UTF-8?q?=E2=80=94=20remove=20unused=20groupby,=20fix=20arg=20order,=20m?= =?UTF-8?q?ove=20imports?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Remove unused groupby parameter from crps() and spread_skill_ratio() - Swap argument order to (ds_reference, ds_prediction) matching plot calling convention and mae() signature - Add preserve_dims support to spread_skill_ratio() for hovmoller plots - Move matplotlib/plot_rank_histogram imports to top of test_plot.py - Add crps and spread_skill_ratio to timeseries/hovmoller test parametrizations - Add da_ensemble_prediction_2d_elapsed fixture for elapsed-time tests - Update test_statistics.py to match new argument order All 38 tests pass. Pre-commit hooks (isort, black, flake8, mypy) clean. --- mllam_verification/operations/statistics.py | 28 ++++++++++------- tests/unit/conftest.py | 20 ++++++++++++ tests/unit/test_plot.py | 35 ++++++++++++--------- tests/unit/test_statistics.py | 14 ++++----- 4 files changed, 63 insertions(+), 34 deletions(-) diff --git a/mllam_verification/operations/statistics.py b/mllam_verification/operations/statistics.py index 183092e..6df6f7c 100644 --- a/mllam_verification/operations/statistics.py +++ b/mllam_verification/operations/statistics.py @@ -209,10 +209,9 @@ def mae( def crps( - ds_prediction: xr.Dataset | xr.DataArray, ds_reference: xr.Dataset | xr.DataArray, + ds_prediction: xr.Dataset | xr.DataArray, ensemble_member_dim: str = "ensemble_member", - groupby: Optional[str] = None, **stats_op_kwargs, ) -> xr.Dataset | xr.DataArray: """Compute the Continuous Ranked Probability Score (CRPS). @@ -226,14 +225,12 @@ def crps( and accepts raw ensemble member trajectories. Args: - ds_prediction: Ensemble forecast dataset or data array. - Must contain `ensemble_member_dim` as a dimension. ds_reference: Reference (observation) dataset or data array. Must NOT contain `ensemble_member_dim`. + ds_prediction: Ensemble forecast dataset or data array. + Must contain `ensemble_member_dim` as a dimension. ensemble_member_dim: Name of the ensemble member dimension in `ds_prediction`. Defaults to ``"ensemble_member"``. - groupby: Optional dimension name to group results by before - computing the statistic. **stats_op_kwargs: Additional keyword arguments forwarded to `scores.probability.crps_for_ensemble`, such as ``reduce_dims`` or ``preserve_dims``. @@ -251,17 +248,18 @@ def crps( Example: >>> da_crps = crps( - ... da_ensemble_prediction, ... da_reference, + ... da_ensemble_prediction, ... ensemble_member_dim="ensemble_member", ... reduce_dims=["x", "y"], ... ) """ + # groupby is accepted from callers but not used by this metric + stats_op_kwargs.pop("groupby", None) stats_op_kwargs["ensemble_member_dim"] = ensemble_member_dim ds_crps = compute_pipeline_statistic( datasets=[ds_prediction, ds_reference], stats_op=scc_prob.crps_for_ensemble, - groupby=groupby, stats_op_kwargs=stats_op_kwargs, ) ds_crps.name = getattr(ds_prediction, "name", "crps") @@ -276,10 +274,9 @@ def crps( def spread_skill_ratio( - ds_prediction: xr.Dataset | xr.DataArray, ds_reference: xr.Dataset | xr.DataArray, + ds_prediction: xr.Dataset | xr.DataArray, ensemble_member_dim: str = "ensemble_member", - groupby: Optional[str] = None, **stats_op_kwargs, ) -> xr.Dataset | xr.DataArray: """Compute the Spread-Skill Ratio (SSR) for ensemble calibration. @@ -294,12 +291,11 @@ def spread_skill_ratio( Skill is the RMSE of the ensemble mean against the reference. Args: + ds_reference: Reference (observation) dataset or data array. ds_prediction: Ensemble forecast dataset or data array. Must contain `ensemble_member_dim` as a dimension. - ds_reference: Reference (observation) dataset or data array. ensemble_member_dim: Name of the ensemble member dimension in `ds_prediction`. Defaults to ``"ensemble_member"``. - groupby: Optional dimension name to group results by. **stats_op_kwargs: Additional keyword arguments forwarded to xarray reduction operations, such as ``reduce_dims``. @@ -313,8 +309,16 @@ def spread_skill_ratio( the RMSE of the Ensemble Mean? https://doi.org/10.1175/MWR-D-14-00037.1 """ + # groupby is accepted from callers but not used by this metric + stats_op_kwargs.pop("groupby", None) + preserve_dims = stats_op_kwargs.pop("preserve_dims", None) reduce_dims = stats_op_kwargs.get("reduce_dims", None) + # Derive reduce_dims from preserve_dims if not explicitly provided + if reduce_dims is None and preserve_dims is not None: + all_dims = [d for d in ds_prediction.dims if d != ensemble_member_dim] + reduce_dims = [d for d in all_dims if d not in preserve_dims] + # Ensemble spread: mean standard deviation across members spread = ds_prediction.std(dim=ensemble_member_dim) if reduce_dims: diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py index 66a602f..ae89fc7 100644 --- a/tests/unit/conftest.py +++ b/tests/unit/conftest.py @@ -195,3 +195,23 @@ def fixture_da_ensemble_prediction_2d_utc( member = da_prediction_2d_utc.copy(data=da_prediction_2d_utc.values + noise) members.append(member.assign_coords(ensemble_member=i)) return xr.concat(members, dim="ensemble_member") + + +@pytest.fixture(name="da_ensemble_prediction_2d_elapsed", scope="session") +def fixture_da_ensemble_prediction_2d_elapsed( + da_prediction_2d_elapsed: xr.DataArray, +) -> xr.DataArray: + """Ensemble version of the 2D elapsed prediction fixture. + + Creates a 10-member ensemble by adding Gaussian noise to the + deterministic prediction fixture using a seeded RNG for reproducibility. + """ + members = [] + rng = np.random.default_rng(seed=42) + for i in range(10): + noise = rng.normal(0, 0.2, da_prediction_2d_elapsed.shape) + member = da_prediction_2d_elapsed.copy( + data=da_prediction_2d_elapsed.values + noise + ) + members.append(member.assign_coords(ensemble_member=i)) + return xr.concat(members, dim="ensemble_member") diff --git a/tests/unit/test_plot.py b/tests/unit/test_plot.py index 89d424e..47d2165 100644 --- a/tests/unit/test_plot.py +++ b/tests/unit/test_plot.py @@ -3,25 +3,31 @@ from datetime import datetime from typing import Callable, Literal, Optional, Tuple +import matplotlib import matplotlib.pyplot as plt import pytest import xarray as xr import mllam_verification.operations.statistics as mlverif_stats from mllam_verification.plot import ( + plot_rank_histogram, plot_single_metric_gridded_map, plot_single_metric_hovmoller, plot_single_metric_timeseries, ) +matplotlib.use("Agg") + @pytest.fixture(name="time_axis_parameters") def fixture_time_type_parameters( request, da_reference_2d_utc, da_prediction_2d_utc, + da_ensemble_prediction_2d_utc, da_reference_2d_elapsed, da_prediction_2d_elapsed, + da_ensemble_prediction_2d_elapsed, ) -> Tuple[xr.DataArray, xr.DataArray, Callable, bool, str, Callable, dict, int]: """Return a tuple of parameters for the test plot functions.""" ( @@ -32,10 +38,18 @@ def fixture_time_type_parameters( time_op_kwargs, expected_num_lines, ) = request.param + + is_ensemble_metric = stats_operation in [ + mlverif_stats.crps, + mlverif_stats.spread_skill_ratio, + ] + if time_axis == "elapsed": return ( da_reference_2d_elapsed, - da_prediction_2d_elapsed, + da_ensemble_prediction_2d_elapsed + if is_ensemble_metric + else da_prediction_2d_elapsed, stats_operation, include_persistence, time_axis, @@ -45,7 +59,7 @@ def fixture_time_type_parameters( ) return ( da_reference_2d_utc, - da_prediction_2d_utc, + da_ensemble_prediction_2d_utc if is_ensemble_metric else da_prediction_2d_utc, stats_operation, include_persistence, time_axis, @@ -79,6 +93,8 @@ class TestPlotSingleMetricTimeseries: ), (mlverif_stats.mae, False, "groupedby.hour.0", None, {}, 1), (mlverif_stats.rmse, False, "groupedby.hour", mlverif_stats.mean, {}, 1), + (mlverif_stats.crps, False, "UTC", None, {}, 1), + (mlverif_stats.spread_skill_ratio, False, "UTC", None, {}, 1), ], indirect=True, ) @@ -248,6 +264,8 @@ class TestPlotSingleMetricHovmoller: (mlverif_stats.mae, None, "UTC", None, {}, None), (mlverif_stats.mae, None, "groupedby.hour", mlverif_stats.mean, {}, None), (mlverif_stats.rmse, None, "groupedby.hour.0", None, {}, None), + (mlverif_stats.crps, None, "UTC", None, {}, None), + (mlverif_stats.spread_skill_ratio, None, "UTC", None, {}, None), ], indirect=True, ) @@ -338,17 +356,11 @@ def test_returns_axes( da_reference_2d_utc: xr.DataArray, ): """plot_rank_histogram() should return matplotlib Axes.""" - import matplotlib - - matplotlib.use("Agg") - from mllam_verification.plot import plot_rank_histogram - axes = plot_rank_histogram( da_reference_2d_utc, da_ensemble_prediction_2d_utc, ensemble_member_dim="ensemble_member", ) - import matplotlib.pyplot as plt assert isinstance(axes, plt.Axes) plt.close("all") @@ -359,13 +371,6 @@ def test_accepts_existing_axes( da_reference_2d_utc: xr.DataArray, ): """plot_rank_histogram() should use provided axes.""" - import matplotlib - - matplotlib.use("Agg") - import matplotlib.pyplot as plt - - from mllam_verification.plot import plot_rank_histogram - fig, ax = plt.subplots() returned_ax = plot_rank_histogram( da_reference_2d_utc, diff --git a/tests/unit/test_statistics.py b/tests/unit/test_statistics.py index 912fb99..ae718c3 100644 --- a/tests/unit/test_statistics.py +++ b/tests/unit/test_statistics.py @@ -42,8 +42,8 @@ def test_crps_returns_dataarray( from mllam_verification.operations.statistics import crps result = crps( - da_ensemble_prediction_2d_utc, da_reference_2d_utc, + da_ensemble_prediction_2d_utc, ensemble_member_dim="ensemble_member", reduce_dims=["x", "y"], ) @@ -58,8 +58,8 @@ def test_crps_ensemble_dim_collapsed( from mllam_verification.operations.statistics import crps result = crps( - da_ensemble_prediction_2d_utc, da_reference_2d_utc, + da_ensemble_prediction_2d_utc, ensemble_member_dim="ensemble_member", reduce_dims=["x", "y"], ) @@ -74,8 +74,8 @@ def test_crps_has_cell_methods( from mllam_verification.operations.statistics import crps result = crps( - da_ensemble_prediction_2d_utc, da_reference_2d_utc, + da_ensemble_prediction_2d_utc, ensemble_member_dim="ensemble_member", reduce_dims=["x", "y"], ) @@ -90,8 +90,8 @@ def test_crps_non_negative( from mllam_verification.operations.statistics import crps result = crps( - da_ensemble_prediction_2d_utc, da_reference_2d_utc, + da_ensemble_prediction_2d_utc, ensemble_member_dim="ensemble_member", reduce_dims=["x", "y"], ) @@ -110,8 +110,8 @@ def test_ssr_returns_dataarray( from mllam_verification.operations.statistics import spread_skill_ratio result = spread_skill_ratio( - da_ensemble_prediction_2d_utc, da_reference_2d_utc, + da_ensemble_prediction_2d_utc, ensemble_member_dim="ensemble_member", reduce_dims=["x", "y"], ) @@ -126,8 +126,8 @@ def test_ssr_positive( from mllam_verification.operations.statistics import spread_skill_ratio result = spread_skill_ratio( - da_ensemble_prediction_2d_utc, da_reference_2d_utc, + da_ensemble_prediction_2d_utc, ensemble_member_dim="ensemble_member", reduce_dims=["x", "y"], ) @@ -155,8 +155,8 @@ def test_ssr_perfect_ensemble_near_one(self): ] ensemble = xr.concat(members, dim="ensemble_member") result = spread_skill_ratio( - ensemble, truth, + ensemble, ensemble_member_dim="ensemble_member", reduce_dims=["time"], ) From 13e76f211eafc5fd24112e161b21fe0cd78bedef Mon Sep 17 00:00:00 2001 From: GiGiKoneti Date: Thu, 23 Apr 2026 13:00:21 +0530 Subject: [PATCH 03/10] test(plot): add probabilistic metrics to unexpected input tests and enable groupby support --- mllam_verification/operations/statistics.py | 29 ++++++++------ tests/unit/test_plot.py | 42 +++++++++++++++++++++ 2 files changed, 59 insertions(+), 12 deletions(-) diff --git a/mllam_verification/operations/statistics.py b/mllam_verification/operations/statistics.py index 6df6f7c..eb9e743 100644 --- a/mllam_verification/operations/statistics.py +++ b/mllam_verification/operations/statistics.py @@ -254,22 +254,24 @@ def crps( ... reduce_dims=["x", "y"], ... ) """ - # groupby is accepted from callers but not used by this metric - stats_op_kwargs.pop("groupby", None) + groupby = stats_op_kwargs.pop("groupby", None) stats_op_kwargs["ensemble_member_dim"] = ensemble_member_dim ds_crps = compute_pipeline_statistic( datasets=[ds_prediction, ds_reference], stats_op=scc_prob.crps_for_ensemble, stats_op_kwargs=stats_op_kwargs, + groupby=groupby, ) - ds_crps.name = getattr(ds_prediction, "name", "crps") - reduce_dims = list(set(ds_reference.dims) - set(ds_crps.dims)) - new_cell_methods = [",".join(reduce_dims) + ": crps"] - if isinstance(ds_crps, xr.DataArray): - update_cell_methods(ds_crps, new_cell_methods) - elif isinstance(ds_crps, xr.Dataset): - for _, da_var in ds_crps.items(): - update_cell_methods(da_var, new_cell_methods) + + if isinstance(ds_crps, (xr.DataArray, xr.Dataset)): + ds_crps.name = getattr(ds_prediction, "name", "crps") + reduce_dims = list(set(ds_reference.dims) - set(ds_crps.dims)) + new_cell_methods = [",".join(reduce_dims) + ": crps"] + if isinstance(ds_crps, xr.DataArray): + update_cell_methods(ds_crps, new_cell_methods) + elif isinstance(ds_crps, xr.Dataset): + for _, da_var in ds_crps.items(): + update_cell_methods(da_var, new_cell_methods) return ds_crps @@ -309,8 +311,7 @@ def spread_skill_ratio( the RMSE of the Ensemble Mean? https://doi.org/10.1175/MWR-D-14-00037.1 """ - # groupby is accepted from callers but not used by this metric - stats_op_kwargs.pop("groupby", None) + groupby = stats_op_kwargs.pop("groupby", None) preserve_dims = stats_op_kwargs.pop("preserve_dims", None) reduce_dims = stats_op_kwargs.get("reduce_dims", None) @@ -335,6 +336,9 @@ def spread_skill_ratio( ds_ssr = spread / skill ds_ssr.name = getattr(ds_prediction, "name", "spread_skill_ratio") + if groupby: + ds_ssr = ds_ssr.groupby(groupby) + reduce_dims_list = reduce_dims if reduce_dims else [] new_cell_methods = [",".join(reduce_dims_list) + ": spread_skill_ratio"] if isinstance(ds_ssr, xr.DataArray): @@ -342,6 +346,7 @@ def spread_skill_ratio( elif isinstance(ds_ssr, xr.Dataset): for _, da_var in ds_ssr.items(): update_cell_methods(da_var, new_cell_methods) + return ds_ssr diff --git a/tests/unit/test_plot.py b/tests/unit/test_plot.py index 47d2165..f618ecb 100644 --- a/tests/unit/test_plot.py +++ b/tests/unit/test_plot.py @@ -144,8 +144,26 @@ def test_expected_input( {"dim": "start_time"}, 1, ), + ( + mlverif_stats.crps, + True, + "elapsed", + mlverif_stats.mean, + {"dim": "start_time"}, + 1, + ), + ( + mlverif_stats.spread_skill_ratio, + True, + "elapsed", + mlverif_stats.mean, + {"dim": "start_time"}, + 1, + ), (mlverif_stats.mae, False, "grouped.hour.0", None, {}, 1), + (mlverif_stats.crps, False, "grouped.hour.0", None, {}, 1), (mlverif_stats.rmse, False, "groupedby.hour", None, {}, 1), + (mlverif_stats.spread_skill_ratio, False, "groupedby.hour", None, {}, 1), (mlverif_stats.rmse, False, "groupedby.time.hour.0", None, {}, 1), ], indirect=True, @@ -307,7 +325,23 @@ def test_expected_input( {}, None, ), + ( + mlverif_stats.crps, + None, + "elapsed", + None, + {}, + None, + ), (mlverif_stats.mae, None, "groupedby.hour", None, {}, None), + ( + mlverif_stats.spread_skill_ratio, + None, + "groupedby.hour", + None, + {}, + None, + ), ( mlverif_stats.rmse, None, @@ -316,6 +350,14 @@ def test_expected_input( {}, None, ), + ( + mlverif_stats.crps, + None, + "groupedby.hour.0", + mlverif_stats.mean, + {}, + None, + ), ], indirect=True, ) From 9fc4d5ef37dbb507e7c8b0027baf917ff67524dd Mon Sep 17 00:00:00 2001 From: GiGiKoneti Date: Thu, 23 Apr 2026 14:31:38 +0530 Subject: [PATCH 04/10] docs: add CHANGELOG.md --- CHANGELOG.md | 22 ++++++++++++++++++++++ 1 file changed, 22 insertions(+) create mode 100644 CHANGELOG.md diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 0000000..c09e9e9 --- /dev/null +++ b/CHANGELOG.md @@ -0,0 +1,22 @@ +# Changelog + +All notable changes to this project will be documented in this file. + +The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/), +and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). + +## [Unreleased] + +### Added +- Probabilistic ensemble evaluation metrics: + - `crps()`: Continuous Ranked Probability Score for ensemble forecasts using the fair estimator. + - `spread_skill_ratio()`: Ratio of ensemble spread to the RMSE of the ensemble mean. +- Plotting functions for ensemble verification: + - `plot_rank_histogram()`: Generates Talagrand diagrams to evaluate ensemble calibration. +- Ensemble test fixtures (`da_ensemble_prediction_2d_utc`, `da_ensemble_prediction_2d_elapsed`) in `tests/unit/conftest.py`. +- Full `groupby` support for all statistical metrics (including CRPS and SSR) to enable grouped verification in plotting pipelines. +- Expanded unit tests for ensemble metrics and plotting functions, including unexpected input validation. + +### Changed +- Standardized statistical function signatures to `(ds_reference, ds_prediction)` to match plotting conventions. +- Improved plotting function robustness for `elapsed` and `UTC` time axes. From 827436c7d29cdab6d62c4a65b7ab1295f73eca94 Mon Sep 17 00:00:00 2001 From: GiGiKoneti Date: Thu, 23 Apr 2026 14:33:31 +0530 Subject: [PATCH 05/10] docs: match CHANGELOG style with neural-lam --- CHANGELOG.md | 20 ++++++++++++-------- 1 file changed, 12 insertions(+), 8 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index c09e9e9..cce57f0 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,15 +8,19 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] ### Added + - Probabilistic ensemble evaluation metrics: - - `crps()`: Continuous Ranked Probability Score for ensemble forecasts using the fair estimator. - - `spread_skill_ratio()`: Ratio of ensemble spread to the RMSE of the ensemble mean. + - `crps()`: Continuous Ranked Probability Score for ensemble forecasts using the fair estimator. + - `spread_skill_ratio()`: Ratio of ensemble spread to the RMSE of the ensemble mean. + [\#4](https://github.com/mllam/mllam-verification/pull/4) @GiGiKoneti - Plotting functions for ensemble verification: - - `plot_rank_histogram()`: Generates Talagrand diagrams to evaluate ensemble calibration. -- Ensemble test fixtures (`da_ensemble_prediction_2d_utc`, `da_ensemble_prediction_2d_elapsed`) in `tests/unit/conftest.py`. -- Full `groupby` support for all statistical metrics (including CRPS and SSR) to enable grouped verification in plotting pipelines. -- Expanded unit tests for ensemble metrics and plotting functions, including unexpected input validation. + - `plot_rank_histogram()`: Generates Talagrand diagrams to evaluate ensemble calibration. + [\#4](https://github.com/mllam/mllam-verification/pull/4) @GiGiKoneti +- Ensemble test fixtures (`da_ensemble_prediction_2d_utc`, `da_ensemble_prediction_2d_elapsed`) in `tests/unit/conftest.py` [\#4](https://github.com/mllam/mllam-verification/pull/4) @GiGiKoneti +- Full `groupby` support for all statistical metrics (including CRPS and SSR) to enable grouped verification in plotting pipelines [\#4](https://github.com/mllam/mllam-verification/pull/4) @GiGiKoneti +- Expanded unit tests for ensemble metrics and plotting functions, including unexpected input validation [\#4](https://github.com/mllam/mllam-verification/pull/4) @GiGiKoneti ### Changed -- Standardized statistical function signatures to `(ds_reference, ds_prediction)` to match plotting conventions. -- Improved plotting function robustness for `elapsed` and `UTC` time axes. + +- Standardized statistical function signatures to `(ds_reference, ds_prediction)` to match plotting conventions [\#4](https://github.com/mllam/mllam-verification/pull/4) @GiGiKoneti +- Improved plotting function robustness for `elapsed` and `UTC` time axes [\#4](https://github.com/mllam/mllam-verification/pull/4) @GiGiKoneti From cf8a46c112a57ebe7068cf5bbebbb3e9423c9ae2 Mon Sep 17 00:00:00 2001 From: GiGiKoneti Date: Thu, 23 Apr 2026 19:31:46 +0530 Subject: [PATCH 06/10] style: run black on test_plot.py --- tests/unit/test_plot.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/tests/unit/test_plot.py b/tests/unit/test_plot.py index f618ecb..da18885 100644 --- a/tests/unit/test_plot.py +++ b/tests/unit/test_plot.py @@ -47,9 +47,11 @@ def fixture_time_type_parameters( if time_axis == "elapsed": return ( da_reference_2d_elapsed, - da_ensemble_prediction_2d_elapsed - if is_ensemble_metric - else da_prediction_2d_elapsed, + ( + da_ensemble_prediction_2d_elapsed + if is_ensemble_metric + else da_prediction_2d_elapsed + ), stats_operation, include_persistence, time_axis, From 15bc82aa74e5f4e3e78b7486e35994f2a66e1b85 Mon Sep 17 00:00:00 2001 From: GiGiKoneti Date: Mon, 18 May 2026 15:24:03 +0530 Subject: [PATCH 07/10] chore: re-trigger CI checks for linting From 30a7c2968db8d64c46f122d5cf4f3c7294286bbe Mon Sep 17 00:00:00 2001 From: GiGiKoneti Date: Mon, 18 May 2026 15:51:06 +0530 Subject: [PATCH 08/10] chore: re-trigger CI checks From a4f331c2bfe4233ce5e09bdc967e5bdeba687040 Mon Sep 17 00:00:00 2001 From: GiGiKoneti Date: Wed, 27 May 2026 00:35:33 +0530 Subject: [PATCH 09/10] chore: re-run CI checks after GitHub outage --- .DS_Store | Bin 0 -> 8196 bytes docs/_images/rank_histogram_example.png | Bin 0 -> 15634 bytes mllam_verification/operations/statistics.py | 271 ++++++++++++++++- mllam_verification/plot.py | 93 ++++++ tests/unit/test_plot.py | 60 ++++ tests/unit/test_statistics.py | 304 ++++++++++++++++++++ 6 files changed, 727 insertions(+), 1 deletion(-) create mode 100644 .DS_Store create mode 100644 docs/_images/rank_histogram_example.png diff --git a/.DS_Store b/.DS_Store new file mode 100644 index 0000000000000000000000000000000000000000..324f1e93a5435b2634f867ee6404cba6a2fb4288 GIT binary patch literal 8196 zcmeHLKX21O9DPS!(Ncj*r*60#10Z!_HdUAyTZyR(aZ?&Z97j!3sVw+=@Hr3?18YAD zV+SN=-u zWzoDk*q9W6nDDYG{GWM%iAf!14l9c$l=rmRgL0wDr5MVFvp>~zV&<^2Xv3jwIFuV% zxeP^_(HWN-I#gy+T5&)eXgZ*C_YJ&2j$@31{CzS!+1)=Xi?J8f!Ys;8Cm)tMnS2bZ z?(eS`uiEchl>H^x<3U43LcJ&WfN7v`>&Y9AS@(sSFnRmUX!81f70v$0xc9Hd$IkLVqSiX`rRz#0pjRNNwvv#?5J+pqtfv!*AFWgE*m~;BRIFDNH zJGXAvup@dOQ12LhYJSfBHscBo(GS!D&!{%a3$^0V?XecF@91J~MCU`J^Bhw?6BD*= z-ePk=wF}tunw`H3BSPiZM?~}|^tj|+2%q5O*GsCc`$BE?uJihRCC&cGxLaJ=A^&>B zF3#xQDcATttKOSaVu?|(-sP62vbHiAFF3ZJ6=DE%{WodRtPAvJYGj#g# zkI$S{gcZ)A_B6`$6)$2&^sUxAH`Jweeo~j$8|uNK^E6=iI%15-_#8ya9dKa9YTVZO z|LgYe|94;}B2gR=2maOpm2~zydyICub>lud*G^gAvFf6El|>W6#-!s!la3QzI&Qcr lsX8WeSXty1mVfXeV5#2@k@GEbUI)KwzNAmJ0 zfluhb{>QZ!oNjmT2?X7F+1o%~k5;*>Uv|BaMSIIel`HFNIXd0JzS1y{>uR~Jb!^SO z>HbvPoO@+%Hu)+_&7vnA&iiTBjL`yQ(RI;NJ4ZfaNQ*I%jrIih3btAC*<@;=q&8N# z4GK<61u1$2B7hMe4}n0_{~x^mu9$f!Vc_FaXISH+eE#xvdJ)vrhZ<1?!i-Y%+(Lmt z<=63X6V!}ab-`2muLOkBr%s*fUfjY%&|}5?kLzMM^}0)+Mhuo(Vx^APh^l?|O0@HJ z)W`f!d~*#d3wuh&{ZHycX(jZxmxtLbJ5pokdvdTJpZu}97~gJnA^7MtIp_V~lBqNZ z#GtB&N7bI<*s;Mvu}t&JD;#FOe}3X~Uop(pElO>T=e@08X2sAHP8XLh|B98Jj)H>S z9E*+(qY-O)BlC0wrDnNgbdxfGI0%9GS<{=zKlK5HCq_uR(3{ItYSD(8dStcP+vx94 zA!K*|7!0sy1p)Ymyy)9#4V9=~}v*!D%pPcq4zP&>I{wae=lWL6N(B;0@bQ*z3{nSmL zsgwN7U@ab@kH6tuZj*tj$1ZDBi7Rctp}Bl!1c zJ!QqOY`_vnJewte0vn`~!ABrIoc{AQXgYuH)XsV8im_$3>DH5r_u_`m&m&Bn%VNu`Z1(*c36R>Zz$#_9!;wP>X`X%O7^_*qLUA zC}GFMmYJ6!QOTItE0JOi)mlbEk+tsf6BAL55+)7Wv2A3&x+}X=`}Jdt6I4`8#%mZl z;d|!Gvql|7sg_r0xJb#OeSA+c-}m=$w+t&zI3FJyuZ6rKS-sqT`*X>dQ}G1Z)y2BP z8@k*it0yP%3zN1IQPCgl?A|MBYpnA*Krc|`WLI_w@Dv9e7HYf`t0ao!D;ic|4FyqxmR(>ilr^_Gc?mK!&mNz z43>~cR+A3Mbf}_St9A7LnQx&~ z%(w2l?O$V$O%YrSFEOadG<0i;k!pGVJljyGBR3^=F{Wj~);y=^p^CYQ?^bw;Zc(k`p-%H+vi@Rij#l^MmpeLX;CORJPVNNI5 zYOa>MqhmA~ZNV!mVNL=}HN|4A51n(wYE4N6J-cZ(#A-~t4{08rjMSEjv3AXl zsgLr`IIqsETx+t#ydLExklI;U2orI06jo$i4I^mux_5YZ#KOthy60$hYAPsGPD-3W zN}4@Zl}opU^&;QMeIXU8y7@aocOSUXnUndE3`*T%JUKagLWV)Rx{Yl;HSOTv<=1E# zGeFi}fA!pck{82LJoR(?VaM{>dS8|7-g;f_RMyRl| zD~j<{%eAu^ODgnPp};NJds&TtHkiNK%zyhD9gFw-+V-Ht=sYwZ2Psy>&YS<-Na(5F zB>fS<#)$9qN5fA%#ABLKUIgMJ>1lujA4y%B5CN~|0?&**{i~e>{@GOiC$A9_f6hbI z2M^v%wWk;?^yRa?efyRjZ2~7s`~Cf{m?F)rmv%osFmx?fI*z0Lj`o!?Q7kOTRB_VF zmzkYs+tp_+p||qgBQAOL?L39J$A?e~eueGTah|@sd(q@v`ptd$x+&=KYUdo6`EGfy zofWQgm#8D$mWL9S=ejb&_`K>~Qwg`QrTH*6h0#Pm`Qu#To2gek4b2}dp>c|f;wK2% zCvML7D0*!zOj{K~y z83$x&-62%e2cM^B3DxqOzP-Bo&7u^8<}#T28ba1x=ec3F$Y$#l$W;(uYX%5j+<-TcPFK8d&qWTbu}&X$~W%6 zXD&THyJlvvho>hdmP40xd0-Ll(cOz~G);3)%zdRfn1so!Em5$$#?`j82qrUbrh2Z{ z??e&GA`sF-3t?S_QTpXJn3=XDnuCeS$+j&324eN$E6DrC_CsZuyOOM|-|Z3&q9V?U zdD#Ix5#d}8>-tqD8kGfmrdjD=l)!8HTux3dzKbtUF`QO2U5+67R^Hi**XUs%(YabV ziRiFtY#es1-rIbv&U0ov@XUq8Vb{&jc9Qf_-jkzMfAg+%c?;%C2*ht{^jU@!QP(Kj zfkF%vgF4E3YoV`=nr?fl2P%ujrmr)SM>UQ+Hj+s((IIYq8jBq&x6MLLw}kOg##^OZ z!=6ole@(?o>^Ym_p6NK=fEHJTTYejoD;gQ^(2m0bm<9k{iCiLybR`cyovQ37f$jff zJ&cEt-j)Lhh2WI=fBVX(NK8sf$_Tq@wBHLAZAQg?jsYeYDK#}GEK^(v8CwhNU{8(( z0zt3pEP}6B>sCM@ZtOXIMX&hLv(snKw)7ySxbk(2(DCu{o4G6qL^G;8hHSm_HWB9M zr>E1MZ*P}acF?lx7ACu_ju{;u9hrB&kwL9bDenwihQG(Y6C1ye)$?9FU7%7*;0Ot=eaQK4-m^ze8!hJIfr-FW=*z($y~WQJjWV zcN+F)sPw8)S#?L-!}D6EfT~zLYQ0{(p!2&AWt0WmL`q3{0}eY1sw|i5Wb&mog6I1k3qtoQ`pcgM{^NnFC zUcpk#k5rd(nly-G{M!EL8U|3-&4S;>P@TBpRD?U+|M*GB%a>6tC>}VWy1f>)UC13rI>w3Dh1`&|84)KZ_*Z|mWZD>mp#-1GSNbT(}7WHJn zfu_HJs%FW2_~m(t?Ep6cKK@P2yT$$j{+v4^lNG>QK^dw3(xSOB_&8atW zJOnhxj!9i6xIKO!HVgv`(Ka0f*mSC}1`iRy**Z;cwzD$2KG1l4uRT?QLnw)>%O1!} zL{-%-_kL&qW_Mwod=Iu`p~SIpCfh{$C_UgiXL=LkHdC+Ot6W*(4D9}gE?zh#{Co89 zONt@VZr}-?ldpxR;CN_esa#oGTT^!NhRs{ehZ9fl{a6Iwdc2{2?_lB;kLb3waWHFc zR6KCZZZ3S-hq8Qk=-yQ)^qJ6|3&-I4Kaa_@`T&f#Z7v@lA6qmK3z~kvxOceG(Sy1L zn@(6w9n_k@uXvV-mJ?S=0Q^N+**LV{MG(mSwL_jKM?2GBh@@gF$GlYQ>guQuTu+rk z!Bu(TKCeElruxX{HKkzGQb#OD95j-u|JZMKt?W>0agUprV`wERSt_w3_juqV->St$ z6*gpyy_W^6y>|)|sF(!=w4o)lMzg9%g-h;V_q8EFAWEp?3}lVH=2<*C2RVMUL89xR*2SzB0mYRO7-M?)vj19yF|!!7y5hh3gv|9a2~_ zr1-;Au|Qf|etmgi-uvz@s}*cUf;fIL7RKCVE`5QvowR?BKNcAo863rQZr$6M72_%@ z;Nmizh>!_<-P9qwPIUtmI&NSQU^(nXisQn=!{fFV2e298k~Z76Ki?XO z!}nG%=DvIPE-5W-sugZA(D9;=e}0r3Fn#n$Z6G)-tc(0*Z;jtkO~Ty!`^Aq!wimbP z{13b4eajY`^gKt4%BU1;gISTOoMfzOH>)}Hi)PdOjVtU18DXp4ybRUzXoMXM^A=$$ zaN^6tG>NbaoQDzeO=XBQeW=2o1LiSb9CSeklyX@g~Zn%};49`-7`=GWFLuj!Upxq~gQwTW*jV~RUD+Q#AY%ggrbs-)eAkueHQ zb8@x-Yt_o&kgAaey*IAcKYx`wd11Ywp+PfWhikCT(>Wb|@lRg+*9zsUCkctPEb_YE zT7gw3ixqL2IGuk240^EE)a9#k7_|uXfa_aZNWnpd?lNnwXtENAQBi)w ztMFYX=;daom6R0A#5X%{1C1x+7k4VT-$oTpUYfrao6e8sHmvHLdN}vU1%B^9Bn!aMtly{W+0>GYJ-Jbrj$-c3 z^U!=UvEmFv3vX}RtY#g^-o9-lK@ANr+whXCi|7WC)FE^RIPbx^?a+PqU7KKh84XoB z@&I3o)qAUBeW4i>&1RHX0s4%YBX?y`nN44d+SB8u(vB9WvexW4*R(Mh!4ceC4BOlZ zFnP153*mOGO+GpdSJ-zK7zp#f-iKk_)&MV9(18Bg@~p-bUy*Y$1MDsFKkIIr_pgVf@0e&Qs~VKVda@u5DT60)-jB!zk4LLeJPC&dlTPNihLqXX$MRyY68 z`BdlOv^SrV7PU0s6H#K0t@)2O2%=N?#cMaEx0eP@$#ZLQ^H;A#?)gOu#-h4)mHl40 zw3!GuQCHu$Mcm}oE49G2x*7|>MdpDK$43V_e~zc^OR^Ia6J~&_*N1G2Fn|-(D_fU` zD_O2zS6ZL#Nb|!D{y{PPwe|HVK!ey~(Pyv>5WLWmS6RjTlahq4J$4kw414!c@6KblE5DvY*8{GUWe&{ zZ9F-0WWWF6O|pcyj;QN`+4E^Av`kwKxGWA$V(Y!O_+ItFdS-g7DuGIcgUEkY{XsEm zTBYa4gjzVQ#L8bYP3azSdj8BMTJih+MJA3>NX0PfgzM~m8apIG@YgHPFc4A-M2yzD zCt*dLXJW|MG?V~xYHh*q5Dp1!Q$Qzh@TuBM7&hsED&3&Q#aee2Zs+j{7X!R!Mf#tZ zkPk^rN=3yn?z8dQ=l7>`sKtJ+($(%PRn)<@z5BplbDJ%Dp==q~ivejUYiq}dZ}nr+lk{wvk9W*akYg|yK2%(S3l?S?DV z-C$xs^9k8L2iOk^aa0d(GvZnG2q`HY9gmZflk&Dn)#FbtSh&+nRbn6|&H%|3(zhN7 zFt~#B*qwKeLzn0rOhop$TGMQelp9?7Wf5=Q(Bp>Z;Z~7Jk%giXeh2TM~Wh~L#8K2q(hTIn!qSkD1%V8~ha zB{|m(5co>0dsw400aQdV$TWX`c7cVDPZQU)LE?)5M%MyY>e>7WXygDTg5RZ#(J9bp z0X`Fn1HW*r>BFUttHk%`GS;_;9sep`bo9)bB^M2xT8J4i@wX?{-k8R+O@4pv?6?en z8XY@BFGr>g*I-{BDvwSSw8fopGvI+RcWI?xm~9S0F6oh;%|oIovnE;lxLt zG>nGDkL0x)RAO(G0HQR*wUq(RlDNlR60ls#`mzm2+Xn20RUfMB%(3K^IuD@q`8<9< zNDy-`B#?%Ver6qdh=YY?{Xlz)7!GFR_-Axdx^juAtmVn^5p?otIrD+{dQ7$MtK|Ov z{+0dkCCu&+cUL)?Hcj)Hw^h+W93Y;J*^W54 zHx_m&=NUEY!ICV~TfcYLCfhEr13u>4DX;ZdOW#}_9}k_e9<3?IZb3QF3L%5fGC*~x zmYB>EbPXGl{41o0`FBX6eBXGW4CZ+)fi|u8Qn6HlA{-NpvvL)jccKfjW#1;pbckKy zNJVG-dimdS=~4G_LFf3ogLiY1;30gim-3_6bv-BAXucLsZi#qg3 zKk7)8(5&}zEl#X>Zx#1?CivC`bTWmcPu1ieyFCE|F&^UM?ZAjX&Rx2tFAFo95sqBg z@Fohtn^|)N1Gj|7U&=SORU#?^r+ZiR9LycZv^Wk|w1rNw-^Z2x9so~% z5p(_S=PCPlcz>LY?-B)lTrk5J&11sty|WVQxiOP4Oy~1krj#X93u~*Y3}byn-yLf0 ziA-=eov)G14jf{OAXV8k-|Ym|Ca@f$DeU>ZHm#x!^E_FogdJR%(-utUf+=dx5>m!j z%P53Vvok4#MgYA?`nI`}!UWySte9TTfVRw4w)U3G)(CERz)4m!e$MR>M?SO_Ww?Xw zFEEH+TKrL*F{p{vqXkv0+3dCpn?~~Wo*ZPH>tcUG@})MSy|pH~>1S7sm0~VJ@A%nI zKLw!BtyMZtCV0ki?#>UHr)Ndh`Jl~>&<4#hxT9O?1A$CiKhSB z;BXPwAkH$>=v;u$WNie!vmdTdhP!MP<-W>eEO0>1W6TLu!QFNm-q;VEbRFh>qO0lp z6o%6PqunAG@Q(>$8dv5%C5xe8&59MNgCP=6xy@08-TJg*a#{NnGX)RL3zcbx&r7M zEEcnXOF4|tD~vwQ2a3s*B?$jksFifs$Nw(Q4%fHpBt8T9qrw*)h7 zF&GE@i~t4cjcQHpgpxLw27Af^-a!LZ8E^1q0pb_Z;=4U)X`P3wTw_g|Aoz#5W(=s> z<~x}A&8+aFxjTli1Xwl>+(4lFOPBrs`AUCqmZ8+L1O4%%Gzcw~y`aPZgx2D0wCTH% zd%qaWGEfHwKRkRoSrH$9ZSP?G#9@0D8;1nNC<_)61)9%A4^~){p3@9KXd^y7J%^3w zH8sP{m;!x4>#p=?hV@>3)_b7hfIw-K>ApIK`WI4b;}9MiD$qnBXcMz{Fld>E0d|xH zT#4ea6^@7EIR;#HNc!$86&cqv0|uRGPvL-W75gWd{d@9Uq!HD=aW9W~vMD?cH`U+k zm9{Q}L58q#iXC~w{Go* zpPz2EE?gZd$Ikn@_QGc&+NoII6NiA)s0oaWi)qNk^C>i~vq7WnTz zDU`w>e7uokdUdSc7A049A?LxTKn*NZ^~mWP+slW08`gdruxPGzEEjs1?=7Q?2krP^ zBA8*J^speH(kGo}kbilZgolR*CbdfXkypLP+5}E&*3ufI4Zy)+fB?g01~SkjgqFhA zwzhC;)qg3_S~J>wUf)JcO<|(f+l0rPpAt#Nhjlx|rHA-V?j}cCg z&=V;252@2KN zxi|lU}%_`wVgEyIg0KGlgEFHkH{wp%6uY$VAn5umoO?Q z+QC)Xc+>o4t-k=$g?6Xow5kNlMzU6tx?OSCdg$sZ*R3?#SC0=S5dj?dvfZ6$Oo^nI zG8|AIiXpZ06Fw{GFT+P|hrL>SeIjkYqE%_Hb9TSCOMWj!p*(OgV5#0s%_2jDUd2E< zp!NET%A#G5>4pMkz)1aV4ema58 zre%aei&;o$yCKYb=?5=`utPGAw%|xxGbkgzn zznW%6FPW@QSFXc*EtO~xA8kUzFbb4H)~t@m6t1qWyCp}%>U>U9ch=Un(PKwvbF5<` zDtr!s+RbNsPU9A*JAxU1rDD5FuL)0@??vFVnQu%-A7tUl{%E6`STvJhAM^Ikym(C#s)Gue2y?3m9+8#CTrl`?OIj%aSW$1|-I8-}ia@$O@uO?NFu7pB)`cZ^FOw<+X=`~DiC zNZ9QsPB?DA=4m^jook$5SSR8Zb!fXFx978XqPAD3VT_cZgVqN!Leada_j36_tB6Wf z@{Acp^@2B zQc34{J?hOB*;9YW7n&dq=Op`^x~yk+>`sc74p$83yGz3lQEPQR4X6;i5ns8NZ}t=9^m{z-hDvNO5T zPfpE%Cp+5aOCXEz5Dkp5o2qR)iEQGCiTUy4TJg6gY|pZ&()gz@W~+tEtW6o}ljHuq zOK`-$7Im)51a{MPC~_P%=sTep85vvsXFv`YDxh=QH>q>G^6u$5qC|31Z$uP%qj4ka zK;a>yYnoNEUBc*aMN3m@4t73Gqa`KSAp26yEkiRiv%C#%z;ft!c0i8QRoej;$k+M# zOUm5K%h~-wQt7B{X!l-Del%9hoe}Kkx?GL@zMo%c1LB*`pXYm$HfBL2TSoPrHsBr3 zkpg0#nAe0raF7I5bwy5IjcVkzoVcRfDCNPpS0aLHp|dAsM9s7w>1xO=2<*P9+X_2o zTQRLMEGw6hSC9*ii;LqN{~(VV)9MHsJw3#=W{Jt!PJRVa#T_unaJ>L$$>ON_U)U14DJ=ow|OAscI^Cg9UY^Hkl{*I*ECt0HAC!M2wZdU#w zAz>uK;zx2I6_Oea3@8~WB0l!OZANlAHn^}ah8+@RojkD}_*8#Gr_e@U*1@omRd7+T zs;H@RPZprsbZ@R!ce$-@j(MI=feopy<4D!7GC?_i?Q%(WOk88Y?zxz#&gGNX8p5bc zvl4DoZ}u4cOgB$Qr>$9GWJ@55GWsyrb*_^k`hqD@GR|#A36!}qE=6se!CUpU+uEur z*GM&Oc@oIFZS(4t!*)OMa({REt1|i12tF}|Ch1dHbhJ)~YfRikBx3cI_Qai)y_@kSbI;wwS)zz2Z)I0wJlNb#H1r4WtN6{by>$UfkVMx43k zTpKpl)gMIHCe1%wsT zGTeGbIBwm#pQoMsC{RRwgpiOhmR7+{WKzhMh5?4iqzELc8dAJZe(|bG7ZJ1 zrm%?DT!8penoo&63g;~KbiaM_;y8FAJeyU+CNNK8aAUEkQHX=25_IO6t}g2BmBgPq z_N~dRJg#0kiP{`&CcekHJx$g<#jh^Xh;@ZVc{b?z&X1RtCW{>%)UhOriE6x?hk=eWK;K%cRd8CC}HFdmSs&7mubmMNqLiwU{lSBPfo34pK z$J)5_al+KCp@Q|j*);4QMB)!#h?h&}ol`TjcOY+CD1sYvBTz2eSN@_QXB2vw`VLm<0z0AzVCp)8-#A)*0 z*h0s%ug}%n+gtz0|~(3cPN;KGynzNCVt!n z^kkpSF=KqSJr%wGIyJ@%DS*C&lmJCsihYYpFg8cOJT~yzc|A=NkR)V@qh8A^jG(`N z;xhmFyI9$QL-Tz+E?ROwbyMftYpy&E=48c0)>bMd{Gm9&iA2bnqUO}q0 z?X_5NubZ}1V)HQ_Y<`zd0yH)A)CT1Y#_IOeDAaEi+?{obRxEim>7K1RJN()taoER) zj1EYJE-^qtmQyO2$UhER-(91@>tUZ_A*kI;QkInPdfb>S<*e^!H_DlDw$UQ+)t7VI zEsCqi*i|B$2ovtqxHgl5=y3R#jX=J0ddN}xMW0|$k}MsD4m z88W;`h-gm4mo16>STb$4HfzE*D0TGZn+7(H%S2|zr$lsc@`~SwV4ow`?nyjE1SKT6 z{+p0#W@MX%^C}0P+2xI~a%FMHhX`^ee5Xk8WcR2jC=zIk&PA*MLm?%zj2>V4?p%R) z|0P2WcXw{sBbXIRtKo7twLXvc zfGaRU;6o9mEsR%tOW>0~-s|vhfze?O;&2p?iKKJurqV9xHwgm&nlr;0?W~VGN!S2&x$FDQ zv?g?yn5$BVy1diQ(^mc-PR9e%Vhs3hBKk5Gt?>*HZiuO{A69|T1uokJN(u|4cJrg1YAjj};P8wHuIA}%%pYmOHZA8$bt!w{qf>M!jb9UU$k zbH0$}N}h@b!`-6v!i5X`{$IX)(b^K9$QY5bRI;HPzJhDkITPixJ~d=%-a8XM*~em> zg3(_3ltG4#enw#YW&S4@Uf|iQ6chh2<7<3nKX}iB znm&B8AW1_L{Xq@tAmy)IxvG+-ygJc7qxU)$Uuhpde_yN@N<)HX^=^A^?V;Xw38#;YJ^$b6BQtKxk`%WdzX3=t1Nw$zU=(=r2^I`AEg=1h zKmDgfVu9jSr|#p91~*ESRJ9{T%u=EChoDubG#F7f&wqwpXDG+Y@Ei>1U*k4eYb;}Mu{0)OArU*_HS;0M z;v$$0(mF}+L4-C#{#UxA^jZ27DunpsVD1Q-x!;wb?Vg2j1Cw(G%U4EG_ufP zx6u%qN%4{v4o?k&4n9?N1U3eGplc<+sG($Ue}4!!lOca!S$zyle=zXrQ+uT3US=B> zi>>h7u*lE1Za0V0v{z3QP%fK42H!d!3zBAeO=1NYdV_vPTe{AB&`&BrbWH+jJ~tzE z5}x=N(nt|C-#C3QU#A?GsWd!>CA0_2x40g>s|xU>f1JHUZH0wL9hN7-o}Qj=hl?hyV!%2~Qs*f$tnMB`Q{Wm6JVa4- z(N`@n!UYE%bE?!K3Aj$T3;NQaex-vw^rysSP#I@X34G>u@FdJu6@4aHR~4|MSVPF| zBrHEUNBiM5l(W;phlgM6kP=%lZ8m`Y`DTBQ1vC#24{t;Aby)n_@dM*()HlPGq4N9! z5DF_GFP6B7b89swxB^_R zA1^PnKiXJZE5;=aDj)=709o5MEpZ4pRe%A~yVenL_IH}kddTER?Kbb!Xb42^sFCc0%?hyQhTpq^Eq+*y5`cI z&spgg;881}4-*8qN0dMgrcb)~DXd-uHqbJF~V>Ud%s zcmRx~%xyzVSp}5uuhF}^yUiDP>K^pJhjSEz%XC9>lJUw7?GH6J7O_s=m*QvitQs^w zCx2LNJTb;a3v_y>dK!EW5}m66Wz9ngc^+*{VN4yQ{0^%;)4T#{#N6WLwJwlz#!m^2 z3d4~JZ@!vmTweno|Chsh_nsrLD%yXG@`|#iSe=SJCN8%h78;~OiV zLM4z=xl(2QiB|~|NuedA98UB2lOTVtuJS;B|7OOkGGofR7tfz3l2}kC&-uTlOunyf zG9L3hbRm333n{^(15R?XwR)$tv4IsND_H>>@qbb>Gt1l~={^}L$X>-je~(9t(iAwa zdY6F@G{KIe*knH-E!A@Paai|BBWn63!KT2Q?k=s=-yLfZeTX{Z?!Gt6zn^BR({%>% zk&w#`T1N~7R>W#Vx6T>WHsZQ!pm=_PUg`Hk=Ml(rKpBa&(FBT>MHnRV}2KYaoN&B~lMFKcgUr-ks)%Gvg zAe7F^;0IhBD78eA(a+6HPMVVMcW1SOCxy+&Gh=EcJN1zi7eIm}qmr6hXo565GxOKe zEUOI%+(q~i>Yb$UdpLtpKo?aZ4N_)1paHYa92~is3KRdMdo?cejJU84WG7&*-NThi z$jUgx3O$du2GlR8{QCT?t@wKwP11yE7zO`#H6ve0XSw(~AT(dlpAMLYNb3y%td~tp zOq`8s&0Sv30~bpg#vvPA=p*R)bJm)1TrLQPHZH6VL**q-#|)EZWk`8h<9~8oV%5bs zSYe-}4Nl{iRH}d}Ueg;0Rg|pc-V_z%WMi|%AT`u1{t z%zITtU)H4@NqA#vp}&9wH^J3X#Rm-sR#IRDV?b1yf`>R+rMvpCAX1+s6C==8LuI*; zVs3hG&g;AHvbMeLcAB2jP%ZgZ*J&28?qlZWtjOM34XDZPR9IQ*Vj;{;woD~p6*0d* zSnpkVnO&3lk5r0)UU;}KfTfH|{%R|Taf!I+Qe1L<4%eDyD*O*8>G&W`^QX;9-Ye?y z`%Rfu*XPqr1oO>_L1ojogdseM@DVk1g&#tiI4=!gg`(Ti z4+aPo#5#=DB+zF;co$T)5CNK6z#urw&2v60&J2KqLy=fO89Z|MlLTj+8yg$*0l~rr zJu%~*W__kjVVLpbZDE}Km>7OE;@p9OI>>_VT;jE5Q`t`?JroHA#AKzRgbRFLzpBnC z0S$USuuGhnh9={S2P7T4eGVZU7GFq7W1hjqgVVST7KuK(hK7Nyj(+pkGmf~k_4k92 g|9Af%eC3$vvJKvh>fJOYpv4Gfc@4Qj*$09D3+EaLv;Y7A literal 0 HcmV?d00001 diff --git a/mllam_verification/operations/statistics.py b/mllam_verification/operations/statistics.py index eb9e743..96f590e 100644 --- a/mllam_verification/operations/statistics.py +++ b/mllam_verification/operations/statistics.py @@ -1,7 +1,8 @@ from types import FunctionType -from typing import List, Mapping, Optional, Tuple +from typing import List, Mapping, Optional, Tuple, Union import numpy as np +import scores.categorical as scc_cat import scores.continuous as scc_cont import scores.probability as scc_prob import xarray as xr @@ -350,6 +351,274 @@ def spread_skill_ratio( return ds_ssr +def brier_score( + ds_reference: xr.Dataset | xr.DataArray, + ds_prediction: xr.Dataset | xr.DataArray, + ensemble_member_dim: str = "ensemble_member", + thresholds: Union[float, List[float]] = 0.5, + **stats_op_kwargs, +) -> xr.Dataset | xr.DataArray: + """Compute the Brier Score for ensemble forecasts. + + Evaluates the accuracy of probabilistic predictions for binary + events defined by threshold exceedance. The ensemble forecast is + converted to a probability (fraction of members exceeding the + threshold) and compared against the observed binary outcome. + + A Brier Score of 0 indicates a perfect forecast; 1 indicates the + worst possible forecast. Lower is better. + + Args: + ds_reference: Reference (observation) dataset or data array. + Must NOT contain ``ensemble_member_dim``. + ds_prediction: Ensemble forecast dataset or data array. + Must contain ``ensemble_member_dim`` as a dimension. + ensemble_member_dim: Name of the ensemble member dimension + in ``ds_prediction``. Defaults to ``"ensemble_member"``. + thresholds: Event threshold(s). A single float or a list of + floats defining the exceedance events to evaluate. + Defaults to ``0.5``. + **stats_op_kwargs: Additional keyword arguments forwarded to + ``scores.probability.brier_score_for_ensemble``. + + Returns: + Dataset or DataArray with Brier Score values. + + References: + Ferro, C. A. T. (2013). Fair scores for ensemble forecasts. + Quarterly Journal of the Royal Meteorological Society, + 140(683), 1917-1923. https://doi.org/10.1002/qj.2270 + + Example: + >>> da_bs = brier_score( + ... da_reference, + ... da_ensemble_prediction, + ... ensemble_member_dim="ensemble_member", + ... thresholds=[0.5, 1.0, 2.0], + ... ) + """ + groupby = stats_op_kwargs.pop("groupby", None) + stats_op_kwargs["ensemble_member_dim"] = ensemble_member_dim + stats_op_kwargs["event_thresholds"] = thresholds + + ds_bs = compute_pipeline_statistic( + datasets=[ds_prediction, ds_reference], + stats_op=scc_prob.brier_score_for_ensemble, + stats_op_kwargs=stats_op_kwargs, + groupby=groupby, + ) + + if isinstance(ds_bs, (xr.DataArray, xr.Dataset)): + ds_bs.name = getattr(ds_prediction, "name", "brier_score") + reduce_dims = list(set(ds_reference.dims) - set(ds_bs.dims)) + new_cell_methods = [",".join(reduce_dims) + ": brier_score"] + if isinstance(ds_bs, xr.DataArray): + update_cell_methods(ds_bs, new_cell_methods) + elif isinstance(ds_bs, xr.Dataset): + for _, da_var in ds_bs.items(): + update_cell_methods(da_var, new_cell_methods) + return ds_bs + + +def equitable_threat_score( + ds_reference: xr.Dataset | xr.DataArray, + ds_prediction: xr.Dataset | xr.DataArray, + threshold: float = 0.5, + reduce_dims: Optional[Union[str, List[str]]] = "all", + **stats_op_kwargs, +) -> xr.Dataset | xr.DataArray: + """Compute the Equitable Threat Score (ETS) for categorical forecasts. + + Also known as Gilbert Skill Score. Evaluates how well the forecast + "yes" events correspond to the observed "yes" events, accounting + for hits due to chance. + + Binary events are defined by exceedance of ``threshold``: + a grid point is considered a "hit" if both forecast and observation + exceed the threshold. + + Args: + ds_reference: Reference (observation) dataset or data array. + ds_prediction: Deterministic forecast dataset or data array. + threshold: Event threshold. Grid points where the value exceeds + this threshold are classified as events. Defaults to ``0.5``. + reduce_dims: Dimensions to reduce when computing the + contingency table. Defaults to ``"all"`` (scalar output). + Pass a list of dimension names to preserve other dimensions. + **stats_op_kwargs: Additional keyword arguments. + + Returns: + Dataset or DataArray with ETS values. Range: -1/3 to 1. + 0 indicates no skill; 1 indicates a perfect score. + + References: + Gilbert, G.K., 1884. Finley's tornado predictions. + American Meteorological Journal, 1(5), pp.166-172. + + Hogan, R.J. et al., 2010. Equitability revisited: Why the + "equitable threat score" is not equitable. Weather and + Forecasting, 25(2), pp.710-726. + https://doi.org/10.1175/2009WAF2222350.1 + + Example: + >>> da_ets = equitable_threat_score( + ... da_reference, + ... da_prediction, + ... threshold=0.5, + ... reduce_dims=["x", "y"], + ... ) + """ + groupby = stats_op_kwargs.pop("groupby", None) + + # Create binary events based on threshold exceedance + obs_events = ds_reference >= threshold + fcst_events = ds_prediction >= threshold + + # Build contingency table using scores library + bcm = scc_cat.BinaryContingencyManager( + fcst_events=fcst_events, + obs_events=obs_events, + ) + + # Reduce dimensions to compute the contingency counts + basic_cm = bcm.transform(reduce_dims=reduce_dims) + + # Calculate ETS from the contingency table + ds_ets = basic_cm.equitable_threat_score() + ds_ets.name = getattr(ds_prediction, "name", "equitable_threat_score") + + if groupby: + ds_ets = ds_ets.groupby(groupby) + + new_cell_methods = [f"threshold({threshold}): equitable_threat_score"] + if isinstance(ds_ets, xr.DataArray): + update_cell_methods(ds_ets, new_cell_methods) + elif isinstance(ds_ets, xr.Dataset): + for _, da_var in ds_ets.items(): + update_cell_methods(da_var, new_cell_methods) + + return ds_ets + + +def fractions_skill_score( + ds_reference: xr.Dataset | xr.DataArray, + ds_prediction: xr.Dataset | xr.DataArray, + threshold: float, + window_size: int, + spatial_dims: Optional[List[str]] = None, + **stats_op_kwargs, +) -> xr.Dataset | xr.DataArray: + """Compute the Fractions Skill Score (FSS) for spatial forecasts. + + FSS evaluates the spatial accuracy of a forecast by comparing + fractional coverage of events (values exceeding a threshold) within + spatial neighborhoods of a given size. It addresses the "double + penalty" problem inherent in point-wise verification of + high-resolution models. + + FSS = 1 - MSE / MSE_ref + + where MSE is the mean squared difference of fractional coverages, + and MSE_ref is the worst-case reference (no skill). + + Args: + ds_reference: Reference (observation) dataset or data array. + Must contain spatial dimensions. + ds_prediction: Forecast dataset or data array. Must contain + spatial dimensions matching ``ds_reference``. + threshold: Event threshold. Grid points where the value + exceeds this threshold are classified as events. + window_size: Size of the spatial neighborhood window (must + be an odd integer). The window is applied as a square + ``window_size x window_size`` rolling mean over the + spatial dimensions. + spatial_dims: Names of the two spatial dimensions. + Defaults to ``["x", "y"]``. Compatible with + ``mlwp-data-specs`` variants such as ``["xc", "yc"]`` + or ``["longitude", "latitude"]``. + **stats_op_kwargs: Additional keyword arguments. + + Returns: + Dataset or DataArray with FSS value(s). Range: 0 to 1. + FSS = 1 indicates a perfect forecast; FSS = 0 indicates no + skill. FSS > 0.5 is generally considered "useful" at the + given spatial scale. + + References: + Roberts, N.M. and Lean, H.W. (2008). Scale-Selective + Verification of Rainfall Accumulations from High-Resolution + Forecasts of Convective Events. Monthly Weather Review, + 136(1), pp.78-97. https://doi.org/10.1175/2007MWR2123.1 + + Example: + >>> da_fss = fractions_skill_score( + ... da_reference, + ... da_prediction, + ... threshold=1.0, + ... window_size=11, + ... spatial_dims=["x", "y"], + ... ) + """ + if spatial_dims is None: + spatial_dims = ["x", "y"] + + if len(spatial_dims) != 2: + raise ValueError( + f"spatial_dims must have exactly 2 elements, got {len(spatial_dims)}" + ) + + if window_size % 2 == 0: + raise ValueError(f"window_size must be odd, got {window_size}") + + groupby = stats_op_kwargs.pop("groupby", None) + + # Convert to binary fields based on threshold exceedance + obs_binary = (ds_reference >= threshold).astype(float) + fcst_binary = (ds_prediction >= threshold).astype(float) + + # Compute fractional coverage using rolling mean over spatial dims + dim_x, dim_y = spatial_dims + obs_frac = ( + obs_binary.rolling({dim_x: window_size, dim_y: window_size}, center=True) + .mean() + .dropna(dim=dim_x) + .dropna(dim=dim_y) + ) + fcst_frac = ( + fcst_binary.rolling({dim_x: window_size, dim_y: window_size}, center=True) + .mean() + .dropna(dim=dim_x) + .dropna(dim=dim_y) + ) + + # MSE of fractional coverages + mse = ((obs_frac - fcst_frac) ** 2).mean(dim=spatial_dims) + + # Worst-case MSE reference (no overlap between obs and fcst fractions) + mse_ref = (obs_frac**2 + fcst_frac**2).mean(dim=spatial_dims) + + # FSS = 1 - MSE / MSE_ref + # When mse_ref == 0, both fields have zero fractional coverage + # (or are identical), so the forecast is trivially perfect → FSS = 1.0 + ds_fss = xr.where(mse_ref > 0, 1.0 - mse / mse_ref, 1.0) + ds_fss.name = getattr(ds_prediction, "name", "fractions_skill_score") + + if groupby: + ds_fss = ds_fss.groupby(groupby) + + new_cell_methods = [ + f"{','.join(spatial_dims)}: fractions_skill_score" + f" (threshold: {threshold}, window: {window_size})" + ] + if isinstance(ds_fss, xr.DataArray): + update_cell_methods(ds_fss, new_cell_methods) + elif isinstance(ds_fss, xr.Dataset): + for _, da_var in ds_fss.items(): + update_cell_methods(da_var, new_cell_methods) + + return ds_fss + + def mean(ds: xr.Dataset | xr.DataArray, **stats_op_kwargs) -> xr.Dataset | xr.DataArray: """Compute the mean across specified dimensions. diff --git a/mllam_verification/plot.py b/mllam_verification/plot.py index 2ad75ae..d720a9f 100644 --- a/mllam_verification/plot.py +++ b/mllam_verification/plot.py @@ -606,3 +606,96 @@ def plot_single_metric_hovmoller( # noqa: C901 ) return axes + + +def plot_fss_scale( + da_reference: xr.DataArray, + da_prediction: xr.DataArray, + threshold: float, + window_sizes: list[int], + spatial_dims: Optional[list[str]] = None, + axes: Optional[plt.Axes] = None, + hue: Optional[str] = "datasource", + **stats_op_kwargs, +) -> plt.Axes: + """Plot Fractions Skill Score (FSS) across multiple spatial scales. + + Computes the FSS for a given threshold across a list of window + sizes and plots the resulting score vs. scale curve. A horizontal + line at FSS = 0.5 is added to indicate the standard threshold for + a "useful" forecast. + + Args: + da_reference: Reference dataset/dataarray containing observations. + da_prediction: Forecast dataset/dataarray. + threshold: The threshold to define binary events. + window_sizes: List of odd integers defining the neighborhood sizes. + spatial_dims: Names of the spatial dimensions (e.g. ["x", "y"]). + axes: Pre-existing matplotlib axes. If None, creates a new figure. + hue: Dimension to use for coloring multiple lines (e.g. "datasource"). + **stats_op_kwargs: Additional arguments passed to fractions_skill_score + (like `groupby`). + + Returns: + The matplotlib axes containing the plot. + """ + if spatial_dims is None: + spatial_dims = ["x", "y"] + + if axes is None: + _, axes = plt.subplots(figsize=(8, 6)) + + # Calculate FSS for each window size + fss_values = [] + for w in window_sizes: + if w % 2 == 0: + raise ValueError(f"window_sizes must be odd integers, got {w}") + + ds_fss = mlverif_stats.fractions_skill_score( + da_reference, + da_prediction, + threshold=threshold, + window_size=w, + spatial_dims=spatial_dims, + **stats_op_kwargs, + ) + # Assign window size coordinate for concatenation + ds_fss = ds_fss.assign_coords(window_size=w) + fss_values.append(ds_fss) + + # Combine into a single DataArray + da_fss_scale = xr.concat(fss_values, dim="window_size") + + # Plot + if hue in da_fss_scale.dims or hue in da_fss_scale.coords: + for hue_val in da_fss_scale[hue].values: + da_sub = da_fss_scale.sel({hue: hue_val}) + axes.plot( + da_sub["window_size"], + da_sub.values, + marker="o", + label=str(hue_val), + ) + axes.legend(title=hue) + else: + axes.plot( + da_fss_scale["window_size"], + da_fss_scale.values, + marker="o", + color="blue", + ) + + # Standard FSS plot formatting + axes.axhline(y=0.5, color="k", linestyle="--", alpha=0.7, label="Useful (0.5)") + axes.set_ylim(0, 1.05) + axes.set_xlabel("Neighborhood Size (grid points)") + axes.set_ylabel("Fractions Skill Score") + axes.set_title(f"FSS vs Scale (Threshold: {threshold})") + axes.grid(True, linestyle=":", alpha=0.6) + + # Avoid duplicate labels in legend if axhline was drawn after hue loop + handles, labels = axes.get_legend_handles_labels() + by_label = dict(zip(labels, handles)) + axes.legend(by_label.values(), by_label.keys()) + + return axes diff --git a/tests/unit/test_plot.py b/tests/unit/test_plot.py index da18885..c129604 100644 --- a/tests/unit/test_plot.py +++ b/tests/unit/test_plot.py @@ -424,3 +424,63 @@ def test_accepts_existing_axes( ) assert returned_ax is ax plt.close("all") + + +class TestPlotFssScale: + """Unit tests for the plot_fss_scale function.""" + + def test_returns_axes( + self, + da_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """plot_fss_scale() should return a matplotlib Axes object.""" + from mllam_verification.plot import plot_fss_scale + + axes = plot_fss_scale( + da_reference_2d_utc, + da_prediction_2d_utc, + threshold=0.5, + window_sizes=[3, 5, 7], + spatial_dims=["x", "y"], + ) + + assert isinstance(axes, plt.Axes) + plt.close("all") + + def test_accepts_existing_axes( + self, + da_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """plot_fss_scale() should use provided axes.""" + from mllam_verification.plot import plot_fss_scale + + fig, ax = plt.subplots() + returned_ax = plot_fss_scale( + da_reference_2d_utc, + da_prediction_2d_utc, + threshold=0.5, + window_sizes=[3, 5], + spatial_dims=["x", "y"], + axes=ax, + ) + assert returned_ax is ax + plt.close("all") + + def test_rejects_even_window_sizes( + self, + da_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """plot_fss_scale() should raise ValueError for even window sizes.""" + from mllam_verification.plot import plot_fss_scale + + with pytest.raises(ValueError, match="window_sizes must be odd integers"): + plot_fss_scale( + da_reference_2d_utc, + da_prediction_2d_utc, + threshold=0.5, + window_sizes=[2, 4], + spatial_dims=["x", "y"], + ) diff --git a/tests/unit/test_statistics.py b/tests/unit/test_statistics.py index ae718c3..823e124 100644 --- a/tests/unit/test_statistics.py +++ b/tests/unit/test_statistics.py @@ -162,3 +162,307 @@ def test_ssr_perfect_ensemble_near_one(self): ) # For a large well-calibrated ensemble SSR should be near 1.0 assert 0.5 < float(result) < 2.0 + + +class TestBrierScore: + """Tests for the brier_score() function.""" + + def test_brier_score_returns_dataarray( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """brier_score() should return a DataArray.""" + from mllam_verification.operations.statistics import brier_score + + result = brier_score( + da_reference_2d_utc, + da_ensemble_prediction_2d_utc, + ensemble_member_dim="ensemble_member", + thresholds=0.5, + reduce_dims=["x", "y"], + ) + assert isinstance(result, xr.DataArray) + + def test_brier_score_ensemble_dim_collapsed( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """Ensemble member dimension must not appear in output.""" + from mllam_verification.operations.statistics import brier_score + + result = brier_score( + da_reference_2d_utc, + da_ensemble_prediction_2d_utc, + ensemble_member_dim="ensemble_member", + thresholds=0.5, + reduce_dims=["x", "y"], + ) + assert "ensemble_member" not in result.dims + + def test_brier_score_has_cell_methods( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """Output must have cell_methods attribute.""" + from mllam_verification.operations.statistics import brier_score + + result = brier_score( + da_reference_2d_utc, + da_ensemble_prediction_2d_utc, + ensemble_member_dim="ensemble_member", + thresholds=0.5, + reduce_dims=["x", "y"], + ) + assert "cell_methods" in result.attrs + + def test_brier_score_range( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """Brier Score must be in range [0, 1].""" + from mllam_verification.operations.statistics import brier_score + + result = brier_score( + da_reference_2d_utc, + da_ensemble_prediction_2d_utc, + ensemble_member_dim="ensemble_member", + thresholds=0.5, + reduce_dims=["x", "y"], + ) + assert float(result.min()) >= 0 + assert float(result.max()) <= 1 + + def test_brier_score_multiple_thresholds( + self, + da_ensemble_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """Brier Score should work with multiple thresholds.""" + from mllam_verification.operations.statistics import brier_score + + result = brier_score( + da_reference_2d_utc, + da_ensemble_prediction_2d_utc, + ensemble_member_dim="ensemble_member", + thresholds=[0.1, 0.5, 0.9], + reduce_dims=["x", "y"], + ) + assert isinstance(result, xr.DataArray) + assert "threshold" in result.dims + + +class TestEquitableThreatScore: + """Tests for the equitable_threat_score() function.""" + + def test_ets_returns_dataarray( + self, + da_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """equitable_threat_score() should return a DataArray.""" + from mllam_verification.operations.statistics import equitable_threat_score + + result = equitable_threat_score( + da_reference_2d_utc, + da_prediction_2d_utc, + threshold=0.5, + reduce_dims="all", + ) + assert isinstance(result, xr.DataArray) + + def test_ets_has_cell_methods( + self, + da_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """Output must have cell_methods attribute.""" + from mllam_verification.operations.statistics import equitable_threat_score + + result = equitable_threat_score( + da_reference_2d_utc, + da_prediction_2d_utc, + threshold=0.5, + reduce_dims="all", + ) + assert "cell_methods" in result.attrs + + def test_ets_range( + self, + da_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """ETS must be in range [-1/3, 1].""" + from mllam_verification.operations.statistics import equitable_threat_score + + result = equitable_threat_score( + da_reference_2d_utc, + da_prediction_2d_utc, + threshold=0.5, + reduce_dims="all", + ) + assert float(result) >= -1 / 3 + assert float(result) <= 1 + + def test_ets_perfect_forecast_is_one(self): + """A perfect forecast should have ETS = 1.0.""" + import numpy as np + + from mllam_verification.operations.statistics import equitable_threat_score + + # Create identical observation and forecast + data = np.array([0.0, 0.0, 1.0, 1.0, 1.0]) + obs = xr.DataArray(data, dims=["grid_index"]) + fcst = xr.DataArray(data, dims=["grid_index"]) + + result = equitable_threat_score( + obs, + fcst, + threshold=0.5, + reduce_dims="all", + ) + assert float(result) == 1.0 + + def test_ets_no_skill_near_zero(self): + """A random forecast should have ETS near 0.""" + import numpy as np + + from mllam_verification.operations.statistics import equitable_threat_score + + rng = np.random.default_rng(seed=42) + obs = xr.DataArray(rng.choice([0.0, 1.0], size=1000), dims=["grid_index"]) + fcst = xr.DataArray(rng.choice([0.0, 1.0], size=1000), dims=["grid_index"]) + + result = equitable_threat_score( + obs, + fcst, + threshold=0.5, + reduce_dims="all", + ) + # Random forecast should have ETS near 0 (within ±0.15) + assert -0.15 < float(result) < 0.15 + + +class TestFractionsSkillScore: + """Tests for the fractions_skill_score() function.""" + + def test_fss_returns_dataarray( + self, + da_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """fractions_skill_score() should return a DataArray.""" + from mllam_verification.operations.statistics import fractions_skill_score + + result = fractions_skill_score( + da_reference_2d_utc, + da_prediction_2d_utc, + threshold=0.5, + window_size=5, + spatial_dims=["x", "y"], + ) + assert isinstance(result, xr.DataArray) + + def test_fss_has_cell_methods( + self, + da_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """Output must have cell_methods attribute.""" + from mllam_verification.operations.statistics import fractions_skill_score + + result = fractions_skill_score( + da_reference_2d_utc, + da_prediction_2d_utc, + threshold=0.5, + window_size=5, + spatial_dims=["x", "y"], + ) + assert "cell_methods" in result.attrs + + def test_fss_range( + self, + da_prediction_2d_utc: xr.DataArray, + da_reference_2d_utc: xr.DataArray, + ): + """FSS must be in range [0, 1].""" + from mllam_verification.operations.statistics import fractions_skill_score + + result = fractions_skill_score( + da_reference_2d_utc, + da_prediction_2d_utc, + threshold=0.5, + window_size=5, + spatial_dims=["x", "y"], + ) + assert float(result.min()) >= 0 + assert float(result.max()) <= 1 + + def test_fss_perfect_forecast_is_one(self): + """A perfect forecast should have FSS = 1.0.""" + import numpy as np + + from mllam_verification.operations.statistics import fractions_skill_score + + data = np.random.rand(20, 20) + obs = xr.DataArray(data, dims=["x", "y"]) + fcst = xr.DataArray(data, dims=["x", "y"]) + + result = fractions_skill_score( + obs, fcst, threshold=0.5, window_size=3, spatial_dims=["x", "y"] + ) + assert float(result) == 1.0 + + def test_fss_increases_with_window_size(self): + """FSS should generally increase with larger window sizes.""" + import numpy as np + + from mllam_verification.operations.statistics import fractions_skill_score + + rng = np.random.default_rng(seed=42) + obs = xr.DataArray(rng.random((50, 50)), dims=["x", "y"]) + # Spatially shifted forecast + fcst = xr.DataArray(np.roll(obs.values, 3, axis=0), dims=["x", "y"]) + + fss_small = fractions_skill_score( + obs, fcst, threshold=0.5, window_size=3, spatial_dims=["x", "y"] + ) + fss_large = fractions_skill_score( + obs, fcst, threshold=0.5, window_size=11, spatial_dims=["x", "y"] + ) + assert float(fss_large) >= float(fss_small) + + def test_fss_rejects_even_window(self): + """Even window_size should raise ValueError.""" + import numpy as np + import pytest + + from mllam_verification.operations.statistics import fractions_skill_score + + data = np.random.rand(10, 10) + obs = xr.DataArray(data, dims=["x", "y"]) + fcst = xr.DataArray(data, dims=["x", "y"]) + + with pytest.raises(ValueError, match="window_size must be odd"): + fractions_skill_score( + obs, fcst, threshold=0.5, window_size=4, spatial_dims=["x", "y"] + ) + + def test_fss_rejects_wrong_spatial_dims(self): + """spatial_dims with != 2 elements should raise ValueError.""" + import numpy as np + import pytest + + from mllam_verification.operations.statistics import fractions_skill_score + + data = np.random.rand(10, 10) + obs = xr.DataArray(data, dims=["x", "y"]) + fcst = xr.DataArray(data, dims=["x", "y"]) + + with pytest.raises(ValueError, match="spatial_dims must have exactly 2"): + fractions_skill_score( + obs, fcst, threshold=0.5, window_size=3, spatial_dims=["x"] + ) From ea188e6aeb9b5d227dd03e8e4d6cc30ab606e9b1 Mon Sep 17 00:00:00 2001 From: GiGiKoneti Date: Thu, 25 Jun 2026 01:10:39 +0530 Subject: [PATCH 10/10] ci: run pre-commit directly to bypass cache service outages --- .github/workflows/ci-pre-commit.yml | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/.github/workflows/ci-pre-commit.yml b/.github/workflows/ci-pre-commit.yml index 542bd8b..0fac1ad 100644 --- a/.github/workflows/ci-pre-commit.yml +++ b/.github/workflows/ci-pre-commit.yml @@ -12,10 +12,15 @@ jobs: runs-on: ubuntu-latest steps: - - uses: actions/checkout@v2 - - uses: actions/setup-python@v2 + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 with: # don't use python3.12 because flake8 finds extra issues with that # version python-version: "3.11" - - uses: pre-commit/action@v2.0.3 + + - name: Install dependencies + run: python -m pip install --upgrade pip && pip install pre-commit + + - name: Run pre-commit hooks + run: pre-commit run --all-files