"""Compute paper precipitation diagnostics and render a comparison figure.""" import json from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np import yaml ROOT = Path(__file__).resolve().parents[1] def pattern_correlation(target, prediction): target_anomaly = target - target.mean() prediction_anomaly = prediction - prediction.mean() denominator = np.sqrt(np.sum(target_anomaly ** 2) * np.sum(prediction_anomaly ** 2)) return float(np.sum(target_anomaly * prediction_anomaly) / max(float(denominator), 1e-12)) def precipitation_metrics(values, years, wet, extreme, percentile): daily_domain_mean = values.mean(axis=(1, 2, 3)) yearly_p95 = [float(np.percentile(daily_domain_mean[years == year], percentile)) for year in np.unique(years)] return {"P_mean_mm_day": float(values.mean()), "yearly_P95_mm_day": yearly_p95, "P95_mean_over_years_mm_day": float(np.mean(yearly_p95)), "wet_day_fraction_gt_1mm": float(np.mean(values > wet)), "extreme_day_fraction_gt_25mm": float(np.mean(values > extreme))} def main(): config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) data = np.load(ROOT / config["paths"]["inference"]) prediction, target, years = data["prediction"], data["target"], data["years"] if prediction.shape != target.shape or prediction.shape[1:] != (1, 216, 488): raise ValueError("evaluation requires complete [N,1,216,488] fields") evaluation = config["evaluation"] target_metrics = precipitation_metrics(target, years, float(evaluation["wet_day_threshold_mm"]), float(evaluation["extreme_day_threshold_mm"]), float(evaluation["p95_percentile"])) prediction_metrics = precipitation_metrics(prediction, years, float(evaluation["wet_day_threshold_mm"]), float(evaluation["extreme_day_threshold_mm"]), float(evaluation["p95_percentile"])) target_pattern, prediction_pattern = target.mean(0).squeeze(), prediction.mean(0).squeeze() report = {"format_version": str(data["format_version"]), "units": str(data["units"]), "sample_count": len(target), "target": target_metrics, "prediction": prediction_metrics, "bias_mm_day": float((prediction - target).mean()), "pattern_correlation": pattern_correlation(target_pattern, prediction_pattern), "definitions": {"P": "mean precipitation over all days and grid cells", "P95": "95th percentile of daily domain-mean P per year, then averaged", "wet_day": "grid-cell day precipitation > 1 mm/day", "extreme_day": "grid-cell day precipitation > 25 mm/day"}} numbers = [report["bias_mm_day"], report["pattern_correlation"]] numbers += list(target_metrics.values())[:1] + list(prediction_metrics.values())[:1] if not np.isfinite(numbers).all(): raise FloatingPointError("non-finite evaluation metrics") output = ROOT / config["paths"]["evaluation_dir"] output.mkdir(parents=True, exist_ok=True) (output / "metrics.json").write_text(json.dumps(report, indent=2) + "\n") vmax = float(np.percentile(target[0], 99.5)) figure, axes = plt.subplots(1, 3, figsize=(13, 4)) for axis, field, title in zip(axes, (data["bilinear_precipitation"][0, 0], target[0, 0], prediction[0, 0]), ("Bilinear input", "Target", "SRCNN")): image = axis.imshow(field, cmap="Blues", vmin=0, vmax=max(vmax, 1.0), aspect="auto") axis.set_title(title); axis.set_xlabel("x"); axis.set_ylabel("y") figure.colorbar(image, ax=axes, label="mm/day", shrink=0.8) figure.subplots_adjust(left=0.06, right=0.9, bottom=0.12, top=0.88, wspace=0.25) figure.savefig(output / "precipitation_comparison.png", dpi=150); plt.close(figure) print(f"evaluation={output.relative_to(ROOT)} pattern_correlation={report['pattern_correlation']:.4f}") if __name__ == "__main__": main()