#!/usr/bin/env python3 import argparse from pathlib import Path import matplotlib.pyplot as plt import numpy as np from model.metnet_2 import load_config, scores, write_json parser = argparse.ArgumentParser(description="Evaluate and visualize MetNet-2 predictions") parser.add_argument("--config", default="conf/config.yaml") args = parser.parse_args() config = load_config(args.config) with np.load(config["paths"]["predictions"]) as data: probabilities, target, rates = data["probabilities"], data["target"], data["rates"] metrics = {str(int(data["lead_minutes"])): scores(probabilities, target)} if not all(np.isfinite(value) for value in [metrics[next(iter(metrics))]["discrete_crps"], *metrics[next(iter(metrics))]["brier"].values(), *metrics[next(iter(metrics))]["csi"].values()]): raise FloatingPointError("evaluation metrics are not finite") write_json(config["paths"]["evaluation_metrics"], metrics) expected, truth = (probabilities * rates[:, None, None]).sum(0), rates[target] figure, axes = plt.subplots(1, 3, figsize=(11, 3.5), constrained_layout=True) for axis, image, title in zip(axes, (truth, expected, expected - truth), ("Target", "Expected rate", "Error")): plot = axis.imshow(image, cmap="viridis") axis.set_title(title) axis.set_axis_off() figure.colorbar(plot, ax=axis, shrink=.75) comparison = Path(config["paths"]["comparison"]) comparison.parent.mkdir(parents=True, exist_ok=True) figure.savefig(comparison, dpi=140) plt.close(figure) print(config["paths"]["evaluation_metrics"], comparison)