File size: 4,371 Bytes
a00d152
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
68
69
70
71
72
73
74
75
76
77
78
79
"""Evaluate Pearson correlation, RMSE, NSE and relative RMSE."""

import json
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]


def metrics(prediction, target):
    prediction, target = prediction.reshape(-1), target.reshape(-1)
    error = prediction - target
    target_centered = target - target.mean()
    prediction_centered = prediction - prediction.mean()
    denominator = np.sqrt(np.sum(target_centered ** 2) * np.sum(prediction_centered ** 2))
    correlation = float(np.sum(target_centered * prediction_centered) / denominator) if denominator > 0 else 0.0
    rmse = float(np.sqrt(np.mean(error ** 2)))
    variance = float(np.sum(target_centered ** 2))
    nse = float(1.0 - np.sum(error ** 2) / variance) if variance > 0 else 0.0
    surge_range = float(np.ptp(target))
    return {"pearson_correlation": correlation, "rmse_m": rmse, "nse": nse,
            "relative_rmse_percent": float(100 * rmse / surge_range) if surge_range > 0 else 0.0}


def main():
    config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
    data = np.load(ROOT / config["paths"]["inference_dir"] / "predictions.npz")
    if str(data["format_version"]) != config["data"]["format_version"]:
        raise ValueError("incompatible prediction format")
    target = data["targets_m"]
    names = [key.removeprefix("predictions_") for key in data.files if key.startswith("predictions_")]
    by_model = {name: metrics(data[f"predictions_{name}"], target) for name in names}
    by_model["GTSR_baseline"] = metrics(data["gtsr_m"], target)
    threshold = float(np.percentile(target, float(config["evaluation"]["extreme_percentile"])))
    extreme_mask = target[:, 0] > threshold
    extreme = {name: metrics(data[f"predictions_{name}"][extreme_mask], target[extreme_mask]) for name in names}
    best_name = max(names, key=lambda name: (by_model[name]["pearson_correlation"] + by_model[name]["nse"] -
                                              by_model[name]["relative_rmse_percent"] / 100))
    tropical = np.abs(data["latitude_degrees"]) < float(config["evaluation"]["tropical_latitude_limit_degrees"])
    regions = {}
    for region, mask in (("tropical", tropical), ("subtropical_extratropical", ~tropical)):
        regions[region] = metrics(data[f"predictions_{best_name}"][mask], target[mask]) if mask.sum() >= 2 else None
    report = {"samples": int(len(target)), "overall": by_model, "selected_model": best_name,
              "extreme_above_observed_95th_percentile": extreme, "regional_selected_model": regions,
              "evaluation_protocol": {"target": "daily maximum non-tidal residual", "unit": "m",
                                      "extreme_threshold_m": threshold, "aggregation": "all station-day samples",
                                      "missing_value_mask": "none in synthetic data"}}
    numeric = [value for model_metrics in by_model.values() for value in model_metrics.values()]
    if not np.isfinite(numeric).all():
        raise FloatingPointError("evaluation contains non-finite metrics")
    output = ROOT / config["paths"]["evaluation_dir"]
    output.mkdir(parents=True, exist_ok=True)
    (output / "metrics.json").write_text(json.dumps(report, indent=2) + "\n")
    prediction = data[f"predictions_{best_name}"][:, 0]
    order = np.argsort(data["timestamps_unix_s"])
    figure, axes = plt.subplots(1, 2, figsize=(11, 4.2))
    axes[0].plot(target[order, 0] * 100, label="Observed", linewidth=1.5)
    axes[0].plot(prediction[order] * 100, label=best_name, linewidth=1.2)
    axes[0].set(xlabel="Station-day index", ylabel="Daily maximum surge (cm)", title="Synthetic surge sequence")
    axes[0].legend()
    axes[1].scatter(target[:, 0] * 100, prediction * 100, c=np.abs(data["latitude_degrees"]), cmap="viridis", s=28)
    limits = [min(target.min(), prediction.min()) * 100, max(target.max(), prediction.max()) * 100]
    axes[1].plot(limits, limits, "k--", linewidth=1)
    axes[1].set(xlabel="Observed surge (cm)", ylabel="Modeled surge (cm)", title="Observed vs modeled")
    figure.tight_layout()
    figure.savefig(output / "comparison.png", dpi=150)
    plt.close(figure)
    print(f"evaluation={output.relative_to(ROOT)} selected={best_name}")


if __name__ == "__main__":
    main()