File size: 5,723 Bytes
950fc23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
"""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")