Download scripts/result.py from OneScience-Group/NNCAM: direct link, hf CLI and curl.
- Browser
- Download file 2.4 kB
-
https://huggingface.co/OneScience-Group/NNCAM/resolve/main/scripts/result.py
- Command line
-
hf download hf://OneScience-Group/NNCAM/scripts/result.py
-
curl -L -o result.py https://huggingface.co/OneScience-Group/NNCAM/resolve/main/scripts/result.py
2.4 kB
| #!/usr/bin/env python3 | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| ROOT = Path(__file__).resolve().parents[1] | |
| def score(truth, prediction): | |
| rmse = float(np.sqrt(np.mean((truth - prediction) ** 2))) | |
| denominator = float(np.sum((truth - truth.mean()) ** 2)) | |
| return {"rmse": rmse, "r2": float(1.0 - np.sum((truth - prediction) ** 2) / denominator) if denominator else None} | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Evaluate NNCAM predictions.") | |
| parser.add_argument("--input", type=Path, default=ROOT / "result/output/predictions.npz") | |
| parser.add_argument("--metrics", type=Path, default=ROOT / "result/evaluation/metrics.json") | |
| parser.add_argument("--figure", type=Path, default=ROOT / "result/evaluation/comparison.png") | |
| args = parser.parse_args() | |
| with np.load(args.input) as data: | |
| truth, prediction = data["truth"], data["prediction"] | |
| if truth.shape != prediction.shape or truth.ndim != 2 or truth.shape[1] != 65: | |
| raise ValueError(f"expected matching [N,65] arrays, got {truth.shape}, {prediction.shape}") | |
| groups = {"dT": slice(0, 30), "dQ": slice(30, 60), "SW": slice(60, 62), "LW": slice(62, 64), "P": slice(64, 65)} | |
| metrics = {name: score(truth[:, indices], prediction[:, indices]) for name, indices in groups.items()} | |
| metrics["overall"] = score(truth, prediction) | |
| values = np.array([value for group in metrics.values() for value in group.values() if value is not None]) | |
| if not np.isfinite(values).all(): | |
| raise RuntimeError("evaluation metrics contain non-finite values") | |
| args.metrics.parent.mkdir(parents=True, exist_ok=True) | |
| args.metrics.write_text(json.dumps(metrics, indent=2), encoding="utf-8") | |
| names = list(groups) | |
| fig, axes = plt.subplots(1, 2, figsize=(10, 4), constrained_layout=True) | |
| axes[0].bar(names, [metrics[name]["rmse"] for name in names]) | |
| axes[0].set_title("Grouped RMSE") | |
| axes[1].scatter(truth[:, 64], prediction[:, 64], s=12, alpha=0.7) | |
| axes[1].set(xlabel="True precipitation", ylabel="Predicted precipitation", title="Precipitation comparison") | |
| fig.savefig(args.figure, dpi=150) | |
| plt.close(fig) | |
| print(f"saved {args.metrics} and {args.figure}; shape={prediction.shape}, finite=true") | |
| if __name__ == "__main__": | |
| main() | |