| |
| 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) |
|
|