PPNN / scripts /result.py
zhangrenchao's picture
Upload folder using huggingface_hub
6ff9439 verified
Raw
History Blame Contribute Delete
3.51 kB
"""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()