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