| """ClimODE evaluation metrics and output serialization helpers.""" |
|
|
| from __future__ import annotations |
|
|
| import json |
| import math |
| from pathlib import Path |
| from typing import Sequence |
|
|
| import numpy as np |
|
|
| try: |
| from scipy.special import erf as _erf |
| except ImportError: |
| _erf = np.vectorize(math.erf) |
|
|
|
|
| VARIABLES = ("z", "t", "t2m", "u10", "v10") |
|
|
|
|
| def latitude_weights(lat2d: np.ndarray) -> np.ndarray: |
| lat = np.asarray(lat2d, dtype=np.float64) |
| if lat.ndim == 2: |
| lat = lat[:, 0] |
| weights = np.cos(np.deg2rad(lat)) |
| weights = weights / np.mean(weights) |
| return weights[:, None] |
|
|
|
|
| def _check_arrays( |
| predictions: np.ndarray, |
| targets: np.ndarray, |
| std: np.ndarray | None, |
| valid_lengths: Sequence[int] | None = None, |
| ) -> None: |
| if predictions.shape != targets.shape: |
| raise ValueError(f"predictions {predictions.shape} != targets {targets.shape}") |
| if predictions.ndim != 6: |
| raise ValueError("Expected [samples, lead, years, channels, height, width]") |
| if predictions.shape[3] != len(VARIABLES): |
| raise ValueError(f"Expected {len(VARIABLES)} channels, got {predictions.shape[3]}") |
| if std is not None and std.shape != predictions.shape: |
| raise ValueError(f"std {std.shape} != predictions {predictions.shape}") |
| if valid_lengths is not None: |
| lengths = np.asarray(valid_lengths, dtype=np.int64) |
| if lengths.shape != (predictions.shape[0],): |
| raise ValueError(f"valid_lengths {lengths.shape} != ({predictions.shape[0]},)") |
| if np.any(lengths < 1) or np.any(lengths > predictions.shape[1]): |
| raise ValueError("valid_lengths must be within the lead dimension") |
|
|
|
|
| def _lead_mask( |
| predictions: np.ndarray, |
| valid_lengths: Sequence[int] | None, |
| ) -> np.ndarray: |
| lengths = ( |
| np.full(predictions.shape[0], predictions.shape[1], dtype=np.int64) |
| if valid_lengths is None |
| else np.asarray(valid_lengths, dtype=np.int64) |
| ) |
| return (np.arange(predictions.shape[1])[None, :] < lengths[:, None]).reshape( |
| predictions.shape[0], predictions.shape[1], 1, 1, 1, 1 |
| ) |
|
|
|
|
| def _weighted_mean(values: np.ndarray, weights: np.ndarray) -> np.ndarray: |
| |
| weighted = values * weights[None, None, None, None, :, :] |
| return weighted.mean(axis=(-1, -2)) |
|
|
|
|
| def latitude_weighted_rmse( |
| predictions: np.ndarray, |
| targets: np.ndarray, |
| lat2d: np.ndarray, |
| valid_lengths: Sequence[int] | None = None, |
| ) -> np.ndarray: |
| weights = latitude_weights(lat2d) |
| error = np.square(np.nan_to_num(predictions - targets, nan=0.0)) |
| per_field = np.sqrt(_weighted_mean(error, weights)) |
| valid = _lead_mask(predictions, valid_lengths)[..., 0, 0, 0, 0] |
| valid_fields = np.broadcast_to(valid[:, :, None, None], per_field.shape) |
| return (per_field * valid_fields).sum(axis=(0, 2)) / np.maximum( |
| valid_fields.sum(axis=(0, 2)), 1.0 |
| ) |
|
|
|
|
| def anomaly_correlation( |
| predictions: np.ndarray, |
| targets: np.ndarray, |
| lat2d: np.ndarray, |
| valid_lengths: Sequence[int] | None = None, |
| ) -> np.ndarray: |
| weights = latitude_weights(lat2d) |
| valid = _lead_mask(predictions, valid_lengths) |
| valid_broadcast = np.broadcast_to(valid, targets.shape) |
| target_clean = np.nan_to_num(targets, nan=0.0) |
| valid_count = valid_broadcast.sum(axis=(0, 1)) |
| |
| climatology = (target_clean * valid_broadcast).sum(axis=(0, 1)) / np.maximum( |
| valid_count, 1.0 |
| ) |
| pred_anomaly = np.nan_to_num(predictions, nan=0.0) - climatology[None, None] |
| target_anomaly = target_clean - climatology[None, None] |
| pred_anomaly -= pred_anomaly.mean(axis=(-1, -2), keepdims=True) |
| target_anomaly -= target_anomaly.mean(axis=(-1, -2), keepdims=True) |
| weighted_mask = weights[None, None, None, None] * valid |
| numerator = (pred_anomaly * target_anomaly * weighted_mask).sum(axis=(-1, -2)) |
| pred_norm = np.sqrt((np.square(pred_anomaly) * weighted_mask).sum(axis=(-1, -2))) |
| target_norm = np.sqrt((np.square(target_anomaly) * weighted_mask).sum(axis=(-1, -2))) |
| per_field = numerator / np.maximum(pred_norm * target_norm, 1.0e-12) |
| valid_fields = np.broadcast_to(valid[..., 0, 0], per_field.shape) |
| return (per_field * valid_fields).sum(axis=(0, 2)) / np.maximum( |
| valid_fields.sum(axis=(0, 2)), 1.0 |
| ) |
|
|
|
|
| def _normal_crps( |
| observations: np.ndarray, |
| means: np.ndarray, |
| scales: np.ndarray, |
| ) -> np.ndarray: |
| """Closed-form CRPS for a Gaussian predictive distribution.""" |
|
|
| scales = np.maximum(np.asarray(scales, dtype=np.float64), 1.0e-6) |
| z = (np.asarray(observations, dtype=np.float64) - means) / scales |
| phi = np.exp(-0.5 * np.square(z)) / math.sqrt(2.0 * math.pi) |
| cdf = 0.5 * (1.0 + _erf(z / math.sqrt(2.0))) |
| return scales * (z * (2.0 * cdf - 1.0) + 2.0 * phi - 1.0 / math.sqrt(math.pi)) |
|
|
|
|
| def gaussian_crps( |
| predictions: np.ndarray, |
| targets: np.ndarray, |
| std: np.ndarray, |
| valid_lengths: Sequence[int] | None = None, |
| ) -> np.ndarray: |
| values = np.nan_to_num(_normal_crps(targets, predictions, std), nan=0.0) |
| mask = np.broadcast_to(_lead_mask(predictions, valid_lengths), values.shape) |
| return (values * mask).sum(axis=(0, 2, 4, 5)) / np.maximum( |
| mask.sum(axis=(0, 2, 4, 5)), 1.0 |
| ) |
|
|
|
|
| def evaluate( |
| predictions: np.ndarray, |
| targets: np.ndarray, |
| lat2d: np.ndarray, |
| std: np.ndarray | None = None, |
| crps_predictions: np.ndarray | None = None, |
| crps_targets: np.ndarray | None = None, |
| crps_std: np.ndarray | None = None, |
| valid_lengths: Sequence[int] | None = None, |
| ) -> dict: |
| _check_arrays(predictions, targets, std, valid_lengths) |
| result = { |
| "variables": list(VARIABLES), |
| "lead_times_hours": [6 * (index + 1) for index in range(predictions.shape[1])], |
| "rmse": latitude_weighted_rmse(predictions, targets, lat2d, valid_lengths).tolist(), |
| "acc": anomaly_correlation(predictions, targets, lat2d, valid_lengths).tolist(), |
| "rmse_space": "physical", |
| "acc_space": "physical", |
| } |
| if std is not None: |
| result["crps"] = gaussian_crps( |
| crps_predictions if crps_predictions is not None else predictions, |
| crps_targets if crps_targets is not None else targets, |
| crps_std if crps_std is not None else std, |
| valid_lengths, |
| ).tolist() |
| result["crps_space"] = "normalized" if crps_predictions is not None else "physical" |
| result["crps_implementation"] = "closed_form_gaussian" |
| return result |
|
|
|
|
| def save_metrics(metrics: dict, path: str | Path) -> None: |
| output = Path(path) |
| output.parent.mkdir(parents=True, exist_ok=True) |
| output.write_text(json.dumps(metrics, indent=2), encoding="utf-8") |
|
|