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