"""Evaluate Gaussian and raw-ensemble forecasts and create calibration plots.""" import json import math import sys from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt 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, ensemble_crps, gaussian_crps def main(): config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) data = np.load(ROOT / config["paths"]["inference"]) if str(data["format_version"]) != FORMAT_VERSION or data["raw_ensemble_t2m"].shape != (len(data["mu"]), 50): raise ValueError("prediction version/shape mismatch") mu = torch.from_numpy(data["mu"]) sigma = torch.from_numpy(data["sigma"]) target = torch.from_numpy(data["targets"]) raw = torch.from_numpy(data["raw_ensemble_t2m"]) model_scores = gaussian_crps(mu, sigma, target) raw_scores = ensemble_crps(raw, target) model_mean, raw_mean = float(model_scores.mean()), float(raw_scores.mean()) pit = 0.5 * (1.0 + torch.erf((target - mu) / sigma / math.sqrt(2.0))).numpy() bins = np.linspace(0.0, 1.0, 11) counts, _ = np.histogram(pit, bins=bins) observed = np.asarray([(pit <= edge).mean() for edge in bins[1:-1]]) nominal = bins[1:-1] calibration_mae = float(np.mean(np.abs(observed - nominal))) rmse = float(torch.sqrt(torch.mean((mu - target).square()))) spread_error_ratio = float(sigma.mean()) / rmse metrics = { "note": "structured virtual-data engineering validation, not paper performance", "samples": len(mu), "station_metadata_count": len(data["station_id"]), "lead_hours": int(data["lead_hours"]), "mean_gaussian_crps": model_mean, "mean_raw_ensemble_crps": raw_mean, "crpss": 1.0 - model_mean / raw_mean, "pit_histogram_counts": counts.tolist(), "pit_bin_edges": bins.tolist(), "calibration_mae": calibration_mae, "rmse_mu": rmse, "mean_sigma": float(sigma.mean()), "spread_error_ratio": spread_error_ratio, } numeric = [model_mean, raw_mean, metrics["crpss"], calibration_mae, rmse, metrics["mean_sigma"], spread_error_ratio] if not np.isfinite(numeric).all() or raw_mean <= 0 or rmse <= 0: raise RuntimeError("evaluation metrics are not finite and valid") output = ROOT / config["paths"]["evaluation_dir"] output.mkdir(parents=True, exist_ok=True) (output / "metrics.json").write_text(json.dumps(metrics, indent=2) + "\n") fig, axes = plt.subplots(1, 3, figsize=(13, 4), constrained_layout=True) axes[0].bar((bins[:-1] + bins[1:]) / 2, counts / counts.sum(), width=0.085) axes[0].axhline(0.1, color="black", linestyle="--", linewidth=1) axes[0].set(xlabel="PIT", ylabel="relative frequency", title="PIT histogram") axes[1].plot(nominal, observed, marker="o"); axes[1].plot([0, 1], [0, 1], "k--") axes[1].set(xlabel="nominal CDF", ylabel="observed frequency", title="Calibration") axes[2].scatter(sigma.numpy(), np.abs(mu.numpy() - target.numpy()), c=data["date_index"], cmap="viridis", s=22) axes[2].set(xlabel="predictive sigma", ylabel="absolute error", title=f"Spread/error ratio = {spread_error_ratio:.2f}") fig.savefig(output / "ppnn_evaluation.png", dpi=150) plt.close(fig) print(f"gaussian_crps={model_mean:.6f} raw_crps={raw_mean:.6f} crpss={metrics['crpss']:.6f} spread_error_ratio={spread_error_ratio:.6f}") if __name__ == "__main__": main()