Download scripts/result.py from OneScience-Group/MetNet-2: direct link, hf CLI and curl.
- Browser
- Download file 1.64 kB
-
https://huggingface.co/OneScience-Group/MetNet-2/resolve/main/scripts/result.py
- Command line
-
hf download hf://OneScience-Group/MetNet-2/scripts/result.py
-
curl -L -o result.py https://huggingface.co/OneScience-Group/MetNet-2/resolve/main/scripts/result.py
1.64 kB
| #!/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) | |