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