| """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() |
|
|