From de3a6ca3bea718ed4d699b0cf90f346d455db40d Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Sun, 29 Mar 2026 20:38:36 +0000 Subject: [PATCH] Fix crps_ens N=1 sum_vars handling and add 4D ensemble_member DataArray export bridge --- neural_lam/metrics.py | 69 +++++++++++++++++++++++++++-- neural_lam/models/ar_model.py | 46 +++++++++++++++++--- tests/test_metrics.py | 82 +++++++++++++++++++++++++++++++++++ 3 files changed, 188 insertions(+), 9 deletions(-) create mode 100644 tests/test_metrics.py diff --git a/neural_lam/metrics.py b/neural_lam/metrics.py index 7db2cca6d..059bfb320 100644 --- a/neural_lam/metrics.py +++ b/neural_lam/metrics.py @@ -157,10 +157,14 @@ def mae(pred, target, pred_std, mask=None, average_grid=True, sum_vars=True): metric_val: One of (...,), (..., d_state), (..., N), (..., N, d_state), depending on reduction arguments. """ - # Replace pred_std with constant ones - return wmae( - pred, target, torch.ones_like(pred_std), mask, average_grid, sum_vars - ) + # Replace pred_std with constant ones. Allow pred_std to be omitted for + # call sites that only need unweighted MAE behavior. + if pred_std is None: + pred_std_ones = torch.ones_like(pred) + else: + pred_std_ones = torch.ones_like(pred_std) + + return wmae(pred, target, pred_std_ones, mask, average_grid, sum_vars) def nll(pred, target, pred_std, mask=None, average_grid=True, sum_vars=True): @@ -227,6 +231,62 @@ def crps_gauss( ) +def crps_ens( + pred, + target, + pred_std, + mask=None, + average_grid=True, + sum_vars=True, + ens_dim=-3, +): + """ + Continuous Ranked Probability Score (CRPS) for finite ensembles. + + (...,) is any number of batch dimensions, potentially different + but broadcastable + pred: (..., N_ens, N, d_state), prediction with ensemble dimension at + ens_dim + target: (..., N, d_state), target + pred_std: unused, kept for metric API compatibility + mask: (N,), boolean mask describing which grid nodes to use in metric + average_grid: boolean, if grid dimension -2 should be reduced (mean over N) + sum_vars: boolean, if variable dimension -1 should be reduced (sum + over d_state) + ens_dim: int, index of ensemble member dimension in pred + + Returns: + metric_val: One of (...,), (..., d_state), (..., N), (..., N, d_state), + depending on reduction arguments. + """ + if pred.shape[ens_dim] == 1: + return mae( + pred.squeeze(ens_dim), + target, + None, + mask=mask, + average_grid=average_grid, + sum_vars=sum_vars, + ) + + pred_ens_first = pred.movedim(ens_dim, 0) + + first_term = torch.mean( + torch.abs(pred_ens_first - target.unsqueeze(0)), dim=0 + ) + + pairwise_abs_diff = torch.abs( + pred_ens_first.unsqueeze(0) - pred_ens_first.unsqueeze(1) + ) + second_term = 0.5 * torch.mean(pairwise_abs_diff, dim=(0, 1)) + + entry_crps = first_term - second_term + + return mask_and_reduce_metric( + entry_crps, mask=mask, average_grid=average_grid, sum_vars=sum_vars + ) + + DEFINED_METRICS = { "mse": mse, "mae": mae, @@ -234,4 +294,5 @@ def crps_gauss( "wmae": wmae, "nll": nll, "crps_gauss": crps_gauss, + "crps_ens": crps_ens, } diff --git a/neural_lam/models/ar_model.py b/neural_lam/models/ar_model.py index a411a3afc..b55497298 100644 --- a/neural_lam/models/ar_model.py +++ b/neural_lam/models/ar_model.py @@ -177,8 +177,9 @@ def _create_dataarray_from_tensor( ---------- tensor : torch.Tensor The tensor to convert to a `xr.DataArray` with dimensions [time, - grid_index, feature]. The tensor will be copied to the CPU if it is - not already there. + grid_index, feature] for deterministic outputs, or [time, + grid_index, ensemble_member, feature] for probabilistic outputs. + The tensor will be copied to the CPU if it is not already there. time : torch.Tensor The time index or indices for the data, given as tensor representing epoch time in nanoseconds. The tensor will be @@ -193,10 +194,45 @@ def _create_dataarray_from_tensor( # provided to ARModel or where to put plotting still needs discussion weather_dataset = WeatherDataset(datastore=self._datastore, split=split) time = np.array(time.cpu(), dtype="datetime64[ns]") - da = weather_dataset.create_dataarray_from_tensor( - tensor=tensor, time=time, category=category + + if tensor.ndim == 4: + da_datastore_category = getattr(weather_dataset, f"da_{category}") + feature_dim_name = f"{category}_feature" + + da = xr.DataArray( + tensor.cpu().numpy(), + dims=( + "time", + "grid_index", + "ensemble_member", + feature_dim_name, + ), + coords={ + "time": time, + "grid_index": da_datastore_category.grid_index, + "ensemble_member": np.arange(tensor.shape[2]), + feature_dim_name: da_datastore_category[feature_dim_name], + }, + ) + + for grid_coord in ["x", "y"]: + if ( + grid_coord in da_datastore_category.coords + and grid_coord not in da.coords + ): + da.coords[grid_coord] = da_datastore_category[grid_coord] + + return da + + if tensor.ndim in (2, 3): + return weather_dataset.create_dataarray_from_tensor( + tensor=tensor, time=time, category=category + ) + + raise ValueError( + "Expected tensor to have 2, 3 or 4 dimensions, " + f"but got {tensor.ndim}." ) - return da def configure_optimizers(self): opt = torch.optim.AdamW( diff --git a/tests/test_metrics.py b/tests/test_metrics.py new file mode 100644 index 000000000..80fa42f60 --- /dev/null +++ b/tests/test_metrics.py @@ -0,0 +1,82 @@ +# Third-party +import numpy as np +import torch + +# First-party +from neural_lam import metrics +from neural_lam.models.ar_model import ARModel +from tests.dummy_datastore import DummyDatastore + + +def test_crps_ens_single_member_shape_matches_multi_member_when_sum_vars_false(): + batch_size = 2 + n_grid_nodes = 4 + n_state_features = 3 + + target = torch.randn(batch_size, n_grid_nodes, n_state_features) + pred_single = torch.randn(batch_size, 1, n_grid_nodes, n_state_features) + pred_multi = torch.cat((pred_single, pred_single + 0.25), dim=1) + + metric_single = metrics.crps_ens( + pred_single, + target, + None, + average_grid=True, + sum_vars=False, + ens_dim=1, + ) + metric_multi = metrics.crps_ens( + pred_multi, + target, + None, + average_grid=True, + sum_vars=False, + ens_dim=1, + ) + + assert metric_single.shape == metric_multi.shape + assert metric_single.shape == (batch_size, n_state_features) + + +class _MinimalARModel: + def __init__(self, datastore): + self._datastore = datastore + + +def test_create_dataarray_from_tensor_supports_ensemble_member_dimension(): + datastore = DummyDatastore(n_grid_points=100, n_timesteps=8) + model = _MinimalARModel(datastore=datastore) + + create_da = ARModel._create_dataarray_from_tensor.__get__( + model, _MinimalARModel + ) + + n_time = 3 + n_ens = 2 + n_grid_nodes = datastore.num_grid_points + n_state_features = datastore.get_num_data_vars(category="state") + + tensor = torch.randn(n_time, n_grid_nodes, n_ens, n_state_features) + time = torch.tensor([0, 1, 2], dtype=torch.int64) + + da = create_da( + tensor=tensor, + time=time, + split="train", + category="state", + ) + + assert da.dims == ( + "time", + "grid_index", + "ensemble_member", + "state_feature", + ) + assert da.sizes["time"] == n_time + assert da.sizes["grid_index"] == n_grid_nodes + assert da.sizes["ensemble_member"] == n_ens + assert da.sizes["state_feature"] == n_state_features + np.testing.assert_array_equal( + da.coords["ensemble_member"].values, + np.arange(n_ens), + )