"""Evaluate four-class forecasts and plot clusters and year-wise predictions.""" import json import sys from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np import yaml ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from model.cesm_seasonal_ml import classification_metrics def main(): config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) predictions = np.load(ROOT / config["paths"]["inference"]) methods = sorted(key.removeprefix("probability_") for key in predictions.files if key.startswith("probability_")) metrics = {method: classification_metrics(predictions["target"], predictions[f"probability_{method}"], config["evaluation"]["random_trials"], config["seed"]) for method in methods} output = ROOT / config["paths"]["evaluation_dir"] output.mkdir(parents=True, exist_ok=True) (output / "metrics.json").write_text(json.dumps({"season": str(predictions["season"]), "class_labels": {str(i + 1): str(name) for i, name in enumerate(predictions["class_names"])}, "models": metrics}, indent=2) + "\n") figure, axes = plt.subplots(1, 4, figsize=(13, 3.2), constrained_layout=True) limit = np.max(np.abs(predictions["centroids"])) for index, axis in enumerate(axes): image = axis.imshow(predictions["centroids"][index], origin="lower", cmap="BrBG", vmin=-limit, vmax=limit, extent=[predictions["longitude"][0], predictions["longitude"][-1], predictions["latitude"][0], predictions["latitude"][-1]], aspect="auto") axis.set_title(f"Class {index + 1}\n{predictions['class_names'][index]}", fontsize=9) axis.set_xlabel("Longitude") axes[0].set_ylabel("Latitude") figure.colorbar(image, ax=axes, label="standardized precipitation anomaly", shrink=0.8) figure.savefig(output / "precipitation_clusters.png", dpi=160) plt.close(figure) primary = str(predictions["primary_model"]) probability = predictions[f"probability_{primary}"] years, target, predicted = predictions["years"], predictions["target"] + 1, probability.argmax(axis=1) + 1 figure, axes = plt.subplots(2, 1, figsize=(12, 6), constrained_layout=True, sharex=True) axes[0].step(years, target, where="mid", label="Observed cluster", linewidth=2) axes[0].step(years, predicted, where="mid", label=f"{primary.upper()} prediction", alpha=0.8) axes[0].set_yticks([1, 2, 3, 4]); axes[0].set_ylabel("Class"); axes[0].legend(ncol=2) image = axes[1].imshow(probability.T, aspect="auto", origin="lower", cmap="viridis", vmin=0, vmax=1, extent=[years[0]-0.5, years[-1]+0.5, 0.5, 4.5]) axes[1].set_yticks([1, 2, 3, 4]); axes[1].set_ylabel("Class probability"); axes[1].set_xlabel("Year") figure.colorbar(image, ax=axes[1], label="Probability") figure.savefig(output / "seasonal_predictions.png", dpi=160) plt.close(figure) figure, axes = plt.subplots(1, 2, figsize=(12, 4.2), constrained_layout=True) x = np.arange(len(methods)); width = 0.34 axes[0].bar(x - width / 2, [metrics[name]["accuracy"] for name in methods], width, label="Accuracy") axes[0].bar(x + width / 2, [metrics[name]["grouped_accuracy"] for name in methods], width, label="Grouped accuracy") axes[0].axhline(metrics[primary]["most_frequent_baseline"], color="black", linestyle="--", label="Most-frequent baseline") axes[0].axhline(metrics[primary]["random_baseline"]["mean"], color="gray", linestyle=":", label="Random baseline") axes[0].set_xticks(x, [name.upper() for name in methods]); axes[0].set_ylim(0, 1); axes[0].legend(fontsize=8) axes[0].set_title("Core forecast metrics") limit = np.max(np.abs(predictions["centroids"])) axes[1].imshow(np.concatenate(list(predictions["centroids"]), axis=1), origin="lower", cmap="BrBG", vmin=-limit, vmax=limit, aspect="auto") axes[1].set_title("Clusters 1-4"); axes[1].set_xlabel("Concatenated longitude"); axes[1].set_ylabel("Latitude") figure.savefig(output / "comparison.png", dpi=160) plt.close(figure) print(f"evaluation={output.relative_to(ROOT)} models={','.join(methods)} primary_accuracy={metrics[primary]['accuracy']:.3f}") if __name__ == "__main__": main()