"""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: # pragma: no cover - only used in minimal environments _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: # values: [N,L,Y,K,H,W], weights: [H,1] 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)) # Official evaluation uses one test-set climatology for each year/channel/grid. 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")