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