File size: 4,224 Bytes
b300acf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
"""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()