"""Dimension-faithful PyTorch implementation of the precipitation DD CNN.""" from __future__ import annotations import json import math import random from pathlib import Path from typing import Iterable import numpy as np import torch import yaml from torch import Tensor, nn DATA_FORMAT_VERSION = "precipdd_v1" INPUT_SHAPE = (1, 55, 160) def load_config(path: str | Path) -> dict: with open(path, "r", encoding="utf-8") as handle: return yaml.safe_load(handle) def seed_all(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) class PrecipDD(nn.Module): """Five-convolution scalar regressor with the paper's 8,960 features.""" def __init__(self, filters: Iterable[int] = (8, 8, 16, 16, 16), dense_units: int = 32): super().__init__() filters = tuple(int(value) for value in filters) if len(filters) != 5 or filters[-1] != 16: raise ValueError("DD requires five convolution layers and 16 final filters") layers: list[nn.Module] = [] channels = 1 for index, width in enumerate(filters): layers.extend((nn.Conv2d(channels, width, 3, padding=1), nn.Tanh())) if index < 2: # TensorFlow SAME pooling is required for 55 -> 28 -> 14 latitude points. layers.append(nn.MaxPool2d(2, stride=2, ceil_mode=True)) channels = width self.features = nn.Sequential(*layers) self.hidden = nn.Linear(16 * 14 * 40, dense_units) self.output = nn.Linear(dense_units, 1) self.reset_parameters() def reset_parameters(self) -> None: for module in self.modules(): if isinstance(module, (nn.Conv2d, nn.Linear)): nn.init.xavier_uniform_(module.weight) fan_in, fan_out = nn.init._calculate_fan_in_and_fan_out(module.weight) bound = math.sqrt(6.0 / (fan_in + fan_out)) nn.init.uniform_(module.bias, -bound, bound) def forward_features(self, precipitation: Tensor) -> Tensor: if precipitation.ndim != 4 or tuple(precipitation.shape[1:]) != INPUT_SHAPE: raise ValueError(f"input must have shape [B,1,55,160], got {tuple(precipitation.shape)}") features = self.features(precipitation) if tuple(features.shape[1:]) != (16, 14, 40): raise RuntimeError(f"feature shape must be [B,16,14,40], got {tuple(features.shape)}") return features def forward(self, precipitation: Tensor) -> Tensor: features = self.forward_features(precipitation).flatten(1) return self.output(torch.sigmoid(self.hidden(features))).squeeze(-1) def validate_archive(archive: np.lib.npyio.NpzFile) -> None: required = {"precipitation", "agmt", "split", "year", "day_of_year", "latitude", "longitude", "format_version"} if missing := required.difference(archive.files): raise ValueError(f"dataset missing fields: {sorted(missing)}") if str(archive["format_version"]) != DATA_FORMAT_VERSION: raise ValueError("dataset format_version mismatch") if archive["precipitation"].ndim != 4 or tuple(archive["precipitation"].shape[1:]) != INPUT_SHAPE: raise ValueError("precipitation must be float data with shape [N,1,55,160]") if archive["agmt"].shape != (len(archive["precipitation"]),): raise ValueError("AGMT must contain one scalar for every daily map") if archive["latitude"].shape != (55,) or archive["longitude"].shape != (160,): raise ValueError("coordinates must contain 55 latitudes and 160 extended longitudes") if not np.all(np.isfinite(archive["precipitation"])) or not np.all(np.isfinite(archive["agmt"])): raise ValueError("dataset contains non-finite values") def load_ensemble(checkpoint_path: str | Path, device: torch.device) -> tuple[list[PrecipDD], dict]: try: checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) except TypeError: checkpoint = torch.load(checkpoint_path, map_location=device) if checkpoint.get("format_version") != DATA_FORMAT_VERSION: raise ValueError("checkpoint format_version mismatch") settings = checkpoint["model_config"] models = [] for state in checkpoint["ensemble_states"]: model = PrecipDD(settings["filters"], settings["dense_units"]).to(device) model.load_state_dict(state) model.eval() models.append(model) return models, checkpoint @torch.no_grad() def ensemble_predict(models: list[PrecipDD], values: Tensor, batch_size: int = 32) -> Tensor: predictions = [] for start in range(0, len(values), batch_size): batch = values[start:start + batch_size] predictions.append(torch.stack([model(batch) for model in models]).mean(0)) return torch.cat(predictions) if predictions else torch.empty(0, device=values.device) def linear_trend(values: np.ndarray, years: np.ndarray) -> float: valid = np.isfinite(values) & np.isfinite(years) if valid.sum() < 2 or np.ptp(years[valid]) == 0: return float("nan") return float(np.polyfit(years[valid], values[valid], 1)[0] * 10.0) def correlation(target: np.ndarray, prediction: np.ndarray) -> float: if len(target) < 2 or np.std(target) == 0 or np.std(prediction) == 0: return float("nan") return float(np.corrcoef(target, prediction)[0, 1]) def write_json(path: str | Path, payload: dict) -> None: path = Path(path) path.parent.mkdir(parents=True, exist_ok=True) path.write_text(json.dumps(payload, indent=2, allow_nan=False) + "\n", encoding="utf-8")