CESM-SeasonalML / scripts /result.py
zhangrenchao's picture
Publish CESM-SeasonalML engineering reproduction
b300acf verified
Raw
History Blame Contribute Delete
4.22 kB
"""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()