diff --git a/neural_lam/metrics.py b/neural_lam/metrics.py index 7db2cca6d..e9211136d 100644 --- a/neural_lam/metrics.py +++ b/neural_lam/metrics.py @@ -3,13 +3,16 @@ def get_metric(metric_name): - """ - Get a defined metric with given name + """Retrieve a metric function by name. - metric_name: str, name of the metric + Args: + metric_name (str): Name of the metric. Returns: - metric: function implementing the metric + Callable: Function implementing the selected metric. + + Raises: + AssertionError: If the metric name is not defined. """ metric_name_lower = metric_name.lower() assert ( @@ -19,62 +22,55 @@ def get_metric(metric_name): def mask_and_reduce_metric(metric_entry_vals, mask, average_grid, sum_vars): - """ - Masks and (optionally) reduces entry-wise metric values + """Apply masking and optional reductions to entry-wise metric values. - (...,) is any number of batch dimensions, potentially different - but broadcastable - metric_entry_vals: (..., N, d_state), prediction - 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) + The tensor shape (..., N, d_state) allows arbitrary batch dimensions + before the grid (N) and variable (d_state) dimensions. + + Args: + metric_entry_vals (torch.Tensor): Entry-wise metric values of shape + (..., N, d_state). + mask (torch.Tensor or None): Boolean mask of shape (N,) selecting + grid nodes to include. + average_grid (bool): If True, reduce the grid dimension (-2) + by taking the mean. + sum_vars (bool): If True, reduce the variable dimension (-1) + by taking the sum. Returns: - metric_val: One of (...,), (..., d_state), (..., N), (..., N, d_state), - depending on reduction arguments. + torch.Tensor: Reduced metric tensor depending on reduction settings. """ - # Only keep grid nodes in mask if mask is not None: - metric_entry_vals = metric_entry_vals[ - ..., mask, : - ] # (..., N', d_state) - - # Optionally reduce last two dimensions - if average_grid: # Reduce grid first - metric_entry_vals = torch.mean( - metric_entry_vals, dim=-2 - ) # (..., d_state) - if sum_vars: # Reduce vars second - metric_entry_vals = torch.sum( - metric_entry_vals, dim=-1 - ) # (..., N) or (...,) + metric_entry_vals = metric_entry_vals[..., mask, :] + + if average_grid: + metric_entry_vals = torch.mean(metric_entry_vals, dim=-2) + + if sum_vars: + metric_entry_vals = torch.sum(metric_entry_vals, dim=-1) return metric_entry_vals def wmse(pred, target, pred_std, mask=None, average_grid=True, sum_vars=True): - """ - Weighted Mean Squared Error - - (...,) is any number of batch dimensions, potentially different - but broadcastable - pred: (..., N, d_state), prediction - target: (..., N, d_state), target - pred_std: (..., N, d_state) or (d_state,), predicted std.-dev. - 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) + """Compute weighted mean squared error. + + Args: + pred (torch.Tensor): Predictions of shape (..., N, d_state). + target (torch.Tensor): Targets of shape (..., N, d_state). + pred_std (torch.Tensor): Predicted standard deviation of shape + (..., N, d_state) or (d_state,). + mask (torch.Tensor, optional): Boolean mask over grid nodes. + average_grid (bool): Whether to average over grid dimension. + sum_vars (bool): Whether to sum over variable dimension. Returns: - metric_val: One of (...,), (..., d_state), (..., N), (..., N, d_state), - depending on reduction arguments. + torch.Tensor: Weighted MSE depending on reductions. """ entry_mse = torch.nn.functional.mse_loss( pred, target, reduction="none" - ) # (..., N, d_state) - entry_mse_weighted = entry_mse / (pred_std**2) # (..., N, d_state) + ) + entry_mse_weighted = entry_mse / (pred_std**2) return mask_and_reduce_metric( entry_mse_weighted, @@ -85,51 +81,44 @@ def wmse(pred, target, pred_std, mask=None, average_grid=True, sum_vars=True): def mse(pred, target, pred_std, mask=None, average_grid=True, sum_vars=True): - """ - (Unweighted) Mean Squared Error - - (...,) is any number of batch dimensions, potentially different - but broadcastable - pred: (..., N, d_state), prediction - target: (..., N, d_state), target - pred_std: (..., N, d_state) or (d_state,), predicted std.-dev. - 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) + """Compute unweighted mean squared error. + + Equivalent to weighted MSE with unit standard deviation. + + Args: + pred (torch.Tensor): Predictions of shape (..., N, d_state). + target (torch.Tensor): Targets of shape (..., N, d_state). + pred_std (torch.Tensor): Dummy standard deviation tensor. + mask (torch.Tensor, optional): Boolean mask over grid nodes. + average_grid (bool): Whether to average over grid dimension. + sum_vars (bool): Whether to sum over variable dimension. Returns: - metric_val: One of (...,), (..., d_state), (..., N), (..., N, d_state), - depending on reduction arguments. + torch.Tensor: MSE depending on reductions. """ - # Replace pred_std with constant ones return wmse( pred, target, torch.ones_like(pred_std), mask, average_grid, sum_vars ) def wmae(pred, target, pred_std, mask=None, average_grid=True, sum_vars=True): - """ - Weighted Mean Absolute Error - - (...,) is any number of batch dimensions, potentially different - but broadcastable - pred: (..., N, d_state), prediction - target: (..., N, d_state), target - pred_std: (..., N, d_state) or (d_state,), predicted std.-dev. - 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) + """Compute weighted mean absolute error. + + Args: + pred (torch.Tensor): Predictions of shape (..., N, d_state). + target (torch.Tensor): Targets of shape (..., N, d_state). + pred_std (torch.Tensor): Predicted standard deviation. + mask (torch.Tensor, optional): Boolean mask over grid nodes. + average_grid (bool): Whether to average over grid dimension. + sum_vars (bool): Whether to sum over variable dimension. Returns: - metric_val: One of (...,), (..., d_state), (..., N), (..., N, d_state), - depending on reduction arguments. + torch.Tensor: Weighted MAE depending on reductions. """ entry_mae = torch.nn.functional.l1_loss( pred, target, reduction="none" - ) # (..., N, d_state) - entry_mae_weighted = entry_mae / pred_std # (..., N, d_state) + ) + entry_mae_weighted = entry_mae / pred_std return mask_and_reduce_metric( entry_mae_weighted, @@ -140,50 +129,42 @@ def wmae(pred, target, pred_std, mask=None, average_grid=True, sum_vars=True): def mae(pred, target, pred_std, mask=None, average_grid=True, sum_vars=True): - """ - (Unweighted) Mean Absolute Error - - (...,) is any number of batch dimensions, potentially different - but broadcastable - pred: (..., N, d_state), prediction - target: (..., N, d_state), target - pred_std: (..., N, d_state) or (d_state,), predicted std.-dev. - 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) + """Compute unweighted mean absolute error. + + Equivalent to weighted MAE with unit standard deviation. + + Args: + pred (torch.Tensor): Predictions of shape (..., N, d_state). + target (torch.Tensor): Targets of shape (..., N, d_state). + pred_std (torch.Tensor): Dummy standard deviation tensor. + mask (torch.Tensor, optional): Boolean mask over grid nodes. + average_grid (bool): Whether to average over grid dimension. + sum_vars (bool): Whether to sum over variable dimension. Returns: - metric_val: One of (...,), (..., d_state), (..., N), (..., N, d_state), - depending on reduction arguments. + torch.Tensor: MAE depending on reductions. """ - # Replace pred_std with constant ones return wmae( pred, target, torch.ones_like(pred_std), mask, average_grid, sum_vars ) def nll(pred, target, pred_std, mask=None, average_grid=True, sum_vars=True): - """ - Negative Log Likelihood loss, for isotropic Gaussian likelihood - - (...,) is any number of batch dimensions, potentially different - but broadcastable - pred: (..., N, d_state), prediction - target: (..., N, d_state), target - pred_std: (..., N, d_state) or (d_state,), predicted std.-dev. - 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) + """Compute negative log-likelihood for Gaussian predictions. + + Args: + pred (torch.Tensor): Predictions of shape (..., N, d_state). + target (torch.Tensor): Targets of shape (..., N, d_state). + pred_std (torch.Tensor): Predicted standard deviation. + mask (torch.Tensor, optional): Boolean mask over grid nodes. + average_grid (bool): Whether to average over grid dimension. + sum_vars (bool): Whether to sum over variable dimension. Returns: - metric_val: One of (...,), (..., d_state), (..., N), (..., N, d_state), - depending on reduction arguments. + torch.Tensor: NLL depending on reductions. """ - # Broadcast pred_std if shaped (d_state,), done internally in Normal class - dist = torch.distributions.Normal(pred, pred_std) # (..., N, d_state) - entry_nll = -dist.log_prob(target) # (..., N, d_state) + dist = torch.distributions.Normal(pred, pred_std) + entry_nll = -dist.log_prob(target) return mask_and_reduce_metric( entry_nll, mask=mask, average_grid=average_grid, sum_vars=sum_vars @@ -193,34 +174,30 @@ def nll(pred, target, pred_std, mask=None, average_grid=True, sum_vars=True): def crps_gauss( pred, target, pred_std, mask=None, average_grid=True, sum_vars=True ): - """ - (Negative) Continuous Ranked Probability Score (CRPS) - Closed-form expression based on Gaussian predictive distribution - - (...,) is any number of batch dimensions, potentially different - but broadcastable - pred: (..., N, d_state), prediction - target: (..., N, d_state), target - pred_std: (..., N, d_state) or (d_state,), predicted std.-dev. - 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) + """Compute negative CRPS for Gaussian predictive distribution. + + Args: + pred (torch.Tensor): Predictions of shape (..., N, d_state). + target (torch.Tensor): Targets of shape (..., N, d_state). + pred_std (torch.Tensor): Predicted standard deviation. + mask (torch.Tensor, optional): Boolean mask over grid nodes. + average_grid (bool): Whether to average over grid dimension. + sum_vars (bool): Whether to sum over variable dimension. Returns: - metric_val: One of (...,), (..., d_state), (..., N), (..., N, d_state), - depending on reduction arguments. + torch.Tensor: CRPS score depending on reductions. """ std_normal = torch.distributions.Normal( - torch.zeros((), device=pred.device), torch.ones((), device=pred.device) + torch.zeros((), device=pred.device), + torch.ones((), device=pred.device), ) - target_standard = (target - pred) / pred_std # (..., N, d_state) + target_standard = (target - pred) / pred_std entry_crps = -pred_std * ( torch.pi ** (-0.5) - 2 * torch.exp(std_normal.log_prob(target_standard)) - target_standard * (2 * std_normal.cdf(target_standard) - 1) - ) # (..., N, d_state) + ) return mask_and_reduce_metric( entry_crps, mask=mask, average_grid=average_grid, sum_vars=sum_vars