Download scripts/result.py from OneScience-Group/ClimODE: direct link, hf CLI and curl.
- Browser
- Download file 4.72 kB
-
https://huggingface.co/OneScience-Group/ClimODE/resolve/main/scripts/result.py
- Command line
-
hf download hf://OneScience-Group/ClimODE/scripts/result.py
-
curl -L -o result.py https://huggingface.co/OneScience-Group/ClimODE/resolve/main/scripts/result.py
4.72 kB
| """Compute metrics and render ClimODE forecast maps from saved outputs.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import yaml | |
| PROJECT_ROOT = Path(__file__).resolve().parents[1] | |
| if str(PROJECT_ROOT) not in sys.path: | |
| sys.path.insert(0, str(PROJECT_ROOT)) | |
| from scripts.metrics import evaluate, save_metrics | |
| def _resolve(path: str | Path) -> Path: | |
| value = Path(path) | |
| return value if value.is_absolute() else PROJECT_ROOT / value | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--config", type=Path, default=PROJECT_ROOT / "conf/config.yaml") | |
| parser.add_argument("--predictions", type=Path, default=None) | |
| parser.add_argument("--targets", type=Path, default=None) | |
| parser.add_argument("--std", type=Path, default=None) | |
| parser.add_argument("--output-dir", type=Path, default=None) | |
| parser.add_argument("--stats-dir", type=Path, default=None) | |
| parser.add_argument("--static-file", type=Path, default=None) | |
| parser.add_argument("--sample", type=int, default=0) | |
| parser.add_argument("--lead", type=int, default=0) | |
| args = parser.parse_args() | |
| with args.config.open("r", encoding="utf-8") as handle: | |
| config = yaml.safe_load(handle) | |
| output_dir = args.output_dir or _resolve(config["data"]["output_dir"]) | |
| predictions = np.load(args.predictions or output_dir / "predictions.npy") | |
| targets = np.load(args.targets or output_dir / "targets.npy") | |
| std_path = args.std or output_dir / "std.npy" | |
| std = np.load(std_path) if std_path.is_file() else None | |
| lengths_path = output_dir / "valid_lengths.npy" | |
| valid_lengths = np.load(lengths_path) if lengths_path.is_file() else None | |
| static_file = _resolve(args.static_file or config["data"]["static_file"]) | |
| import h5py | |
| with h5py.File(static_file, "r") as handle: | |
| lat2d = handle["lat2d"][:] | |
| stats_dir = _resolve( | |
| args.stats_dir | |
| or config["data"].get("stats_dir", Path(config["data"]["data_dir"]) / "static") | |
| ) | |
| minimum = np.load(stats_dir / "min_values.npy").reshape(1, 1, 1, 5, 1, 1) | |
| maximum = np.load(stats_dir / "max_values.npy").reshape(1, 1, 1, 5, 1, 1) | |
| scale = maximum - minimum | |
| metrics = evaluate( | |
| predictions * scale + minimum, | |
| targets * scale + minimum, | |
| lat2d, | |
| std * scale if std is not None else None, | |
| crps_predictions=predictions if std is not None else None, | |
| crps_targets=targets if std is not None else None, | |
| crps_std=std, | |
| valid_lengths=valid_lengths, | |
| ) | |
| metrics_path = output_dir.parent / "metrics.json" | |
| save_metrics(metrics, metrics_path) | |
| figure_dir = output_dir / "figures" | |
| figure_dir.mkdir(parents=True, exist_ok=True) | |
| try: | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| except ImportError as exc: | |
| raise RuntimeError("Visualization requires matplotlib in the active environment") from exc | |
| if not 0 <= args.sample < predictions.shape[0]: | |
| raise IndexError(f"sample must be in [0,{predictions.shape[0] - 1}]") | |
| if not 0 <= args.lead < predictions.shape[1]: | |
| raise IndexError(f"lead must be in [0,{predictions.shape[1] - 1}]") | |
| if valid_lengths is not None and args.lead >= int(valid_lengths[args.sample]): | |
| raise IndexError( | |
| f"lead {args.lead} is padding for sample {args.sample}; " | |
| f"valid length is {int(valid_lengths[args.sample])}" | |
| ) | |
| names = ["z", "t", "t2m", "u10", "v10"] | |
| for channel, name in enumerate(names): | |
| prediction = predictions[args.sample, args.lead, 0, channel] | |
| target = targets[args.sample, args.lead, 0, channel] | |
| difference = prediction - target | |
| figure, axes = plt.subplots(1, 3, figsize=(12, 3.4), constrained_layout=True) | |
| for axis, image, title in zip( | |
| axes, | |
| (prediction, target, difference), | |
| ("prediction", "target", "difference"), | |
| ): | |
| cmap = "RdBu_r" if title == "difference" else "viridis" | |
| plot = axis.imshow(image, cmap=cmap, origin="upper", aspect="auto") | |
| axis.set_title(title) | |
| axis.set_xlabel("longitude index") | |
| axis.set_ylabel("latitude index") | |
| figure.colorbar(plot, ax=axis, shrink=0.8) | |
| figure.suptitle(f"ClimODE {name}, lead={(args.lead + 1) * 6} h") | |
| figure.savefig(figure_dir / f"{name}_lead_{(args.lead + 1) * 6:03d}h.png", dpi=150) | |
| plt.close(figure) | |
| print(json.dumps({"metrics": str(metrics_path), "figures": str(figure_dir)})) | |
| if __name__ == "__main__": | |
| main() | |