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