Download model/precipdd.py from OneScience-Group/PrecipDD: direct link, hf CLI and curl.
- Browser
- Download file 5.72 kB
-
https://huggingface.co/OneScience-Group/PrecipDD/resolve/main/model/precipdd.py
- Command line
-
hf download hf://OneScience-Group/PrecipDD/model/precipdd.py
-
curl -L -o precipdd.py https://huggingface.co/OneScience-Group/PrecipDD/resolve/main/model/precipdd.py
5.72 kB
| """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 | |
| 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") | |