File size: 2,720 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 | """Run all PPNN replicas and save the complete verification payload."""
import sys
from pathlib import Path
import numpy as np
import torch
import yaml
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.ppnn import FORMAT_VERSION, PPNN, ensemble_features
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
data = np.load(ROOT / config["data"]["file"])
checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location="cpu", weights_only=False)
if str(data["format_version"]) != FORMAT_VERSION or checkpoint.get("format_version") != FORMAT_VERSION:
raise ValueError("data/checkpoint format version mismatch")
if data["ensemble"].shape[1:] != (50, 18) or len(data["station_id"]) != 537 or int(data["lead_hours"]) != 48:
raise ValueError("inference requires 50 members, 18 variables, 537 stations, and 48h lead")
scaling = checkpoint["scaling"]
continuous = ensemble_features(torch.from_numpy(data["ensemble"]), torch.from_numpy(data["auxiliary"]))
continuous = (continuous - scaling["feature_mean"]) / scaling["feature_std"]
station = torch.from_numpy(data["station_index"])
mus, sigmas = [], []
for state in checkpoint["model"]:
model = PPNN(int(checkpoint["model_config"]["hidden_size"]), eps=float(checkpoint["model_config"]["sigma_epsilon"]))
model.load_state_dict(state)
model.eval()
with torch.no_grad():
mu_scaled, sigma_scaled = model(continuous, station)
mus.append(mu_scaled * scaling["target_std"] + scaling["target_mean"])
sigmas.append(sigma_scaled * scaling["target_std"])
# The paper averages Gaussian parameters from random networks; this is not a mixture distribution.
mu = torch.stack(mus).mean(0).numpy()
sigma = torch.stack(sigmas).mean(0).numpy()
target = data["target"]
raw_t2m = data["ensemble"][:, :, 0]
if not all(np.isfinite(x).all() for x in (mu, sigma, target, raw_t2m)) or not np.all(sigma > 0):
raise RuntimeError("inference output is non-finite or has non-positive sigma")
output = ROOT / config["paths"]["inference"]
output.parent.mkdir(parents=True, exist_ok=True)
np.savez_compressed(output, mu=mu, sigma=sigma, targets=target, raw_ensemble_t2m=raw_t2m,
station_index=data["station_index"], station_id=data["station_id"], date_index=data["date_index"],
dates=data["dates"], lead_hours=data["lead_hours"], format_version=np.asarray(FORMAT_VERSION))
print(f"saved={output.relative_to(ROOT)} samples={len(mu)} replicas={len(mus)} raw_ensemble={raw_t2m.shape}")
if __name__ == "__main__":
main()
|