Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
69 changes: 65 additions & 4 deletions neural_lam/metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -227,11 +231,68 @@ 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,
"wmse": wmse,
"wmae": wmae,
"nll": nll,
"crps_gauss": crps_gauss,
"crps_ens": crps_ens,
}
46 changes: 41 additions & 5 deletions neural_lam/models/ar_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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(
Expand Down
82 changes: 82 additions & 0 deletions tests/test_metrics.py
Original file line number Diff line number Diff line change
@@ -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),
)