"""Climate metrics and conservation diagnostics from the paper equations.""" from __future__ import annotations import torch def _weights(area_weights: torch.Tensor, field: torch.Tensor) -> torch.Tensor: weights = area_weights.to(device=field.device, dtype=field.dtype) if weights.ndim == 1: weights = weights[:, None] if weights.shape != field.shape[-2:]: raise ValueError(f"area weights must match spatial shape {field.shape[-2:]}, got {tuple(weights.shape)}") return weights / weights.sum().clamp_min(torch.finfo(field.dtype).eps) def area_weighted_global_mean(field: torch.Tensor, area_weights: torch.Tensor) -> torch.Tensor: weights = _weights(area_weights, field) return (field * weights).sum(dim=(-2, -1)) def time_mean_pattern_rmse(pred: torch.Tensor, truth: torch.Tensor, area_weights: torch.Tensor) -> torch.Tensor: if pred.shape != truth.shape: raise ValueError("pred and truth must have equal shapes") error = (pred - truth).mean(dim=0) if pred.ndim == 4 else (pred - truth).mean(dim=1) weights = _weights(area_weights, error) return torch.sqrt((error.square() * weights).sum(dim=(-2, -1)).clamp_min(0.0)) def global_time_mean_bias(pred: torch.Tensor, truth: torch.Tensor, area_weights: torch.Tensor) -> torch.Tensor: if pred.shape != truth.shape: raise ValueError("pred and truth must have equal shapes") pred_gm = area_weighted_global_mean(pred, area_weights) truth_gm = area_weighted_global_mean(truth, area_weights) return (pred_gm - truth_gm).mean(dim=0 if pred.ndim == 4 else 1) def total_water_path(qt: torch.Tensor, dp: torch.Tensor, gravity: float = 9.80665) -> torch.Tensor: """Compute TWP=(1/g) sum_k qT_k dp_k; qt/dp are [...,8,H,W].""" if qt.shape != dp.shape or qt.shape[-3] != 8: raise ValueError("qt and dp must have shape [...,8,H,W]") return (qt * dp).sum(dim=-3) / gravity def moisture_budget_violation(twp_t: torch.Tensor, twp_next: torch.Tensor, evaporation: torch.Tensor, precipitation: torch.Tensor, advective_tendency: torch.Tensor) -> torch.Tensor: """Eq.(1) residual using a one-step tendency.""" return (twp_next - twp_t) - (evaporation - precipitation + advective_tendency) def dry_air_surface_pressure(surface_pressure: torch.Tensor, twp: torch.Tensor, gravity: float = 9.80665) -> torch.Tensor: return surface_pressure - gravity * twp def forecast_metrics(pred: torch.Tensor, truth: torch.Tensor, area_weights: torch.Tensor) -> dict[str, float]: """Return aggregate Eq.(5-7) metrics for [T,C,H,W] or [B,T,C,H,W].""" if pred.shape != truth.shape: raise ValueError("pred and truth must have equal shapes") if pred.ndim == 5: pred = pred[0] truth = truth[0] rmse = time_mean_pattern_rmse(pred, truth, area_weights) bias = global_time_mean_bias(pred, truth, area_weights) return { "mean_time_mean_rmse": float(rmse.mean().item()), "mean_global_time_mean_bias": float(bias.mean().item()), "T7_time_mean_rmse": float(rmse[7].item()), "precipitation_time_mean_rmse": float(rmse[40].item()), }