CorrDiff / scripts /result.py
zhangrenchao's picture
Update CorrDiff model package
5b3329b verified
Raw
History Blame Contribute Delete
3.15 kB
"""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()