File size: 3,630 Bytes
6ff9439 | 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 | """PPNN NN-aux-emb model and proper scoring rules."""
from __future__ import annotations
import math
import torch
from torch import nn
FORMAT_VERSION = "ppnn_nn_aux_emb_v1"
MEMBER_COUNT = 50
VARIABLE_COUNT = 18
STATION_COUNT = 537
LEAD_HOURS = 48
AUX_FEATURE_COUNT = 4
CONTINUOUS_FEATURE_COUNT = VARIABLE_COUNT * 2 + AUX_FEATURE_COUNT
EMBEDDING_DIM = 2
INPUT_FEATURE_COUNT = CONTINUOUS_FEATURE_COUNT + EMBEDDING_DIM
class PPNN(nn.Module):
"""NN-aux-emb: ensemble moments, auxiliary predictors, and station embedding."""
def __init__(self, hidden_size: int = 32, station_count: int = STATION_COUNT, eps: float = 1e-4):
super().__init__()
if station_count != STATION_COUNT:
raise ValueError(f"station_count must be the paper value {STATION_COUNT}")
self.station_count = station_count
self.eps = float(eps)
self.station_embedding = nn.Embedding(station_count, EMBEDDING_DIM)
self.hidden = nn.Linear(INPUT_FEATURE_COUNT, hidden_size)
self.output = nn.Linear(hidden_size, 2)
def forward(self, continuous: torch.Tensor, station_index: torch.Tensor):
if continuous.ndim != 2 or continuous.shape[-1] != CONTINUOUS_FEATURE_COUNT:
raise ValueError(f"continuous features must have shape [batch,{CONTINUOUS_FEATURE_COUNT}]")
if station_index.ndim != 1 or station_index.shape[0] != continuous.shape[0]:
raise ValueError("station_index must have shape [batch]")
if station_index.numel() and (station_index.min() < 0 or station_index.max() >= self.station_count):
raise ValueError("station index outside [0, 536]")
features = torch.cat((continuous, self.station_embedding(station_index)), dim=-1)
mu, raw_sigma = self.output(torch.relu(self.hidden(features))).unbind(dim=-1)
return mu, raw_sigma.abs() + self.eps
def ensemble_features(ensemble: torch.Tensor, auxiliary: torch.Tensor) -> torch.Tensor:
"""Reduce [batch, 50, 18] to 36 moments and append four auxiliary features."""
if ensemble.ndim != 3 or ensemble.shape[1:] != (MEMBER_COUNT, VARIABLE_COUNT):
raise ValueError(f"ensemble must have shape [batch,{MEMBER_COUNT},{VARIABLE_COUNT}]")
if auxiliary.ndim != 2 or auxiliary.shape != (ensemble.shape[0], AUX_FEATURE_COUNT):
raise ValueError(f"auxiliary must have shape [batch,{AUX_FEATURE_COUNT}]")
moments = torch.cat((ensemble.mean(dim=1), ensemble.std(dim=1, correction=1)), dim=-1)
return torch.cat((moments, auxiliary), dim=-1)
def gaussian_crps(mu: torch.Tensor, sigma: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
"""Closed-form CRPS for a Gaussian predictive distribution."""
if torch.any(sigma <= 0):
raise ValueError("sigma must be strictly positive")
z = (target - mu) / sigma
pdf = torch.exp(-0.5 * z.square()) / math.sqrt(2.0 * math.pi)
cdf = 0.5 * (1.0 + torch.erf(z / math.sqrt(2.0)))
return sigma * (z * (2.0 * cdf - 1.0) + 2.0 * pdf - 1.0 / math.sqrt(math.pi))
def ensemble_crps(ensemble: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
"""Empirical ensemble CRPS without constructing a member-pair matrix."""
if ensemble.shape[-1] != MEMBER_COUNT:
raise ValueError(f"last ensemble dimension must be {MEMBER_COUNT}")
sorted_members = ensemble.sort(dim=-1).values
weights = (2 * torch.arange(1, MEMBER_COUNT + 1, device=ensemble.device) - MEMBER_COUNT - 1).to(ensemble.dtype)
pair_term = (sorted_members * weights).sum(dim=-1) / (MEMBER_COUNT * MEMBER_COUNT)
return (ensemble - target.unsqueeze(-1)).abs().mean(dim=-1) - pair_term
|