| """Compute deterministic and probabilistic CorrDiff metrics and plots.""" |
|
|
| import argparse |
| 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 crps_ensemble(ensemble, target): |
| first = np.abs(ensemble - target[None]).mean(0) |
| sorted_members = np.sort(ensemble, axis=0) |
| m = ensemble.shape[0] |
| weights = (2 * np.arange(1, m + 1) - m - 1).reshape(m, 1, 1, 1, 1) |
| return first - (sorted_members * weights).sum(0) / m**2 |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--config", default=str(ROOT / "conf/config.yaml")) |
| args = parser.parse_args() |
| config = yaml.safe_load(Path(args.config).read_text(encoding="utf-8")) |
| archive = np.load(ROOT / config["paths"]["predictions"]) |
| if "protocol" not in archive or archive["protocol"].ndim != 0 or str(archive["protocol"].item()) != config["data"]["protocol"]: |
| raise ValueError("Prediction protocol does not match the configured protocol") |
| if "data_source" not in archive or archive["data_source"].ndim != 0 or not str(archive["data_source"].item()): |
| raise ValueError("Prediction data_source must be a non-empty scalar") |
| protocol = str(archive["protocol"].item()) |
| data_source = str(archive["data_source"].item()) |
| ensemble, target = archive["ensemble"], archive["target"] |
| mean, spread = ensemble.mean(0), ensemble.std(0) |
| axes = (0, 2, 3) |
| mae = np.abs(mean - target).mean(axis=axes) |
| rmse = np.sqrt(((mean - target) ** 2).mean(axis=axes)) |
| crps = crps_ensemble(ensemble, target).mean(axis=axes) |
| spread_value = spread.mean(axis=axes) |
| names = config["data"]["target_variables"] |
| metrics = {name: {"mae": float(mae[i]), "rmse": float(rmse[i]), "crps": float(crps[i]), |
| "ensemble_spread": float(spread_value[i])} for i, name in enumerate(names)} |
| metrics["aggregate"] = {key: float(np.mean([metrics[n][key] for n in names])) |
| for key in ("mae", "rmse", "crps", "ensemble_spread")} |
| output = ROOT / config["paths"]["evaluation_dir"] |
| output.mkdir(parents=True, exist_ok=True) |
| payload = {"metrics": metrics, "protocol": protocol, "data_source": data_source} |
| (output / "metrics.json").write_text(json.dumps(payload, indent=2) + "\n") |
| figure, plot_axes = plt.subplots(len(names), 4, figsize=(13, 3 * len(names))) |
| for channel, name in enumerate(names): |
| fields = (target[0, channel], mean[0, channel], spread[0, channel], mean[0, channel] - target[0, channel]) |
| titles = ("target", "ensemble mean", "ensemble spread", "mean error") |
| for axis, field, title in zip(plot_axes[channel], fields, titles): |
| axis.imshow(field, cmap="coolwarm" if title == "mean error" else "viridis") |
| axis.set_title(f"{name}: {title}"); axis.axis("off") |
| figure.tight_layout(); figure.savefig(output / "ensemble_diagnostics.png", dpi=120); plt.close(figure) |
| print(json.dumps(payload, indent=2)); print(f"evaluation={output}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|