Spaces:
Running on Zero
Running on Zero
Download scripts/evaluation/evaluate_flood_detection.py from ageraustine/River_Network: direct link, hf CLI and curl.
- Browser
- Download file 168 kB
-
https://huggingface.co/spaces/ageraustine/River_Network/resolve/main/scripts/evaluation/evaluate_flood_detection.py
- Command line
-
hf download hf://spaces/ageraustine/River_Network/scripts/evaluation/evaluate_flood_detection.py
-
curl -L -o evaluate_flood_detection.py https://huggingface.co/spaces/ageraustine/River_Network/resolve/main/scripts/evaluation/evaluate_flood_detection.py
168 kB
| """ | |
| Evaluates whether the currently trained spatiotemporal GNN would have | |
| flagged the real historical high-flow events identified by | |
| identify_flood_events.py. This is the test that actually matters for | |
| this project's stated purpose -- beating a naive mean-baseline says | |
| the model learned real structure, but says nothing about whether it | |
| would catch a real flood. Those are different questions, and only | |
| this one answers the second. | |
| Reuses the exact data pipeline from train_spatiotemporal_gnn.py (same | |
| subgraph selection, same standardization) rather than rebuilding it | |
| independently -- the model's saved weights only make sense against the | |
| exact input shapes and standardization statistics they were trained | |
| with, so any mismatch here would silently produce meaningless | |
| predictions rather than a clean error. | |
| For every real (anchor_date, node, horizon) combination in the held-out | |
| test split, checks whether the model's own discharge prediction -- | |
| converted back to real L/s, not left in standardized units -- exceeds | |
| that station's real event threshold (the same threshold | |
| identify_flood_events.py used to define a "high-flow day" there). | |
| Confusion matrix is computed per horizon: catching a real event a day | |
| out and catching it a week out are different questions, and averaging | |
| them together would hide which lead times the model is actually | |
| useful at. | |
| Usage: | |
| python -m scripts.evaluation.evaluate_flood_detection --data-root datasets | |
| """ | |
| import argparse | |
| from pathlib import Path | |
| from typing import Dict, List | |
| import numpy as np | |
| import pandas as pd | |
| import torch | |
| try: | |
| from scripts.training.train_spatiotemporal_gnn import ( | |
| BASIN_FILE_NAMES, HORIZONS, QUANTILES, load_basin_graph, combine_basins, | |
| build_combined_dynamic_tensors, standardize_dynamic_tensors, standardize_static_features, | |
| extract_catchment_area_km2, | |
| compute_rolling_precip_sum, fit_specific_discharge_coefficient, estimate_discharge_from_precip, | |
| build_forecast_tensor, compute_days_since_last_real, compress_days_since_last_real, | |
| get_dynamic_channel_index, ROUTING_HORIZON_COUNT, FORECAST_LEAD_TIMES, | |
| ) | |
| from src.graph.subgraph_selection import build_collapsed_subgraph | |
| from src.graph.spatiotemporal_gnn import prepare_graph_training_windows, SpatiotemporalGNN | |
| from src.graph.physics_losses import compute_per_node_historical_median | |
| from src.graph.dynamic_features import compute_forward_filled_tensor | |
| from src.graph.flood_events import ( | |
| build_event_lookup, is_real_event_day, days_to_nearest_real_event, build_station_thresholds, | |
| ) | |
| except ImportError: | |
| import sys | |
| sys.path.insert(0, str(Path(__file__).resolve().parent.parent.parent)) | |
| from scripts.training.train_spatiotemporal_gnn import ( | |
| BASIN_FILE_NAMES, HORIZONS, QUANTILES, load_basin_graph, combine_basins, | |
| build_combined_dynamic_tensors, standardize_dynamic_tensors, standardize_static_features, | |
| extract_catchment_area_km2, | |
| compute_rolling_precip_sum, fit_specific_discharge_coefficient, estimate_discharge_from_precip, | |
| build_forecast_tensor, compute_days_since_last_real, compress_days_since_last_real, | |
| get_dynamic_channel_index, ROUTING_HORIZON_COUNT, FORECAST_LEAD_TIMES, | |
| ) | |
| from src.graph.subgraph_selection import build_collapsed_subgraph | |
| from src.graph.spatiotemporal_gnn import prepare_graph_training_windows, SpatiotemporalGNN | |
| from src.graph.physics_losses import compute_per_node_historical_median | |
| from src.graph.dynamic_features import compute_forward_filled_tensor | |
| from src.graph.flood_events import ( | |
| build_event_lookup, is_real_event_day, days_to_nearest_real_event, build_station_thresholds, | |
| ) | |
| # build_event_lookup / is_real_event_day / days_to_nearest_real_event / | |
| # build_station_thresholds moved to src/graph/flood_events.py so | |
| # train_spatiotemporal_gnn.py can import the same real event-bookkeeping | |
| # logic for its own event-reweighted loss, without a circular import | |
| # (this script already imports FROM scripts.training.train_spatiotemporal_gnn, | |
| # so the reverse import isn't possible) -- see that module's own | |
| # docstring. | |
| def plot_confusion_matrices_by_horizon(confusion: dict, horizons: list, output_path: Path) -> None: | |
| """ | |
| A real 2x2 confusion matrix heatmap per horizon, arranged in a grid | |
| -- gives an immediate visual read on how tp/fp/fn/tn actually shift | |
| across lead times, rather than reading raw counts out of a printed | |
| table. | |
| """ | |
| import matplotlib.pyplot as plt | |
| n_horizons = len(horizons) | |
| n_cols = min(5, n_horizons) | |
| n_rows = (n_horizons + n_cols - 1) // n_cols | |
| fig, axes = plt.subplots(n_rows, n_cols, figsize=(3.2 * n_cols, 3.2 * n_rows)) | |
| axes = axes.flatten() if n_horizons > 1 else [axes] | |
| for idx, h in enumerate(horizons): | |
| ax = axes[idx] | |
| c = confusion[h] | |
| # rows = actual (event, no event), cols = predicted (flag, no flag) | |
| matrix = [[c["tp"], c["fn"]], [c["fp"], c["tn"]]] | |
| im = ax.imshow(matrix, cmap="Blues", vmin=0) | |
| for i in range(2): | |
| for j in range(2): | |
| value = matrix[i][j] | |
| # Text color flips for readability against the | |
| # colormap's own darker high-value cells. | |
| text_color = "white" if value > (max(max(row) for row in matrix) / 2) else "black" | |
| ax.text(j, i, str(value), ha="center", va="center", color=text_color, fontsize=11) | |
| ax.set_xticks([0, 1]) | |
| ax.set_yticks([0, 1]) | |
| ax.set_xticklabels(["flagged", "not flagged"], fontsize=8) | |
| ax.set_yticklabels(["event", "no event"], fontsize=8) | |
| ax.set_title(f"{h}-day", fontsize=10) | |
| for idx in range(n_horizons, len(axes)): | |
| axes[idx].axis("off") | |
| fig.suptitle("Flood-detection confusion matrix by lead time (real held-out test data)") | |
| fig.tight_layout() | |
| fig.savefig(output_path, dpi=150) | |
| plt.close(fig) | |
| def plot_precision_recall_by_horizon(result_df: "object", output_path: Path) -> None: | |
| """ | |
| Precision, recall, and the real base rate, all against horizon on | |
| one plot -- deliberately including base rate directly rather than | |
| leaving it as a separate manual calculation, since this project has | |
| repeatedly found that a horizon's precision only means something | |
| when read against how often a real event actually occurs there (a | |
| horizon with very few real events can show inflated-looking | |
| precision/recall from a handful of lucky flags). | |
| """ | |
| import matplotlib.pyplot as plt | |
| horizons = result_df["horizon_days"].values | |
| precision = result_df["precision"].values | |
| recall = result_df["recall"].values | |
| base_rate = (result_df["tp"] + result_df["fn"]) / ( | |
| result_df["tp"] + result_df["fp"] + result_df["fn"] + result_df["tn"] | |
| ) | |
| fig, ax = plt.subplots(figsize=(9, 5)) | |
| x_pos = range(len(horizons)) # evenly spaced positions -- HORIZONS jumps unevenly | |
| # (5 -> 10 -> 30 -> 60 -> 90 -> 180), so a real numeric x-axis would | |
| # compress the 1-10 day range (where most of the real signal lives) | |
| # into a small sliver of the plot. | |
| ax.plot(x_pos, precision, marker="o", label="precision", color="#2563eb") | |
| ax.plot(x_pos, recall, marker="o", label="recall", color="#16a34a") | |
| ax.plot(x_pos, base_rate, marker="o", label="real event base rate", color="#94a3b8", linestyle="--") | |
| ax.set_xticks(list(x_pos)) | |
| ax.set_xticklabels([str(h) for h in horizons]) | |
| ax.set_xlabel("Lead time (days)") | |
| ax.set_ylabel("Rate") | |
| ax.set_ylim(-0.02, 1.02) | |
| ax.set_title("Precision/recall vs. lead time, against the real event base rate") | |
| ax.legend(loc="best") | |
| ax.grid(True, alpha=0.3) | |
| fig.tight_layout() | |
| fig.savefig(output_path, dpi=150) | |
| plt.close(fig) | |
| def plot_single_trace( | |
| horizons: list, quantiles: list, quantile_values_real: "object", | |
| threshold: float, routing_horizon_count: int, station_code: str, anchor_str: str, | |
| output_path: Path, | |
| ) -> None: | |
| """ | |
| Real predicted discharge, every quantile, across every real | |
| horizon, for one specific (anchor, station) example -- with the | |
| real station threshold and the real ROUTING_HORIZON_COUNT boundary | |
| both marked directly on the plot. Built specifically to check for a | |
| real visual discontinuity right at the physics-loss boundary, which | |
| an aggregate statistic (mean/ratio across many examples) could | |
| smooth over or miss entirely. | |
| """ | |
| import matplotlib.pyplot as plt | |
| fig, ax = plt.subplots(figsize=(10, 6)) | |
| x_pos = range(len(horizons)) | |
| colors = {0.5: "#2563eb", 0.9: "#f59e0b", 0.95: "#dc2626", 0.99: "#7c3aed"} | |
| for q_idx, q in enumerate(quantiles): | |
| ax.plot(x_pos, quantile_values_real[:, q_idx], marker="o", | |
| label=f"quantile {q}", color=colors.get(q, "#333333")) | |
| ax.axhline(threshold, color="#16a34a", linestyle="--", linewidth=1.5, label="real station threshold") | |
| # Real ROUTING_HORIZON_COUNT boundary -- physics losses (routing, | |
| # soft-DTW, water balance) apply up through this horizon and no | |
| # further; marked as a real, vertical reference line, not just | |
| # mentioned in the title, so it's directly visible against whatever | |
| # the quantile lines actually do there. | |
| if routing_horizon_count in horizons: | |
| boundary_x = horizons.index(routing_horizon_count) | |
| ax.axvline(boundary_x, color="#94a3b8", linestyle=":", linewidth=1.5, | |
| label=f"ROUTING_HORIZON_COUNT boundary ({routing_horizon_count}d)") | |
| ax.set_xticks(list(x_pos)) | |
| ax.set_xticklabels([str(h) for h in horizons]) | |
| ax.set_xlabel("Lead time (days)") | |
| ax.set_ylabel("Discharge (real L/s)") | |
| ax.set_title(f"Real predicted discharge quantiles -- station {station_code}, anchor {anchor_str}") | |
| ax.legend(loc="best", fontsize=9) | |
| ax.grid(True, alpha=0.3) | |
| fig.tight_layout() | |
| fig.savefig(output_path, dpi=150) | |
| plt.close(fig) | |
| def compute_continuous_metrics(observed, predicted) -> dict: | |
| """ | |
| Real, standard hydrology model-evaluation metrics -- NSE, KGE, | |
| RMSE, MAE, PBIAS -- computed only on real (non-NaN) observed days, | |
| matching this project's established "missing stays missing" | |
| discipline rather than fabricating a comparison point where no real | |
| observation exists. | |
| NSE (Nash-Sutcliffe Efficiency): 1 is perfect, 0 means "no better | |
| than always predicting the real observed mean", negative means | |
| worse than that. Real formula: | |
| 1 - sum((obs-pred)^2) / sum((obs-obs_mean)^2). | |
| KGE (Kling-Gupta Efficiency): 1 is perfect. Decomposes skill into | |
| real correlation (r), real variability ratio (alpha = | |
| pred_std/obs_std), and real bias ratio (beta = pred_mean/obs_mean) | |
| -- a real, single number can hide very different real failure | |
| modes (e.g. good correlation but systematically too smooth), which | |
| is exactly why KGE is reported alongside NSE, not instead of it. | |
| RMSE, MAE: real absolute error in the same real units as the data | |
| (L/s here) -- RMSE weights large real errors more heavily than MAE | |
| does, so reporting both, not just one, shows whether error is | |
| dominated by a few large real misses or spread more evenly. | |
| PBIAS (Percent Bias): positive means the model under-predicts on | |
| average (real observed exceeds real predicted), negative means it | |
| over-predicts -- directly relevant given this station's real plot | |
| showed both patterns at different points in time (median under- | |
| predicting real peaks, 0.95 quantile over-predicting elsewhere). | |
| """ | |
| observed = np.asarray(observed, dtype=float) | |
| predicted = np.asarray(predicted, dtype=float) | |
| mask = ~np.isnan(observed) & ~np.isnan(predicted) | |
| n_real = int(mask.sum()) | |
| if n_real < 2: | |
| return {"n_real": n_real, "NSE": float("nan"), "KGE": float("nan"), | |
| "RMSE": float("nan"), "MAE": float("nan"), "PBIAS": float("nan")} | |
| obs, pred = observed[mask], predicted[mask] | |
| obs_mean = obs.mean() | |
| denom_nse = np.sum((obs - obs_mean) ** 2) | |
| nse = float(1.0 - np.sum((obs - pred) ** 2) / denom_nse) if denom_nse > 0 else float("nan") | |
| if obs.std() > 0 and pred.std() > 0 and obs_mean != 0: | |
| r = float(np.corrcoef(obs, pred)[0, 1]) | |
| alpha = float(pred.std() / obs.std()) | |
| beta = float(pred.mean() / obs_mean) | |
| kge = float(1.0 - np.sqrt((r - 1) ** 2 + (alpha - 1) ** 2 + (beta - 1) ** 2)) | |
| else: | |
| kge = float("nan") | |
| rmse = float(np.sqrt(np.mean((obs - pred) ** 2))) | |
| mae = float(np.mean(np.abs(obs - pred))) | |
| pbias = float(100.0 * np.sum(obs - pred) / np.sum(obs)) if np.sum(obs) != 0 else float("nan") | |
| return {"n_real": n_real, "NSE": nse, "KGE": kge, "RMSE": rmse, "MAE": mae, "PBIAS": pbias} | |
| def compute_station_model_metrics( | |
| dates_list: list, observed, predicted, station_code: str, event_lookup: dict, threshold: float, | |
| ) -> dict: | |
| """ | |
| Real per-station, per-model metrics for the comparison table: | |
| everything compute_continuous_metrics already gives (NSE, KGE, | |
| RMSE, MAE, PBIAS), plus Peak RMSE, POD, and FAR -- added to match | |
| the real multi-model comparison layout this project is moving to | |
| (one metrics row per model, not a single model-vs-naive text box). | |
| Peak RMSE: RMSE restricted to real dates that fall inside a real | |
| labeled flood event for this station (is_real_event_day) -- "how | |
| well does this model do specifically on the days that matter for | |
| flood detection", distinct from the whole-period RMSE above, which | |
| is dominated by the much more common quiet days by sheer count. | |
| POD/FAR: reuses the SAME real event-window definition | |
| (is_real_event_day) as this script's own main confusion matrix, | |
| not a simplified "observed > threshold" proxy -- so these numbers | |
| are directly comparable to the aggregate confusion-matrix output | |
| elsewhere in this script, not a second, subtly different | |
| definition of "event" living only on this plot. A real day counts | |
| as a predicted event whenever this model's own prediction reaches | |
| that station's real threshold (station_thresholds), matching the | |
| same flagging rule the confusion matrix uses. | |
| Returns NaN for Peak RMSE/POD/FAR (not an error) when this station | |
| has no real labeled events at all, or none fall in the real dates | |
| being plotted -- a real, meaningful "not evaluable on this axis" | |
| rather than a misleading 0. | |
| """ | |
| base = compute_continuous_metrics(observed, predicted) | |
| observed_arr = np.asarray(observed, dtype=float) | |
| predicted_arr = np.asarray(predicted, dtype=float) | |
| dates_arr = pd.DatetimeIndex(dates_list) | |
| is_event = np.array([is_real_event_day(station_code, d, event_lookup) is not None for d in dates_arr]) | |
| valid = ~np.isnan(observed_arr) & ~np.isnan(predicted_arr) | |
| peak_mask = valid & is_event | |
| if peak_mask.sum() >= 2: | |
| peak_rmse = float(np.sqrt(np.mean((observed_arr[peak_mask] - predicted_arr[peak_mask]) ** 2))) | |
| else: | |
| peak_rmse = float("nan") | |
| if valid.sum() >= 1 and np.isfinite(threshold): | |
| predicted_flag = valid & (predicted_arr >= threshold) | |
| tp = int((predicted_flag & is_event & valid).sum()) | |
| fn = int((~predicted_flag & is_event & valid).sum()) | |
| fp = int((predicted_flag & ~is_event & valid).sum()) | |
| pod = tp / (tp + fn) if (tp + fn) > 0 else float("nan") | |
| far = fp / (fp + tp) if (fp + tp) > 0 else float("nan") | |
| else: | |
| pod, far = float("nan"), float("nan") | |
| base.update({"Peak_RMSE": peak_rmse, "POD": pod, "FAR": far, | |
| "Status": "EVALUABLE" if base["n_real"] >= 2 else "INSUFFICIENT DATA"}) | |
| return base | |
| def plot_station_comparison( | |
| dates_list: list, real_observed: "object", model_predictions: Dict[str, "object"], | |
| station_code: str, horizon: int, output_path: Path, | |
| model_metrics: Dict[str, dict] = None, threshold: float = None, train_q90: float = None, | |
| split_label: str = "VALIDATION", basin_label: str = None, | |
| ) -> None: | |
| """ | |
| Real observed discharge vs. one or more real models' real | |
| predictions, laid out as: a full-width real time-series panel on | |
| top, a real observed-vs-simulated 1:1 scatter panel on the bottom | |
| left, and a real per-model metrics table on the bottom right -- | |
| replacing plot_station_timeseries's single-model-vs-naive-baseline | |
| layout with one built for a real, growing roster of models | |
| (cnn_lstm_gnn, gru_gcn, lstm_gcn, stgnn, hydro_pihgnn, ...), not | |
| just this project's one ST-GNN checkpoint. | |
| No naive baseline anywhere in this layout, by real, explicit | |
| request -- the real comparison that matters now is model vs. model | |
| vs. real observed, not model vs. a constant-value strawman. | |
| model_predictions: {model_name: [values...]}, one real point- | |
| prediction series per model, same real dates/order as | |
| real_observed. A model with QUANTILE output (like this project's | |
| own ST-GNN) should pass its real median here -- this layout is | |
| deliberately point-prediction-first, matching the real screenshot | |
| it was built from, not this project's earlier quantile-fan style; | |
| that quantile detail is better read from --check-quantile- | |
| calibration's own real numbers than crowded onto this plot | |
| alongside several other real models' lines. | |
| model_metrics: {model_name: dict}, each real dict from | |
| compute_station_model_metrics (NSE, KGE, RMSE, MAE, PBIAS, | |
| Peak_RMSE, POD, FAR, Status, n_real) -- rendered as one real table | |
| row per model, not annotated text, so it stays readable regardless | |
| of how many real models are being compared. Falls back to | |
| compute_continuous_metrics's smaller real dict shape gracefully | |
| (missing keys render as "--", not a crash), so this still works | |
| before every real model has POD/FAR-capable metrics wired up. | |
| threshold/train_q90: optional real horizontal reference lines on | |
| the top panel only (this project's real per-station flood | |
| threshold, and/or the real training-period q90, matching the real | |
| screenshot's own "Train q90=... m3/s" dashed line) -- both None by | |
| default, since not every real caller has both readily available. | |
| """ | |
| import matplotlib.pyplot as plt | |
| from matplotlib.gridspec import GridSpec | |
| model_names = list(model_predictions.keys()) | |
| palette = plt.cm.tab10.colors | |
| fig = plt.figure(figsize=(14, 10)) | |
| gs = GridSpec(2, 2, height_ratios=[1.3, 1.0], figure=fig, hspace=0.55, wspace=0.28) | |
| ax_ts = fig.add_subplot(gs[0, :]) | |
| ax_scatter = fig.add_subplot(gs[1, 0]) | |
| ax_table = fig.add_subplot(gs[1, 1]) | |
| ax_table.axis("off") | |
| # --- Top: full-width time series, observed + every real model --- | |
| ax_ts.plot(dates_list, real_observed, color="#1d4ed8", linewidth=2.0, | |
| label="Observed", zorder=10) | |
| for i, name in enumerate(model_names): | |
| ax_ts.plot(dates_list, model_predictions[name], color=palette[i % len(palette)], | |
| linewidth=1.2, alpha=0.9, label=name) | |
| if train_q90 is not None and np.isfinite(train_q90): | |
| ax_ts.axhline(train_q90, color="#0ea5e9", linestyle="--", linewidth=1.2, | |
| label=f"Train q90={train_q90:.2f}") | |
| if threshold is not None and np.isfinite(threshold): | |
| ax_ts.axhline(threshold, color="#7c3aed", linestyle="--", linewidth=1.2, | |
| label="Real station threshold") | |
| basin_part = f"{basin_label} | " if basin_label else "" | |
| ax_ts.set_title(f"{split_label} | {basin_part}{station_code} (gauged) | horizon {horizon} d") | |
| ax_ts.set_ylabel("Discharge (m³/s)") | |
| ax_ts.legend(loc="upper left", fontsize=8.5, ncol=3, framealpha=0.9) | |
| ax_ts.grid(True, alpha=0.25) | |
| # --- Bottom-left: observed vs. simulated, 1:1 scatter --- | |
| real_observed_arr = np.asarray(real_observed, dtype=float) | |
| finite_vals = [real_observed_arr[~np.isnan(real_observed_arr)]] if np.isfinite(real_observed_arr).any() else [] | |
| for i, name in enumerate(model_names): | |
| pred_arr = np.asarray(model_predictions[name], dtype=float) | |
| mask = ~np.isnan(real_observed_arr) & ~np.isnan(pred_arr) | |
| ax_scatter.scatter(real_observed_arr[mask], pred_arr[mask], s=10, alpha=0.35, | |
| color=palette[i % len(palette)], label=name, edgecolors="none") | |
| if mask.any(): | |
| finite_vals.append(pred_arr[mask]) | |
| if finite_vals: | |
| lo = min(float(np.nanmin(v)) for v in finite_vals if len(v)) | |
| hi = max(float(np.nanmax(v)) for v in finite_vals if len(v)) | |
| ax_scatter.plot([lo, hi], [lo, hi], color="black", linestyle="--", linewidth=1.2, zorder=1) | |
| ax_scatter.set_title("Observed vs simulated (1:1)") | |
| ax_scatter.set_xlabel("Observed (m³/s)") | |
| ax_scatter.set_ylabel("Simulated (m³/s)") | |
| ax_scatter.legend(loc="upper left", fontsize=8) | |
| ax_scatter.grid(True, alpha=0.25) | |
| # --- Bottom-right: one metrics row per real model --- | |
| ax_table.set_title("Hydrological diagnostics", fontsize=11, pad=14) | |
| columns = ["Model", "NSE", "KGE", "RMSE", "PBIAS%", "Peak RMSE", "POD", "FAR", "Status"] | |
| rows = [] | |
| for name in model_names: | |
| m = (model_metrics or {}).get(name, {}) | |
| def cell(key, fmt="{:.3f}"): | |
| v = m.get(key) | |
| return fmt.format(v) if v is not None and np.isfinite(v) else "--" | |
| rows.append([ | |
| name, cell("NSE"), cell("KGE"), cell("RMSE", "{:.3f}"), cell("PBIAS", "{:.3f}"), | |
| cell("Peak_RMSE", "{:.3f}"), cell("POD", "{:.3f}"), cell("FAR", "{:.3f}"), | |
| m.get("Status", "--"), | |
| ]) | |
| if rows: | |
| tbl = ax_table.table(cellText=rows, colLabels=columns, loc="center", cellLoc="center") | |
| tbl.auto_set_font_size(False) | |
| tbl.set_fontsize(8.5) | |
| # Column widths sized to their real content (not left equal- | |
| # width, matplotlib's own default) -- equal widths were | |
| # confirmed directly to overlap real neighboring cells once | |
| # "Status" (EVALUABLE/INSUFFICIENT DATA) sat next to several | |
| # narrow numeric columns in the same real table. | |
| tbl.auto_set_column_width(col=list(range(len(columns)))) | |
| tbl.scale(1.0, 1.8) | |
| # NOT fig.autofmt_xdate() -- confirmed directly as the real cause | |
| # of a real reported bug ("can't see date on first plot"): | |
| # autofmt_xdate() assumes every real axes in the figure shares one | |
| # x-axis stacked in rows, and HIDES x tick labels (and clears the | |
| # xlabel) on every real axes that isn't in the last row. Here, | |
| # ax_ts is row 0 and ax_scatter/ax_table (whose x-axes are | |
| # "Observed (m3/s)" and nothing, not dates) are row 1, so | |
| # autofmt_xdate() was silently hiding ax_ts's real date labels | |
| # entirely, thinking it was an interior row of a shared-x stack. | |
| # Rotating ax_ts's own real tick labels directly, instead, is the | |
| # real fix -- scoped to the one real axes that actually has dates. | |
| for label in ax_ts.get_xticklabels(): | |
| label.set_rotation(30) | |
| label.set_ha("right") | |
| ax_ts.set_xlabel("Date") | |
| fig.savefig(output_path, dpi=150, bbox_inches="tight") | |
| plt.close(fig) | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description="Evaluate flood-event detection against the trained model") | |
| parser.add_argument("--data-root", type=Path, default=Path("datasets")) | |
| parser.add_argument("--model-path", type=Path, default=None, | |
| help="Defaults to <data-root>/spatiotemporal_gnn/model.pt") | |
| parser.add_argument("--events-path", type=Path, default=None, | |
| help="Defaults to <data-root>/flood_events_discharge.csv") | |
| parser.add_argument("--lookback-days", type=int, default=30) | |
| parser.add_argument("--train-end", type=str, default="2021-12-31") | |
| parser.add_argument("--val-end", type=str, default="2023-12-31") | |
| parser.add_argument("--test-end", type=str, default="2026-12-31") | |
| parser.add_argument("--max-nodes", type=int, default=100) | |
| parser.add_argument("--temporal-type", type=str, default="gru", choices=["gru", "lstm", "cnn"], | |
| help="MUST match the real --temporal-type the checkpoint being evaluated was " | |
| "actually trained with -- a mismatch means model.load_state_dict() below " | |
| "fails (or worse, silently loads the wrong shapes) since the temporal " | |
| "encoder's own real parameters differ by type. Not auto-detected from the " | |
| "checkpoint file itself; the caller must know and pass this correctly.") | |
| parser.add_argument("--flag-quantile", type=float, default=0.95, | |
| help="Which of the model's predicted quantiles counts as a real flood flag, " | |
| "for any horizon not covered by --flag-quantile-overrides. Defaults to " | |
| "0.95, matching identify_flood_events.py's own default percentile " | |
| "threshold -- a coherent, matched comparison.") | |
| parser.add_argument("--flag-quantile-overrides", type=str, default=None, | |
| help="Per-horizon overrides, format 'h1=q1,h2=q2,...', e.g. " | |
| "'1=0.9,2=0.9,3=0.9,4=0.9,5=0.9'. Real, measured evidence in this " | |
| "project shows the right operating point genuinely differs by horizon: " | |
| "at short horizons (1-5 days) here, a real predicted quantile is a real " | |
| "VALUE compared against a fixed real threshold -- since higher quantiles " | |
| "are, by the model's own monotonicity guarantee, always larger numbers, " | |
| "a HIGHER quantile makes the threshold EASIER to exceed, not harder (the " | |
| "opposite of what \"more confidence required\" might suggest). Lowering " | |
| "toward 0.9 measurably improved precision at short horizons in real " | |
| "testing (e.g. day 1: 63.8%% -> 67.9%%, day 5: 9.5%% -> 42.4%%) but " | |
| "collapsed recall to near-zero at longer horizons (day 30 recall hit " | |
| "0%% at 0.9, back when 30 was still a real predicted horizon -- HORIZONS " | |
| "has since been narrowed to max out at 21 days, but the real underlying " | |
| "finding stands) -- there's real evidence against one global value being " | |
| "right for every horizon. Horizons not listed here fall back to " | |
| "--flag-quantile.") | |
| parser.add_argument("--check-fp-precipitation", action="store_true", | |
| help="Diagnostic: compares real recent precipitation (summed over the last " | |
| "ROUTING_HORIZON_COUNT real days of each example's own input window) " | |
| "across false positives vs. true negatives. Checked a real, specific " | |
| "hypothesis -- that water_balance_loss, with ET hardcoded to 0.0 and no " | |
| "storage term, has no way to represent real infiltration/storage " | |
| "absorbing part of a rain event, and may be pushing the model toward " | |
| "over-predicting discharge whenever real precipitation is high. Real " | |
| "result from actual testing: the difference was real but tiny (8.80mm " | |
| "vs 8.67mm mean, identical medians) -- too small to explain the real " | |
| "precision problem, and fp counts don't track precipitation the way " | |
| "they'd need to for this to be the main driver. Kept as a real, useful " | |
| "check for future runs, but this specific hypothesis is not well " | |
| "supported by the evidence so far -- see --check-quantile-spread for the " | |
| "hypothesis that superseded it.") | |
| parser.add_argument("--check-quantile-spread", action="store_true", | |
| help="Diagnostic: reports the real, un-standardized gap between the model's " | |
| "predicted 0.95 and 0.5 (median) discharge quantiles, per horizon. Real " | |
| "result from actual clean testing (single global --flag-quantile, no " | |
| "per-horizon overrides -- an earlier run with overrides active was " | |
| "confounded, since switching the flagging quantile itself at the same " | |
| "horizon boundary being examined can't be cleanly separated from a real " | |
| "model-behavior effect): the spread grows smoothly and continuously with " | |
| "horizon (~1.3-1.5x per step throughout), but real false positives jump " | |
| "sharply at one specific point (day 4->5: 74->314 in one real run) that " | |
| "the smooth spread growth does NOT explain (day 4->5 spread only grew " | |
| "~1.3x, same as every other step) -- this specific hypothesis is NOT well " | |
| "supported by the real evidence. The real day 4->5 jump instead lines up " | |
| "with day 5 being one of the few horizons with real ECMWF forecast " | |
| "coverage while day 4 has none -- see --check-forecast-fp-correlation for " | |
| "the hypothesis that superseded this one.") | |
| parser.add_argument("--check-forecast-fp-correlation", action="store_true", | |
| help="Diagnostic: for horizons with real ECMWF forecast coverage (currently " | |
| "[1,2,3,5,7], from FORECAST_LEAD_TIMES), reports what fraction of real " | |
| "false positives at that horizon had real (non-missing) forecast " | |
| "precipitation available, vs. the real base rate of forecast availability " | |
| "across every real example at that horizon. Real result from actual " | |
| "testing: a real but weak correlation (ratios 1.06-1.15x, all below the " | |
| "1.5x flag threshold) that does NOT single out day 5 specifically (day 1-3 " | |
| "showed similar ratios) -- this specific hypothesis is NOT well supported " | |
| "either. See --dump-fp-details for direct inspection of the real false " | |
| "positives themselves, which superseded this and the two hypotheses " | |
| "before it (--check-fp-precipitation, --check-quantile-spread).") | |
| parser.add_argument("--dump-fp-details", type=int, default=None, | |
| help="Exploratory, not hypothesis-testing: dumps every real false positive at " | |
| "the given horizon (station, date, predicted value, real threshold, " | |
| "margin, days to the nearest real event) to a CSV, plus a quick printed " | |
| "summary -- for direct inspection after three real, specific mechanisms " | |
| "(precipitation, quantile-spread, forecast-availability) were each tested " | |
| "and none explained the sharp real fp jump at day 5. Pass the horizon to " | |
| "dump, e.g. '--dump-fp-details 5'.") | |
| parser.add_argument("--dump-fp-output", type=Path, default=None, | |
| help="Defaults to <data-root>/evaluation/fp_details_horizon_<h>.csv") | |
| parser.add_argument("--plot-single-trace", type=str, default=None, | |
| help="Exploratory: plots one real (anchor, station) example's predicted " | |
| "discharge quantiles across all real horizons, with the station's real " | |
| "threshold and ROUTING_HORIZON_COUNT's boundary both marked -- a genuinely " | |
| "different check than the aggregate statistics tried so far (precipitation " | |
| "correlation, quantile-spread growth, forecast availability, forecast date " | |
| "alignment; none fully explained the real day-5 false-positive pattern). A " | |
| "visual discontinuity right at the physics-loss boundary would be " | |
| "immediately obvious here in a way an aggregate number isn't. Format: " | |
| "'station_code=anchor_date', e.g. 'H605641201=2026-01-01'.") | |
| parser.add_argument("--plot-stations", type=str, default=None, | |
| help="Comma-separated real station codes, e.g. 'H605641201,H404021101' -- for " | |
| "each, plots real observed discharge, the model's real predicted median + " | |
| "0.95 quantile, and a real per-station naive baseline (that station's own " | |
| "training-period empirical median/0.95, not the global naive baseline used " | |
| "for the overall skill comparison -- a station-specific comparison is more " | |
| "meaningful here), all against real test-period dates, at one fixed " | |
| "horizon (--plot-stations-horizon).") | |
| parser.add_argument("--plot-stations-horizon", type=int, default=1, | |
| help="Which real horizon to plot station time series at -- defaults to 1 (the " | |
| "horizon with the strongest real, established skill in this project).") | |
| parser.add_argument("--model-name", type=str, default="stgnn", | |
| help="This checkpoint's real display name/legend label/table row on " | |
| "--plot-stations' comparison plot -- e.g. 'stgnn', 'cnn_lstm_gnn'. The " | |
| "PRIMARY checkpoint's name -- flood-flagging, calibration, and every other " | |
| "diagnostic in this script besides --plot-stations still run against this " | |
| "one model only; see --extra-model to add more real checkpoints to the " | |
| "comparison PLOT specifically.") | |
| parser.add_argument("--no-auto-extra-models", action="store_true", | |
| help="By default, --plot-stations auto-detects every OTHER real trained " | |
| "temporal_type sitting next to the primary checkpoint -- i.e. it checks " | |
| "<data-root>/spatiotemporal_gnn/model.pt (gru), <data-root>/" | |
| "spatiotemporal_gnn/lstm/model.pt, and <data-root>/spatiotemporal_gnn/cnn/" | |
| "model.pt (whichever of these ISN'T the primary --temporal-type), and adds " | |
| "each real checkpoint found there to the comparison plot automatically, " | |
| "named after its own temporal_type (e.g. 'lstm', 'cnn') -- no flag needed to " | |
| "pick them up once train_spatiotemporal_gnn.py has saved them. Pass this " | |
| "flag to disable that and plot the primary checkpoint alone.") | |
| parser.add_argument("--extra-model", action="append", default=None, metavar="NAME=TEMPORAL_TYPE[=PATH]", | |
| help="Adds another real trained checkpoint's own real predictions to --plot-" | |
| "stations' comparison plot, ON TOP OF whatever --no-auto-extra-models' own " | |
| "real auto-detection already found -- only needed for a real checkpoint " | |
| "living outside the standard <data-root>/spatiotemporal_gnn/<temporal_type>/ " | |
| "layout, or under a different real display name. e.g. '--extra-model " | |
| "lstm_v2=lstm=datasets/lstm_v2/model.pt'. Repeatable. Format is " | |
| "'name=temporal_type' (path defaults the same way auto-detection does) or " | |
| "'name=temporal_type=path' for a real checkpoint anywhere else. Each real " | |
| "extra model (auto-detected or explicit) is scored/blended with the exact " | |
| "same --blend-quantiles/--blend-mode/--blend-upper-quantile settings as the " | |
| "primary model, so the comparison stays apples-to-apples. Only affects " | |
| "--plot-stations -- flood-flagging, --check-quantile-calibration and every " | |
| "other diagnostic still run against the primary checkpoint only.") | |
| parser.add_argument("--blend-quantiles", action="store_true", | |
| help="Instead of plotting/scoring the raw 0.95 quantile, plot/score a real " | |
| "interplay between the 0.5 (median) and 0.95 quantiles: predicted = " | |
| "median + alpha*(q95-median), where alpha = spread/(spread+tau) and " | |
| "spread = max(q95-median, 0) for that real (station, date). alpha -> 0 " | |
| "(pure median) when q95 and median are close together, alpha -> 1 (pure " | |
| "q95) when they're far apart -- a real, continuous saturating blend, not a " | |
| "hard switch. tau sets the spread scale at which alpha=0.5; see " | |
| "--blend-scale. Affects --plot-stations' plotted series/metrics and, when " | |
| "--trace-fp-inputs is also set, its false-spike trace and the low/high-flow " | |
| "stratified check -- all switch from pure q0.95 to this blended series " | |
| "together so they stay a real, consistent comparison.") | |
| parser.add_argument("--blend-upper-quantile", type=float, default=0.95, | |
| help="Which real predicted quantile --blend-quantiles uses as its upper anchor " | |
| "-- predicted = median + alpha*(q_upper-median) -- instead of always q0.95. " | |
| "Must be one of this model's real predicted quantiles (0.9/0.95/0.99 here). " | |
| "Real testing via --compare-all-quantiles showed q0.9 ALONE (no blending at " | |
| "all) already beats both q0.95 alone and the median/q0.95 blend on NSE/RMSE/" | |
| "PBIAS at station H404021101 (NSE=0.874 vs 0.861 blended vs 0.838 pure " | |
| "q0.95), at a real cost of POD dropping from 1.000 to 0.867 -- so a median/" | |
| "q0.9 blend (set this to 0.9) is a real, concrete next thing worth testing " | |
| "instead of assuming q0.95 is always the right upper anchor.") | |
| parser.add_argument("--blend-scale", type=float, default=None, | |
| help="tau in the --blend-quantiles formula, in m3/s. If omitted, tau defaults " | |
| "to that station's own real median (q95-median) spread over its real " | |
| "plotted (test) period -- a real, per-station, data-driven scale, not a " | |
| "fixed global constant. Set this explicitly to compare stations on the " | |
| "same real scale, or if a station's own median spread is degenerate " | |
| "(e.g. zero or near-zero).") | |
| parser.add_argument("--blend-mode", type=str, default="spread", choices=["spread", "level"], | |
| help="How --blend-quantiles picks alpha (the real median/q0.95 mixing weight; " | |
| "the blended VALUE is always median + alpha*(q95-median) either way). " | |
| "'spread' (default, original behavior): alpha = spread/(spread+tau), driven " | |
| "by the model's OWN real (q0.95-median) gap for that exact (station, date, " | |
| "horizon) -- alpha->1 whenever the model's two quantiles happen to disagree " | |
| "a lot, including on an ordinary day if the model's internal spread just " | |
| "widens there. 'level' instead drives alpha from how ELEVATED the model's " | |
| "own real median PREDICTION itself is, relative to that station's own real " | |
| "training-period flow range: alpha = clip((median-train_median)/(train_q90-" | |
| "train_median), 0, 1). A real --plot-stations comparison (median vs. q0.5/" | |
| "0.9/0.95/0.99, station H420021012) showed the real median tracking real " | |
| "observed discharge closely through ordinary flow and only really " | |
| "undershooting during the two real tallest peaks -- 'level' targets that " | |
| "real pattern directly: alpha stays near 0 (pure median) through ordinary " | |
| "real flow, the same real regime the median already handles well, and only " | |
| "climbs toward 1 (leaning on q0.95) once the real median prediction itself " | |
| "enters that station's own real high-flow range -- rather than reacting to " | |
| "the model's internal spread, which the real evidence so far (--blend-" | |
| "quantiles' first real run) showed can spike on ordinary quiet days too, " | |
| "not only real peaks. 'level' also needs no preliminary real tau-fitting " | |
| "forward pass -- train_median/train_q90 come straight from real training " | |
| "data -- so it's cheaper to run than 'spread'; --blend-scale is ignored " | |
| "in this mode.") | |
| parser.add_argument("--blend-quantiles-max-horizon", type=int, default=None, | |
| help="Restricts --blend-quantiles' effect on FLOOD-FLAGGING (the confusion " | |
| "matrix / --dump-fp-details, not --plot-stations' own comparison plot) to " | |
| "real horizons at or below this value; longer real horizons fall back to " | |
| "the raw --flag-quantile selection for flagging instead. Real testing on " | |
| "this project's own data showed the blend cuts real false alarms sharply " | |
| "at short horizons with almost no real recall cost, but from around h=7 " | |
| "onward it also suppresses real true positives (the model's own predicted " | |
| "spread compresses at long lead times during real events, not just quiet " | |
| "ones, so alpha collapses toward the median then too) -- a real, substantial " | |
| "recall loss at h=14/21 specifically. Leave unset to blend every horizon " | |
| "(the original, real all-or-nothing behavior); set e.g. 5 to blend only " | |
| "h in {1,2,3,4,5} and keep h in {7,10,14,21} on raw q0.95. If " | |
| "--plot-stations-horizon itself is above this cutoff, --plot-stations' own " | |
| "comparison plot/metrics fall back to raw q0.95 too, so the plot and the " | |
| "confusion matrix never silently disagree about which real horizons are " | |
| "blended.") | |
| parser.add_argument("--compare-all-quantiles", action="store_true", | |
| help="For each real station in --plot-stations, at --plot-stations-horizon: " | |
| "prints Peak_RMSE/POD/FAR/NSE/PBIAS for EVERY one of this model's real " | |
| "predicted quantiles (0.5/0.9/0.95/0.99, not just whichever one --flag-" | |
| "quantile or --blend-quantiles picked), each compared against that " | |
| "station's own real threshold, side by side -- a direct, real answer to " | |
| "'which quantile actually catches peaks', instead of reading it off a " | |
| "single plot showing only one or two of them. Uses the exact real quantile " | |
| "predictions already cached for --plot-stations (station_timeseries_data), " | |
| "so this needs no extra real model forward pass.") | |
| parser.add_argument("--trace-fp-inputs", type=int, default=None, | |
| help="Exploratory: for each real station in --plot-stations, identifies the N " | |
| "real dates with the largest real gap between the predicted 0.95 quantile " | |
| "and real observed discharge, among dates where the real observation was " | |
| "at or below that station's own real median (the specific real pattern a " | |
| "real plot showed: the 0.95 quantile spiking repeatedly while real observed " | |
| "discharge stayed flat and low) -- then traces the real input features " | |
| "(recent real discharge/precipitation, their real missingness/days-since-" | |
| "last-real-observation channels, and real forecast precipitation) the model " | |
| "actually saw for each one, to check for a real, repeating artifact in the " | |
| "real input data rather than guessing from a plot.") | |
| parser.add_argument("--ablate-missingness-flag", action="store_true", | |
| help="Diagnostic, added to directly answer the question 'is the discharge " | |
| "missingness flag alone reaching the model with meaningful causal effect, " | |
| "or is the model reacting almost entirely to the filled discharge value's " | |
| "own magnitude regardless of whether it was flagged missing?'. Requires " | |
| "--trace-fp-inputs and --plot-stations to also be set (it reuses the same " | |
| "traced false-spike examples). For the single worst false-spike example " | |
| "per requested station (rank 0 -- the largest real gap between predicted " | |
| "0.95 and real observed, among low-flow days), runs three forward passes " | |
| "through the SAME already-loaded model: (1) unmodified, (2) with the " | |
| "discharge missingness flag flipped 0 for that one node across the " | |
| "lookback window (fill VALUE left untouched), and (3) with both the " | |
| "missingness flag AND the days-since-last-real-observation channel zeroed " | |
| "for that node (fully presenting the fill value as if it were a genuine, " | |
| "same-day real observation). If (2)/(3) barely move the predicted 0.95 " | |
| "quantile away from (1), that's direct evidence the flag has little to no " | |
| "causal effect and the false spike is driven by the fill value's own " | |
| "magnitude, not by the model failing to notice it's a fill. If (2)/(3) move " | |
| "the prediction substantially toward the real observed value, the flag IS " | |
| "doing real work and the false spike is more about the fill value's " | |
| "magnitude being implausible than about the flag being ignored.") | |
| parser.add_argument("--check-quantile-calibration", action="store_true", | |
| help="Diagnostic, added after the forward-fill fix moved the worst real " | |
| "false-spike cases (per --trace-fp-inputs) away from fill-value gaps and " | |
| "onto fully-observed real dates with a recent real precipitation pulse " | |
| "(e.g. H404021101, 2026-06-03: real observed stayed at 1623 L/s, " | |
| "missingness flag 0.0 throughout, yet predicted 0.95=4298.7 L/s after " | |
| "11.3mm of real rain the day before) -- i.e. the remaining problem looks " | |
| "like a genuine quantile-calibration issue on real data, not a fill " | |
| "artifact. This checks that directly two ways. (1) Global calibration: " | |
| "across every real (gauge, date, horizon) combination in the held-out test " | |
| "set, reports empirical coverage per quantile level -- the real fraction of " | |
| "real observed values at or below that quantile's prediction -- against the " | |
| "nominal target (e.g. the 0.95 quantile should cover ~95% of real " | |
| "observations; a much higher real empirical coverage means Q0.95 is " | |
| "systematically too wide/high, not just occasionally). (2) Precipitation " | |
| "stratification: restricted to real low-flow days only (observed at or " | |
| "below that station's own real training-period median, the same " | |
| "restriction --trace-fp-inputs already uses), splits the real Q0.95 " | |
| "coverage and the real (predicted 0.95 - observed) gap by whether real " | |
| "precipitation summed over the preceding ROUTING_HORIZON_COUNT real days " | |
| "exceeded --calibration-precip-threshold-mm. A real, substantially larger " | |
| "gap and higher over-coverage on the elevated-precipitation side is direct " | |
| "evidence a recent rain pulse is a systematic driver of this project's " | |
| "remaining false-spike problem, not a coincidence of the handful of " | |
| "examples --trace-fp-inputs happened to surface.") | |
| parser.add_argument("--calibration-precip-threshold-mm", type=float, default=5.0, | |
| help="Real precipitation (mm, summed over the preceding ROUTING_HORIZON_COUNT " | |
| "real days) above which --check-quantile-calibration classifies a real " | |
| "low-flow day as 'elevated precipitation' rather than 'quiet'. Default 5.0mm " | |
| "chosen as a real, moderate rain event -- well above typical real background " | |
| "noise in this dataset's precipitation channel but well below a real flood-" | |
| "triggering event, so it isolates the specific 'moderate rain, no real flow " | |
| "response yet' regime the false-spike examples showed.") | |
| parser.add_argument("--calibration-precip-staleness-days", type=float, default=30.0, | |
| help="Further splits --check-quantile-calibration's 'elevated precipitation' " | |
| "low-flow bucket by how stale the real precipitation reading behind it is -- " | |
| "added after --trace-fp-inputs surfaced two real false-spike examples " | |
| "(2026-05-16, 2026-05-02) whose real 'recent precipitation' input was a " | |
| "forward-filled value from over a real year earlier (the real " | |
| "days-since-last-real-precipitation-observation channel sat exactly at " | |
| "compress_days_since_last_real's real 365-day sentinel ceiling, " | |
| "log1p(365)=5.903, not a genuinely recent reading) -- i.e. the 'elevated " | |
| "precipitation' label may itself be a stale-forward-fill artifact, not real " | |
| "recent rain. This checks that directly: within the elevated-precipitation, " | |
| "low-flow bucket, splits by whether the real days-since-last-real-" | |
| "precipitation-observation at the most recent real lookback day exceeds this " | |
| "many real days ('stale') or not ('recent'), and reports each side's real " | |
| "Q0.95 gap and coverage separately, plus what real fraction of the elevated-" | |
| "precipitation bucket is stale. Default 30.0 real days -- well past any real " | |
| "short reporting gap, but well short of the real 365-day sentinel ceiling, so " | |
| "it cleanly separates 'genuinely recent rain' from 'a stale forward-filled " | |
| "reading, regardless of whether it happens to have hit the ceiling yet.'") | |
| parser.add_argument("--exclude-stale-precip-tail", action="store_true", | |
| help="Drops every real test example whose anchor date falls at or after this " | |
| "run's real structurally-dead precipitation cutoff (real precipitation-data-" | |
| "coverage end date + --calibration-precip-staleness-days), from EVERY " | |
| "diagnostic in this script, not just the precipitation-specific ones -- " | |
| "added after a real run showed real precipitation coverage stopping 119 " | |
| "real days before the real test period ended, meaning any example anchored " | |
| "in that suffix has a structurally stale precipitation input for EVERY real " | |
| "station, not just ones with their own reporting gap. Without this, that " | |
| "suffix silently degrades every stat this script reports (calibration " | |
| "coverage, the confusion matrix, NSE/RMSE/PBIAS) with a data-availability " | |
| "artifact that has nothing to do with real model quality. No-op if real " | |
| "precipitation coverage already extends through the real test period.") | |
| args = parser.parse_args() | |
| # Every real output this script produces (CSVs, every plot) lives | |
| # under this dedicated subfolder, kept organized and separate from | |
| # the real input data (model checkpoint, real events, real forecast | |
| # archive) this script reads FROM, which stay at their existing | |
| # real locations under data_root directly. | |
| eval_output_dir = args.data_root / "evaluation" | |
| eval_output_dir.mkdir(parents=True, exist_ok=True) | |
| # Same real, per-temporal_type storage logic as training -- "gru" | |
| # keeps the original, unchanged default path; "lstm"/"cnn" look in | |
| # their own real subfolder, matching exactly where training saved | |
| # them. Only applies when --model-path isn't explicitly given. | |
| default_model_dir = args.data_root / "spatiotemporal_gnn" | |
| if args.temporal_type != "gru": | |
| default_model_dir = default_model_dir / args.temporal_type | |
| model_path = args.model_path or (default_model_dir / "model.pt") | |
| events_path = args.events_path or (args.data_root / "flood_events_discharge.csv") | |
| # --extra-model: parsed here (real, once, before the real data | |
| # pipeline runs) so a real typo/bad temporal_type is caught early | |
| # rather than after the real per-example loop below has already run. | |
| # Actual real checkpoint LOADING happens later, only if | |
| # --plot-stations produced real requested_stations to use them on. | |
| extra_model_specs: List[tuple] = [] | |
| for spec in (args.extra_model or []): | |
| parts = spec.split("=") | |
| if len(parts) == 2: | |
| extra_name, extra_temporal_type = parts | |
| extra_path = None | |
| elif len(parts) == 3: | |
| extra_name, extra_temporal_type, extra_path = parts | |
| else: | |
| print(f"--extra-model '{spec}' isn't 'name=temporal_type' or 'name=temporal_type=path' " | |
| "-- skipping.") | |
| continue | |
| if extra_temporal_type not in ("gru", "lstm", "cnn"): | |
| print(f"--extra-model '{spec}': temporal_type '{extra_temporal_type}' isn't gru/lstm/cnn " | |
| "-- skipping.") | |
| continue | |
| if extra_path is None: | |
| extra_default_dir = args.data_root / "spatiotemporal_gnn" | |
| if extra_temporal_type != "gru": | |
| extra_default_dir = extra_default_dir / extra_temporal_type | |
| extra_path = extra_default_dir / "model.pt" | |
| else: | |
| extra_path = Path(extra_path) | |
| extra_model_specs.append((extra_name, extra_temporal_type, extra_path)) | |
| # Auto-detection (default on): checks the two OTHER real | |
| # temporal_types' own standard save locations next to the primary | |
| # checkpoint -- <data-root>/spatiotemporal_gnn/model.pt for gru, | |
| # <data-root>/spatiotemporal_gnn/<type>/model.pt otherwise, the | |
| # exact same real layout train_spatiotemporal_gnn.py already saves | |
| # to -- and adds whichever real checkpoints actually exist there, | |
| # named after their own temporal_type. Skips a temporal_type already | |
| # covered by the primary model or an explicit --extra-model above | |
| # (by name), so nothing real gets loaded twice. | |
| if not args.no_auto_extra_models: | |
| already_named = {args.temporal_type} | {name for name, _, _ in extra_model_specs} | |
| for candidate_type in ("gru", "lstm", "cnn"): | |
| if candidate_type == args.temporal_type or candidate_type in already_named: | |
| continue | |
| candidate_dir = args.data_root / "spatiotemporal_gnn" | |
| if candidate_type != "gru": | |
| candidate_dir = candidate_dir / candidate_type | |
| candidate_path = candidate_dir / "model.pt" | |
| if candidate_path.exists(): | |
| extra_model_specs.append((candidate_type, candidate_type, candidate_path)) | |
| already_named.add(candidate_type) | |
| auto_found = [name for name in ("gru", "lstm", "cnn") | |
| if name != args.temporal_type and any(n == name for n, _, _ in extra_model_specs)] | |
| if auto_found: | |
| print(f"Auto-detected real extra checkpoint(s) alongside the primary {args.temporal_type} " | |
| f"model (see --no-auto-extra-models to disable): {auto_found}") | |
| events_df = pd.read_csv(events_path) | |
| if events_df.empty: | |
| print(f"{events_path} has no real events -- run identify_flood_events.py first.") | |
| return | |
| event_lookup = build_event_lookup(events_df) | |
| station_thresholds = build_station_thresholds(events_df) | |
| print(f"Loaded {len(events_df)} real event(s) across {events_df['station_code'].nunique()} station(s)") | |
| # Rebuilding the exact same real data pipeline the model was | |
| # trained on -- same subgraph, same node ordering, same | |
| # standardization statistics. The saved weights are only meaningful | |
| # against these exact conditions. | |
| max_nodes_per_basin = args.max_nodes // len(BASIN_FILE_NAMES) | |
| basin_data = [] | |
| for basin_id, file_key in BASIN_FILE_NAMES.items(): | |
| nodes_df, edges_df = load_basin_graph(args.data_root, file_key) | |
| nodes_df, edges_df = build_collapsed_subgraph(nodes_df, edges_df, max_nodes=max_nodes_per_basin) | |
| basin_data.append((nodes_df, edges_df, basin_id)) | |
| # Must match train_spatiotemporal_gnn.py's own shared_static_cols/ | |
| # gauge_static_cols exactly -- the saved checkpoint's input layer | |
| # shape is fixed to whatever these were at training time. A mismatch | |
| # here isn't cosmetic; it's the same real shape-mismatch bug class | |
| # already caught twice before in this file (estimated_discharge, | |
| # dt_tensors). | |
| shared_static_cols = ["elevation_m", "catchment_area_km2", "idpr_value", | |
| "avg_groundwater_level_m", "distance_to_nearest_cavity_km", "n_cavities_within_20km"] | |
| gauge_static_cols: List[str] = [] | |
| combined = combine_basins(basin_data, shared_static_cols, gauge_static_cols) | |
| station_codes = combined["station_codes"] | |
| dynamic_tensors, target_tensor, waterlevel_target_tensor, dates = build_combined_dynamic_tensors( | |
| combined, basin_data, args.data_root, ("2013-01-01", args.test_end) | |
| ) | |
| # Real precipitation data-coverage check -- added after | |
| # --calibration-precip-staleness-days found 94% of the "elevated | |
| # precipitation" low-flow bucket was a stale forward-filled reading | |
| # (>30 real days old), a surprisingly high real fraction worth a | |
| # cheap sanity check before concluding the fix is a precipitation- | |
| # specific fill-decay policy: if the real Météo-France precipitation | |
| # feed simply stops before the real test period starts, EVERY test- | |
| # period example is structurally stale regardless of true rainfall | |
| # or any individual station's reporting pattern -- that's a real | |
| # data-freshness/pipeline gap ("get more recent precipitation | |
| # data"), a completely different fix than per-node reporting | |
| # sparsity ("decay the fill toward zero/climatology"). Checked here, | |
| # raw, right after the real precipitation tensor is built and before | |
| # standardization -- computed once, printed once, doesn't gate | |
| # anything else in this script. | |
| # Set below whenever some suffix of the real test period is | |
| # structurally guaranteed stale (every real station, not just some) | |
| # -- None means no such suffix exists (full real coverage, or no | |
| # real precipitation at all). Used by --exclude-stale-precip-tail | |
| # below to drop that suffix from every diagnostic in this script, | |
| # not just the precipitation-specific ones, since a structurally | |
| # stale precipitation input silently degrades discharge predictions | |
| # too, not only the precip-stratified checks that happened to | |
| # surface it. | |
| structurally_dead_precip_from = None | |
| if "precipitation" in dynamic_tensors: | |
| precip_is_real_anywhere = ~np.all(np.isnan(dynamic_tensors["precipitation"]), axis=0) | |
| dates_arr = np.array(dates) | |
| real_precip_dates = dates_arr[precip_is_real_anywhere] | |
| if len(real_precip_dates) == 0: | |
| print("Real precipitation data coverage: NO real (non-NaN) precipitation reading exists " | |
| "anywhere in this run's date range -- the 'precipitation' channel is present but " | |
| "entirely unobserved, so every real forward-filled/sentinel value downstream reflects " | |
| "that, not a per-node gap.") | |
| else: | |
| min_real_precip_date = pd.Timestamp(real_precip_dates.min()) | |
| max_real_precip_date = pd.Timestamp(real_precip_dates.max()) | |
| val_end_ts = pd.Timestamp(args.val_end) | |
| test_end_ts = pd.Timestamp(args.test_end) | |
| print(f"Real precipitation data coverage: {min_real_precip_date.date()} to " | |
| f"{max_real_precip_date.date()} (at least one real station reading on that date) -- " | |
| f"the real test period runs from {val_end_ts.date()} (exclusive) to {test_end_ts.date()}.") | |
| if max_real_precip_date < val_end_ts: | |
| structurally_dead_precip_from = val_end_ts | |
| print(f" WARNING: real precipitation coverage ENDS ({max_real_precip_date.date()}) before " | |
| f"the real test period even starts ({val_end_ts.date()}) -- EVERY real precipitation " | |
| f"reading anywhere in the test period is necessarily a stale forward-filled value " | |
| f"from before the test period began, structurally, regardless of true local " | |
| f"rainfall or any individual gauge's own reporting pattern. Any '--calibration-" | |
| f"precip-staleness-days' staleness finding above reflects this real coverage gap, " | |
| f"not per-node sparsity -- the real fix is fresher real precipitation data for the " | |
| f"test period, not a fill-decay policy change.") | |
| elif max_real_precip_date < test_end_ts: | |
| days_short = (test_end_ts - max_real_precip_date).days | |
| # Everything anchored at or after this date has a real | |
| # input precip window guaranteed to be more than | |
| # --calibration-precip-staleness-days stale for EVERY | |
| # real station, not just ones with their own reporting | |
| # gap -- that's the real point at which the structural | |
| # cutoff, not per-node sparsity, takes over completely. | |
| structurally_dead_precip_from = max_real_precip_date + pd.Timedelta( | |
| days=args.calibration_precip_staleness_days) | |
| print(f" Real precipitation coverage stops {days_short} real day(s) before the real test " | |
| f"period ends -- the later part of the real test period is stale for this same " | |
| f"structural reason; only the earlier part can meaningfully reflect real per-node/" | |
| f"per-gauge reporting gaps. Real examples anchored at or after " | |
| f"{structurally_dead_precip_from.date()} ({args.calibration_precip_staleness_days:.0f} " | |
| f"real days past the real coverage cutoff) are structurally stale for EVERY real " | |
| f"station -- see --exclude-stale-precip-tail to drop them from every diagnostic " | |
| f"below.") | |
| else: | |
| print(" Real precipitation coverage extends through the real test period -- the " | |
| "staleness found by --calibration-precip-staleness-days reflects real, per-node/" | |
| "per-gauge reporting gaps, not a systemic real data-pipeline cutoff.") | |
| # Per-real-gauge-station breakdown -- the global figure | |
| # above only says at least one node somewhere had a real | |
| # reading on a given date; it can't tell a real, uniform | |
| # pipeline cutoff (every station stops the same day) apart | |
| # from one or more real gauges having their OWN long- | |
| # standing reporting gap while others stay current. That | |
| # distinction matters directly: a shared cutoff is a real | |
| # data-freshness problem (get fresher data); a station- | |
| # specific gap is a real per-gauge ingestion problem (fix | |
| # that station's own feed), and only the second kind is | |
| # what the fill-decay-toward-zero fix would actually help | |
| # with. Restricted to real gauge stations | |
| # (combined["is_gauged"]) -- the ~4,500-node graph is mostly | |
| # virtual/confluence nodes that never carry a real | |
| # precipitation reading at all, so auditing every node | |
| # would just print noise. | |
| gauge_mask_cov = np.asarray(combined["is_gauged"], dtype=bool) | |
| is_real_precip_per_node = ~np.isnan(dynamic_tensors["precipitation"]) | |
| station_gaps = [] | |
| for node_idx_cov in np.where(gauge_mask_cov)[0]: | |
| code_cov = station_codes[node_idx_cov] | |
| real_dates_this_node = dates_arr[is_real_precip_per_node[node_idx_cov]] | |
| if len(real_dates_this_node) == 0: | |
| station_gaps.append((code_cov, None, float("inf"))) | |
| else: | |
| last_real_this_node = pd.Timestamp(real_dates_this_node.max()) | |
| station_gaps.append((code_cov, last_real_this_node, | |
| (max_real_precip_date - last_real_this_node).days)) | |
| station_gaps.sort(key=lambda row: row[2], reverse=True) | |
| n_never = sum(1 for _, last, _ in station_gaps if last is None) | |
| n_matches_global_max = sum(1 for _, _, gap in station_gaps if gap == 0) | |
| print(f"\nReal per-gauge-station precipitation coverage ({len(station_gaps)} real gauge " | |
| f"station(s)), most-stale-relative-to-the-global-max first:") | |
| for code_cov, last_real_this_node, gap_days in station_gaps[:10]: | |
| if last_real_this_node is None: | |
| print(f" {code_cov}: NO real precipitation reading anywhere in this run's date range.") | |
| else: | |
| print(f" {code_cov}: last real reading {last_real_this_node.date()} " | |
| f"({gap_days} real day(s) behind this run's global max, " | |
| f"{max_real_precip_date.date()})") | |
| print(f" -> {n_matches_global_max}/{len(station_gaps)} real gauge station(s) share the exact " | |
| f"same most-recent real reading as the global max -- consistent with a single shared " | |
| f"pipeline feed cutting off together for all of them, not independent per-station gaps. " | |
| f"{n_never} real gauge station(s) have NO real precipitation reading anywhere in this " | |
| f"run's date range at all -- a real, standing per-station ingestion gap, distinct from " | |
| f"the shared cutoff above.") | |
| # Mirrors train_spatiotemporal_gnn.py's main() exactly -- confirmed | |
| # as a real, necessary fix, not an optional refinement: the saved | |
| # checkpoint was trained with this channel present (added to | |
| # main() several turns ago), and without replicating it here, the | |
| # reconstructed model has 2 fewer input channels than the real | |
| # checkpoint, causing a genuine shape-mismatch error on load. Same | |
| # real reasoning as train_spatiotemporal_gnn.py's own version -- | |
| # see that file for the full explanation. | |
| catchment_area_km2 = extract_catchment_area_km2(basin_data, combined) | |
| n_real_catchment = int((~np.isnan(catchment_area_km2)).sum()) | |
| if "precipitation" in dynamic_tensors and n_real_catchment > 0: | |
| try: | |
| rolling_precip_sum = compute_rolling_precip_sum(dynamic_tensors["precipitation"], window_days=args.lookback_days) | |
| specific_discharge_k = fit_specific_discharge_coefficient( | |
| target_tensor, rolling_precip_sum, catchment_area_km2, dates, args.train_end, | |
| ) | |
| estimated_discharge = estimate_discharge_from_precip(rolling_precip_sum, catchment_area_km2, specific_discharge_k) | |
| dynamic_tensors["estimated_discharge"] = estimated_discharge | |
| except ValueError: | |
| pass # matches train_spatiotemporal_gnn.py's own graceful skip | |
| dynamic_tensors, target_tensor, waterlevel_target_tensor, standardization_stats = standardize_dynamic_tensors( | |
| dynamic_tensors, target_tensor, waterlevel_target_tensor, dates, args.train_end, | |
| optional_vars=["precipitation", "estimated_discharge"], | |
| ) | |
| discharge_mean, discharge_std = standardization_stats["discharge"] | |
| # Mirrors train_spatiotemporal_gnn.py's own dt_tensors construction | |
| # exactly -- the real checkpoint was trained with this channel | |
| # present, so reconstructing the model without it here would cause | |
| # the same real shape-mismatch error found once already for | |
| # estimated_discharge. | |
| dt_tensors = { | |
| var: compress_days_since_last_real(compute_days_since_last_real(tensor)) | |
| for var, tensor in dynamic_tensors.items() | |
| } | |
| # For the real precipitation/false-positive correlation check | |
| # (--check-fp-precipitation): the real channel index and real | |
| # un-standardization stats, captured here since var_names' | |
| # insertion order (and channels_per_var=3, matching the real dt | |
| # channel this project added) has to match prepare_graph_training_ | |
| # windows' own real channel layout exactly, or this would silently | |
| # read the wrong channel entirely. | |
| var_names_for_channels = list(dynamic_tensors.keys()) | |
| has_real_precip = "precipitation" in var_names_for_channels | |
| if has_real_precip: | |
| precip_value_idx, _ = get_dynamic_channel_index(var_names_for_channels, "precipitation", channels_per_var=3) | |
| precip_mean, precip_std = standardization_stats["precipitation"] | |
| # [value, missing_flag, days_since] per var (channels_per_var=3 | |
| # above) -- days_since is the 3rd of that var's own 3 channels, | |
| # same offset --trace-fp-inputs already uses for discharge/ | |
| # precip's own dt channels. Used by --calibration-precip- | |
| # staleness-days to tell a genuinely recent rain pulse apart | |
| # from a stale forward-filled reading being counted as one. | |
| precip_dt_idx = precip_value_idx + 2 | |
| # Real, per-node fill values -- MUST match training's own fix | |
| # exactly (see prepare_graph_training_windows' own docstring and | |
| # train_spatiotemporal_gnn.py's identical computation), or the | |
| # model would be evaluated on inputs built differently than the | |
| # ones it was actually trained on -- a real, silent train/eval | |
| # mismatch, not just a missed improvement. MEDIAN, not mean -- see | |
| # compute_per_node_historical_median's own docstring for why. | |
| train_end_ts_for_fill = pd.Timestamp(args.train_end) | |
| train_date_mask_for_fill = np.array([d <= train_end_ts_for_fill for d in dates]) | |
| per_node_fill_values = { | |
| var: compute_per_node_historical_median(tensor[:, train_date_mask_for_fill]) | |
| for var, tensor in dynamic_tensors.items() | |
| } | |
| # Same forward-fill fix as training, and for the same reason this | |
| # whole block exists: MUST match train_spatiotemporal_gnn.py's | |
| # identical computation exactly, or the model would be evaluated on | |
| # inputs built differently than the ones it was actually trained | |
| # on. Preferred over per_node_fill_values above whenever a real | |
| # prior observation exists to carry forward -- see | |
| # compute_forward_filled_tensor's own docstring and | |
| # prepare_graph_training_windows' "CRITICAL, CONFIRMED FIX #2" for | |
| # why. No train-only-fit discipline needed (unlike | |
| # per_node_fill_values): forward-fill only ever looks backward. | |
| forward_fill_tensors = { | |
| var: compute_forward_filled_tensor(tensor) | |
| for var, tensor in dynamic_tensors.items() | |
| } | |
| X, Y, Y_mask, anchor_dates = prepare_graph_training_windows( | |
| combined["x_shared_static"], dynamic_tensors, target_tensor, dates, args.lookback_days, HORIZONS, | |
| dt_tensors=dt_tensors, per_node_fill_values=per_node_fill_values, | |
| forward_fill_tensors=forward_fill_tensors, | |
| ) | |
| # Held-out test split only -- evaluating against data the model saw | |
| # during training would overstate real skill, same principle as | |
| # every other evaluation in this project. | |
| anchor_dates_arr = pd.DatetimeIndex(anchor_dates) | |
| test_mask = anchor_dates_arr > pd.Timestamp(args.val_end) | |
| if args.exclude_stale_precip_tail and structurally_dead_precip_from is not None: | |
| # Applied at the SAME point as the val_end split above, not as | |
| # a later per-diagnostic filter -- so every downstream stat in | |
| # this script (calibration, the confusion matrix, NSE/RMSE/ | |
| # PBIAS) is computed on the same, structurally-clean example | |
| # set, rather than some diagnostics quietly including the dead | |
| # tail and others not. | |
| n_before_tail_exclusion = int(test_mask.sum()) | |
| test_mask = test_mask & (anchor_dates_arr < structurally_dead_precip_from) | |
| n_excluded_tail = n_before_tail_exclusion - int(test_mask.sum()) | |
| print(f"--exclude-stale-precip-tail: dropped {n_excluded_tail} real test example(s) anchored at " | |
| f"or after {structurally_dead_precip_from.date()} (real precipitation input structurally " | |
| f"stale for EVERY real station past this point) from every diagnostic below.") | |
| X_test, Y_test, anchor_dates_test = X[test_mask], Y[test_mask], [d for d, keep in zip(anchor_dates, test_mask) if keep] | |
| print(f"Evaluating on {len(X_test)} real held-out test example(s)") | |
| # Real forecasted precipitation for the test examples -- mirrors | |
| # train_spatiotemporal_gnn.py's own wiring exactly, so this | |
| # evaluation reflects the same inputs the model was actually | |
| # trained with, not a stripped-down version of them. | |
| forecast_csv_path = args.data_root / "previous_runs_forecast.csv" | |
| if forecast_csv_path.exists(): | |
| nodes_df_combined = pd.DataFrame({"station_code": station_codes}) | |
| forecast_precip_test, forecast_missing_test = build_forecast_tensor( | |
| nodes_df_combined, forecast_csv_path, anchor_dates_test, HORIZONS, | |
| ) | |
| n_real_forecast = int((forecast_missing_test == 0.0).sum()) | |
| print(f"Real forecast precipitation available for {n_real_forecast} test (example, node, horizon) value(s)") | |
| else: | |
| forecast_precip_test = np.zeros((len(X_test), combined["n_nodes"], len(HORIZONS)), dtype=np.float32) | |
| forecast_missing_test = np.ones((len(X_test), combined["n_nodes"], len(HORIZONS)), dtype=np.float32) | |
| print("No real forecast archive found -- evaluating without forecast precipitation") | |
| x_shared_static_std, _, shared_static_flagged = standardize_static_features( | |
| combined["x_shared_static"], shared_static_cols, | |
| ) | |
| x_gauge_static_std, _, gauge_static_flagged = standardize_static_features( | |
| combined["x_gauge_static"], gauge_static_cols, | |
| ) | |
| print(f"Static feature standardization: shared missingness flags for {shared_static_flagged or 'none'}, " | |
| f"gauge missingness flags for {gauge_static_flagged or 'none'}") | |
| x_shared_static_t = torch.tensor(x_shared_static_std, dtype=torch.float32) | |
| x_gauge_static_t = torch.tensor(x_gauge_static_std, dtype=torch.float32) | |
| is_gauged_t = torch.tensor(combined["is_gauged"], dtype=torch.bool) | |
| edge_index_t = torch.tensor(combined["edge_index"], dtype=torch.long) | |
| edge_attr_t = torch.tensor(combined["edge_attr"], dtype=torch.float32) | |
| basin_id_t = torch.tensor(combined["basin_id"], dtype=torch.long) | |
| horizons_t = torch.tensor(HORIZONS, dtype=torch.float32) | |
| model = SpatiotemporalGNN( | |
| shared_static_dim=x_shared_static_t.shape[1], gauge_static_dim=x_gauge_static_t.shape[1], | |
| gauge_dynamic_dim=X.shape[-1], edge_dim=edge_attr_t.shape[1], | |
| n_basins=len(BASIN_FILE_NAMES), n_targets=2, quantiles=QUANTILES, | |
| temporal_type=args.temporal_type, | |
| ) | |
| model.load_state_dict(torch.load(model_path, map_location="cpu")) | |
| model.eval() | |
| print(f"Loaded trained model from {model_path}") | |
| if args.flag_quantile not in QUANTILES: | |
| print(f"--flag-quantile {args.flag_quantile} isn't one of the model's real predicted quantiles " | |
| f"{QUANTILES} -- pick one of those.") | |
| return | |
| default_quantile_idx = QUANTILES.index(args.flag_quantile) | |
| # Per-horizon quantile_idx map -- every horizon defaults to | |
| # --flag-quantile's index, then real per-horizon overrides replace | |
| # specific entries. Built as an explicit {horizon: idx} dict rather | |
| # than a single shared value, since real testing in this project | |
| # showed the right operating point genuinely differs by horizon | |
| # (see --flag-quantile-overrides' own help text for the real numbers). | |
| quantile_idx_by_horizon = {h: default_quantile_idx for h in HORIZONS} | |
| if args.flag_quantile_overrides: | |
| for pair in args.flag_quantile_overrides.split(","): | |
| pair = pair.strip() | |
| if not pair: | |
| continue | |
| try: | |
| h_str, q_str = pair.split("=") | |
| h_override, q_override = int(h_str), float(q_str) | |
| except ValueError: | |
| print(f"Couldn't parse override '{pair}' -- expected format 'horizon=quantile', e.g. '5=0.9'.") | |
| return | |
| if h_override not in HORIZONS: | |
| print(f"Override horizon {h_override} isn't one of this model's real horizons {HORIZONS}.") | |
| return | |
| if q_override not in QUANTILES: | |
| print(f"Override quantile {q_override} (for horizon {h_override}) isn't one of the " | |
| f"model's real predicted quantiles {QUANTILES} -- pick one of those.") | |
| return | |
| quantile_idx_by_horizon[h_override] = QUANTILES.index(q_override) | |
| if args.blend_quantiles: | |
| print(f"Flagging a real event when a real median/q0.95 BLEND (see --blend-quantiles) exceeds " | |
| f"that station's real threshold, instead of the raw per-horizon quantile below -- the " | |
| f"per-horizon quantile map is kept only as --flag-quantile-overrides' own reference point, " | |
| f"not what's actually compared against the threshold now. Per-horizon quantile (unused " | |
| f"while --blend-quantiles is set): " | |
| f"{ {h: QUANTILES[idx] for h, idx in quantile_idx_by_horizon.items()} }") | |
| else: | |
| print(f"Flagging a real event when the model's predicted quantile exceeds that station's real " | |
| f"threshold -- not the median, since flood risk is a tail question. Per-horizon quantile: " | |
| f"{ {h: QUANTILES[idx] for h, idx in quantile_idx_by_horizon.items()} }") | |
| # Confusion matrix per horizon: {horizon: {"tp":.., "fp":.., "fn":.., "tn":..}} | |
| confusion = {h: {"tp": 0, "fp": 0, "fn": 0, "tn": 0} for h in HORIZONS} | |
| fp_precip_values: List[float] = [] | |
| tn_precip_values: List[float] = [] | |
| # Real (0.95 - median) discharge spread, per horizon, across every | |
| # real example and node regardless of tp/fp/fn/tn classification -- | |
| # unlike the precipitation check, this doesn't depend on the | |
| # flagging outcome, so every real prediction contributes. | |
| spread_by_horizon: Dict[int, List[float]] = {h: [] for h in HORIZONS} | |
| if args.check_quantile_calibration: | |
| # Real, per-station training-period median -- used only to | |
| # restrict the precipitation-stratified half of this check to | |
| # real LOW-flow days, the same restriction --trace-fp-inputs | |
| # already uses (a station's own real typical day, not the | |
| # single global naive baseline used for the overall skill | |
| # comparison elsewhere). | |
| train_end_ts_calib = pd.Timestamp(args.train_end) | |
| train_date_mask_calib = np.array([d <= train_end_ts_calib for d in dates]) | |
| station_train_median: Dict[str, float] = {} | |
| for node_idx_calib, code_calib in enumerate(station_codes): | |
| real_train_vals = (target_tensor[node_idx_calib, train_date_mask_calib] * discharge_std | |
| + discharge_mean) | |
| real_train_vals = real_train_vals[~np.isnan(real_train_vals)] | |
| station_train_median[code_calib] = float(np.median(real_train_vals)) if len(real_train_vals) else float("nan") | |
| # Global per-quantile calibration: does the real empirical | |
| # coverage (fraction of real observed values at or below the | |
| # predicted quantile) match the nominal quantile level, across | |
| # every real (gauge, date, horizon) combination in the held-out | |
| # test set -- not just the handful of examples --trace-fp-inputs | |
| # happens to surface. | |
| calib_hits: Dict[float, int] = {q: 0 for q in QUANTILES} | |
| calib_totals: Dict[float, int] = {q: 0 for q in QUANTILES} | |
| # Real, DIRECT low-flow-vs-high-flow split, for EVERY real | |
| # predicted quantile -- added after a real run showed a real, | |
| # derived asymmetry at q=0.95 (global coverage under nominal, | |
| # but a precipitation-stratified, low-flow-only subset over | |
| # nominal): the global figure and the low-flow-only figure | |
| # together implied a real high-flow-day coverage far below | |
| # nominal, but that was only ever back-calculated by subtraction. | |
| # Generalized to every quantile (not just 0.95) after a real | |
| # station plot (H404021101, full test range) visually showed | |
| # q0.99 detaching from real observed discharge during quiet | |
| # periods far more than q0.95/q0.9 did -- this checks whether | |
| # that's real and consistent, or just how it looked on one plot. | |
| gap_low_flow: Dict[float, List[float]] = {q: [] for q in QUANTILES} | |
| gap_high_flow: Dict[float, List[float]] = {q: [] for q in QUANTILES} | |
| coverage_low_flow: Dict[float, Dict[str, int]] = {q: {"hits": 0, "total": 0} for q in QUANTILES} | |
| coverage_high_flow: Dict[float, Dict[str, int]] = {q: {"hits": 0, "total": 0} for q in QUANTILES} | |
| # Precipitation-stratified, low-flow-only, q=0.95-specific: does | |
| # a recent real rain pulse specifically drive the over- | |
| # prediction gap, or is it evenly spread regardless of real | |
| # recent precipitation. | |
| q95_gap_high_precip: List[float] = [] | |
| q95_gap_low_precip: List[float] = [] | |
| q95_coverage_high_precip = {"hits": 0, "total": 0} | |
| q95_coverage_low_precip = {"hits": 0, "total": 0} | |
| # Staleness split of the elevated-precipitation bucket above -- | |
| # is "elevated precipitation" real recent rain, or a stale | |
| # forward-filled reading from long before the real gap started | |
| # (see --calibration-precip-staleness-days' own docstring for | |
| # the real two examples that motivated this)? Only meaningful | |
| # inside the elevated-precip side, so "low_precip" has no | |
| # staleness counterpart -- a quiet reading that's ALSO stale | |
| # isn't a distinct case worth separately tracking here. | |
| q95_gap_high_precip_stale: List[float] = [] | |
| q95_gap_high_precip_recent: List[float] = [] | |
| q95_coverage_high_precip_stale = {"hits": 0, "total": 0} | |
| q95_coverage_high_precip_recent = {"hits": 0, "total": 0} | |
| if 0.95 not in QUANTILES: | |
| print("--check-quantile-calibration's precipitation-stratified half needs 0.95 in this " | |
| f"model's real predicted quantiles {QUANTILES} -- skipping that half, global " | |
| f"calibration coverage still runs.") | |
| if args.check_quantile_spread: | |
| if 0.5 not in QUANTILES or 0.95 not in QUANTILES: | |
| print("--check-quantile-spread needs both 0.5 and 0.95 in this model's real predicted " | |
| f"quantiles {QUANTILES} -- neither is missing normally, but checked defensively.") | |
| args.check_quantile_spread = False | |
| else: | |
| median_idx, q95_idx = QUANTILES.index(0.5), QUANTILES.index(0.95) | |
| requested_stations: List[str] = [] | |
| station_timeseries_data: Dict[str, Dict[str, list]] = {} | |
| if args.plot_stations: | |
| if 0.5 not in QUANTILES or 0.95 not in QUANTILES: | |
| print("--plot-stations needs both 0.5 and 0.95 in this model's real predicted quantiles " | |
| f"{QUANTILES} -- neither is missing normally, but checked defensively.") | |
| elif args.plot_stations_horizon not in HORIZONS: | |
| print(f"--plot-stations-horizon {args.plot_stations_horizon} isn't one of this model's " | |
| f"real horizons {HORIZONS}.") | |
| else: | |
| requested_stations = [s.strip() for s in args.plot_stations.split(",") if s.strip()] | |
| unknown = [s for s in requested_stations if s not in station_codes] | |
| if unknown: | |
| print(f"Station(s) {unknown} aren't real station codes in this run -- dropping them.") | |
| requested_stations = [s for s in requested_stations if s in station_codes] | |
| median_idx_ts, q95_idx_ts = QUANTILES.index(0.5), QUANTILES.index(0.95) | |
| plot_horizon_idx = HORIZONS.index(args.plot_stations_horizon) | |
| # "median"/"q95" kept as their own keys (used elsewhere by | |
| # the hydrology metrics and --trace-fp-inputs, both of which | |
| # only ever need those two specifically); "quantiles" added | |
| # alongside them to carry EVERY real predicted quantile, for | |
| # plot_station_timeseries to draw the full fan rather than | |
| # just the two hardcoded lines it used to. | |
| station_timeseries_data = { | |
| code: {"dates": [], "observed": [], "median": [], "q95": [], | |
| "quantiles": {q: [] for q in QUANTILES}, "example_idx": []} | |
| for code in requested_stations | |
| } | |
| # {horizon: {"fp_with_forecast": n, "fp_without_forecast": n, | |
| # "all_with_forecast": n, "all_without_forecast": n}} -- | |
| # only meaningful for horizons in FORECAST_LEAD_TIMES, which is | |
| # itself a real, direct subset of HORIZONS this project actually | |
| # has real forecast coverage for. | |
| forecast_fp_stats: Dict[int, Dict[str, int]] = { | |
| h: {"fp_with_forecast": 0, "fp_without_forecast": 0, "all_with_forecast": 0, "all_without_forecast": 0} | |
| for h in HORIZONS if h in FORECAST_LEAD_TIMES | |
| } | |
| fp_details: List[dict] = [] | |
| if args.dump_fp_details is not None and args.dump_fp_details not in HORIZONS: | |
| print(f"--dump-fp-details {args.dump_fp_details} isn't one of this model's real horizons {HORIZONS}.") | |
| return | |
| # --blend-quantiles, applied to flood-FLAGGING (not just the | |
| # --plot-stations comparison plot above, which computes its own, | |
| # independent blend already): flagging needs a real per-station tau | |
| # (the spread at which alpha=0.5) BEFORE the main real per-example | |
| # loop below can decide predicted_flag for even its first real | |
| # example, so this real tau can't be accumulated inside that same | |
| # loop as it runs -- it has to be known upfront. Rather than fabricate | |
| # a fixed/global tau (which would blend very differently for a small | |
| # station than a large one), this real preliminary pass runs the | |
| # model over the exact same real X_test once, purely to measure each | |
| # real station's own median (q0.95-median) spread over its real test | |
| # period -- the identical real tau definition --plot-stations' own | |
| # blend uses, just computed here for every real thresholded station | |
| # instead of only the ones passed to --plot-stations. This is a | |
| # second, real forward pass over the model (not free), but X_test is | |
| # small enough here that the added real cost is minor next to | |
| # everything else this script already does in one run. | |
| blend_tau_by_station: Dict[str, float] = {} | |
| train_median_by_station: Dict[str, float] = {} | |
| train_q90_by_station: Dict[str, float] = {} | |
| if args.blend_quantiles: | |
| if 0.5 not in QUANTILES or args.blend_upper_quantile not in QUANTILES: | |
| print(f"--blend-quantiles needs both 0.5 and --blend-upper-quantile {args.blend_upper_quantile} " | |
| f"in this model's real predicted quantiles {QUANTILES} -- disabling the flagging blend " | |
| "(falls back to --flag-quantile's raw quantile for flagging; --plot-stations' own blend " | |
| "is gated the same way and already disabled itself above if this is missing).") | |
| args.blend_quantiles = False | |
| else: | |
| median_idx_flag, q95_idx_flag = QUANTILES.index(0.5), QUANTILES.index(args.blend_upper_quantile) | |
| thresholded_codes = [c for c in station_codes if c in station_thresholds] | |
| if args.blend_mode == "level": | |
| # No model forward pass needed at all for this mode -- | |
| # train_median/train_q90 come straight from each real | |
| # station's own real training-period discharge (target_ | |
| # tensor), the same real source station_train_median | |
| # (above, under --check-quantile-calibration) and the | |
| # plot loop's own train_q90 already use elsewhere in | |
| # this script. | |
| train_end_ts_blend = pd.Timestamp(args.train_end) | |
| train_date_mask_blend = np.array([d <= train_end_ts_blend for d in dates]) | |
| for node_idx_blend, code_blend in enumerate(station_codes): | |
| real_train_vals_blend = (target_tensor[node_idx_blend, train_date_mask_blend] | |
| * discharge_std + discharge_mean) | |
| real_train_vals_blend = real_train_vals_blend[~np.isnan(real_train_vals_blend)] | |
| if len(real_train_vals_blend): | |
| train_median_by_station[code_blend] = float(np.median(real_train_vals_blend)) | |
| train_q90_by_station[code_blend] = float(np.quantile(real_train_vals_blend, 0.90)) | |
| else: | |
| train_median_by_station[code_blend] = float("nan") | |
| train_q90_by_station[code_blend] = float("nan") | |
| print("--blend-quantiles (flagging, mode=level): real per-station (train_median, " | |
| "train_q90) L/s anchors -- " | |
| + ", ".join(f"{c}=({train_median_by_station[c]:.1f},{train_q90_by_station[c]:.1f})" | |
| for c in thresholded_codes)) | |
| else: | |
| # 'spread' mode: flagging needs a real per-station tau | |
| # (the spread at which alpha=0.5) BEFORE the main real | |
| # per-example loop below can decide predicted_flag for | |
| # even its first real example, so this real tau can't be | |
| # accumulated inside that same loop as it runs -- it has | |
| # to be known upfront. Rather than fabricate a fixed/ | |
| # global tau (which would blend very differently for a | |
| # small station than a large one), this real preliminary | |
| # pass runs the model over the exact same real X_test | |
| # once, purely to measure each real station's own median | |
| # (q0.95-median) spread over its real test period -- the | |
| # identical real tau definition --plot-stations' own | |
| # blend uses, just computed here for every real | |
| # thresholded station instead of only the ones passed to | |
| # --plot-stations. This is a second, real forward pass | |
| # over the model (not free), but X_test is small enough | |
| # here that the added real cost is minor next to | |
| # everything else this script already does in one run. | |
| # | |
| # Restricts the real tau-fitting data itself to the same | |
| # real horizons the blend will actually be USED at (see | |
| # --blend-quantiles-max-horizon) -- e.g. with max- | |
| # horizon=5, a real station's tau is fit only from its | |
| # own real h in {1,2,3,4,5} spread, not diluted by real | |
| # long-horizon spread that'll never drive a real | |
| # flagging decision anyway. | |
| blend_horizon_idx = [ | |
| h_idx for h_idx, h in enumerate(HORIZONS) | |
| if args.blend_quantiles_max_horizon is None or h <= args.blend_quantiles_max_horizon | |
| ] | |
| spread_by_station_flag: Dict[str, List[float]] = {code: [] for code in station_codes} | |
| with torch.no_grad(): | |
| for i_tau in range(len(X_test)): | |
| x_dynamic_seq_tau = torch.tensor(X_test[i_tau], dtype=torch.float32) | |
| forecast_precip_tau = torch.tensor(forecast_precip_test[i_tau], dtype=torch.float32) | |
| forecast_missing_tau = torch.tensor(forecast_missing_test[i_tau], dtype=torch.float32) | |
| pred_tau = model(x_shared_static_t, x_gauge_static_t, x_dynamic_seq_tau, is_gauged_t, | |
| edge_index_t, edge_attr_t, basin_id_t, horizons_t, | |
| forecast_precip_tau, forecast_missing_tau) | |
| Q_std_tau = pred_tau[:, :, 0, :].numpy() | |
| Q_real_tau = Q_std_tau * discharge_std + discharge_mean | |
| spread_tau = np.clip( | |
| Q_real_tau[:, :, q95_idx_flag] - Q_real_tau[:, :, median_idx_flag], 0.0, None) | |
| spread_tau = spread_tau[:, blend_horizon_idx] | |
| for node_idx_tau, code_tau in enumerate(station_codes): | |
| spread_by_station_flag[code_tau].extend(spread_tau[node_idx_tau, :].tolist()) | |
| for code_tau, spreads_tau in spread_by_station_flag.items(): | |
| if args.blend_scale is not None: | |
| blend_tau_by_station[code_tau] = args.blend_scale | |
| continue | |
| finite_tau = np.asarray([s for s in spreads_tau if np.isfinite(s)]) | |
| tau_val = float(np.median(finite_tau)) if len(finite_tau) else float("nan") | |
| blend_tau_by_station[code_tau] = tau_val if np.isfinite(tau_val) and tau_val > 0 else 1.0 | |
| print(f"--blend-quantiles (flagging, mode=spread, upper=q{args.blend_upper_quantile}): real " | |
| f"per-station tau (L/s), this run's real test-period median predicted " | |
| f"(q{args.blend_upper_quantile}-median) spread -- " | |
| + ", ".join(f"{c}={blend_tau_by_station[c]:.1f}" for c in thresholded_codes)) | |
| if args.blend_quantiles_max_horizon is not None: | |
| blended_h = [h for h in HORIZONS if h <= args.blend_quantiles_max_horizon] | |
| raw_h = [h for h in HORIZONS if h > args.blend_quantiles_max_horizon] | |
| print(f"--blend-quantiles-max-horizon {args.blend_quantiles_max_horizon}: flagging blends " | |
| f"real horizons {blended_h}, falls back to raw --flag-quantile for real horizons " | |
| f"{raw_h}.") | |
| # Precomputed once, outside the main per-example loop below, so the | |
| # main loop's 'level'-mode branch is a cheap array lookup rather than | |
| # rebuilding these every (example, horizon) iteration. | |
| train_median_arr_by_node = np.array( | |
| [train_median_by_station.get(c, float("nan")) for c in station_codes]) | |
| train_q90_arr_by_node = np.array( | |
| [train_q90_by_station.get(c, float("nan")) for c in station_codes]) | |
| with torch.no_grad(): | |
| for i in range(len(X_test)): | |
| x_dynamic_seq = torch.tensor(X_test[i], dtype=torch.float32) | |
| anchor = anchor_dates_test[i] | |
| forecast_precip_i = torch.tensor(forecast_precip_test[i], dtype=torch.float32) | |
| forecast_missing_i = torch.tensor(forecast_missing_test[i], dtype=torch.float32) | |
| pred = model(x_shared_static_t, x_gauge_static_t, x_dynamic_seq, is_gauged_t, | |
| edge_index_t, edge_attr_t, basin_id_t, horizons_t, | |
| forecast_precip_i, forecast_missing_i) | |
| # [n_nodes, n_horizons, n_quantiles] -- every predicted | |
| # quantile, not just one: which specific quantile counts as | |
| # "the flood flag" now genuinely varies by horizon (see | |
| # quantile_idx_by_horizon), so the selection has to happen | |
| # per-horizon below, not once for the whole tensor. | |
| Q_pred_all_quantiles_standardized = pred[:, :, 0, :].numpy() | |
| Q_pred_all_quantiles_real = Q_pred_all_quantiles_standardized * discharge_std + discharge_mean | |
| if requested_stations: | |
| # Real observed value at this specific (station, target | |
| # date) -- Y_test is standardized the same way discharge | |
| # predictions are, so the same real mean/std un- | |
| # standardizes it correctly. NaN where no real | |
| # observation exists, which plot_station_timeseries | |
| # renders as a genuine gap, not a fabricated value. | |
| target_date_ts = anchor + pd.Timedelta(days=args.plot_stations_horizon - 1) | |
| Y_real_this_horizon = Y_test[i][:, plot_horizon_idx] * discharge_std + discharge_mean | |
| for code in requested_stations: | |
| node_idx_ts = station_codes.index(code) | |
| station_timeseries_data[code]["dates"].append(target_date_ts) | |
| station_timeseries_data[code]["observed"].append(float(Y_real_this_horizon[node_idx_ts])) | |
| station_timeseries_data[code]["median"].append( | |
| float(Q_pred_all_quantiles_real[node_idx_ts, plot_horizon_idx, median_idx_ts])) | |
| station_timeseries_data[code]["q95"].append( | |
| float(Q_pred_all_quantiles_real[node_idx_ts, plot_horizon_idx, q95_idx_ts])) | |
| for q_idx_ts, q_ts in enumerate(QUANTILES): | |
| station_timeseries_data[code]["quantiles"][q_ts].append( | |
| float(Q_pred_all_quantiles_real[node_idx_ts, plot_horizon_idx, q_idx_ts])) | |
| # Real example index -- needed to trace back to | |
| # X_test[i]'s own real input window for this exact | |
| # example, if --trace-fp-inputs is used to inspect | |
| # what real input features drove a specific false | |
| # spike, rather than guessing from the plot alone. | |
| station_timeseries_data[code]["example_idx"].append(i) | |
| if args.check_quantile_spread: | |
| # Real, un-standardized spread -- computed once per | |
| # example here (all horizons, all nodes), not inside the | |
| # per-horizon flagging loop below, since this doesn't | |
| # depend on any station threshold or event label at all. | |
| spread_real = Q_pred_all_quantiles_real[:, :, q95_idx] - Q_pred_all_quantiles_real[:, :, median_idx] | |
| for h_idx, h in enumerate(HORIZONS): | |
| spread_by_horizon[h].extend(spread_real[:, h_idx].tolist()) | |
| if (args.check_fp_precipitation or args.check_quantile_calibration) and has_real_precip: | |
| # Real recent precipitation, per node, for THIS example's | |
| # own input window -- one value per node, shared across | |
| # every horizon below (the input window itself doesn't | |
| # change with horizon, only the target date does), same | |
| # ROUTING_HORIZON_COUNT-day real window water_balance_loss | |
| # itself uses during training, not an arbitrarily chosen one. | |
| recent_window = X_test[i][-ROUTING_HORIZON_COUNT:, :, precip_value_idx] # [days, n_nodes], standardized | |
| # STALE COMMENT, CORRECTED: this used to say missing days | |
| # were filled with 0.0 (standardized), true before this | |
| # project's forward-fill fix. Precipitation is now one of | |
| # forward_fill_tensors' vars (same dict comprehension over | |
| # ALL dynamic_tensors, not just discharge), so a missing | |
| # precip day here instead carries forward the real, most- | |
| # recent REAL reading -- which can be an arbitrarily old | |
| # one if this node's real precip coverage has a long gap, | |
| # not a neutral value. That matters directly for THIS sum: | |
| # a real recent rain pulse and a real year-old stale | |
| # reading both just look like "some nonzero mm" to this | |
| # sum alone -- see precip_dt_recent_real_per_node below, | |
| # added specifically to tell them apart. | |
| recent_precip_sum_real_per_node = recent_window.sum(axis=0) * precip_std + ROUTING_HORIZON_COUNT * precip_mean | |
| # Real days-since-last-real-precipitation-observation, at | |
| # the single most recent real lookback day (index -1) -- | |
| # the freshest staleness read available for this example, | |
| # same day compress_days_since_last_real's own forward- | |
| # accumulation would report as of "now". dt channels are | |
| # used directly as log1p-compressed real values (NOT | |
| # further z-scored -- confirmed directly: a real traced | |
| # example's compressed value landed exactly on | |
| # log1p(365)=5.9026..., compress_days_since_last_real's | |
| # own real sentinel ceiling, which a z-scored value | |
| # couldn't land on exactly), so this inverts with expm1 | |
| # to recover real days for the human-readable staleness | |
| # threshold below. | |
| precip_dt_recent_real_per_node = np.expm1(X_test[i][-1, :, precip_dt_idx]) | |
| for h_idx, h in enumerate(HORIZONS): | |
| target_date = anchor + pd.Timedelta(days=h - 1) | |
| quantile_idx = quantile_idx_by_horizon[h] | |
| Q_pred_real_this_horizon = Q_pred_all_quantiles_real[:, h_idx, quantile_idx] | |
| blend_active_this_horizon = ( | |
| args.blend_quantiles | |
| and (args.blend_quantiles_max_horizon is None or h <= args.blend_quantiles_max_horizon) | |
| ) | |
| if blend_active_this_horizon: | |
| # Real per-node blend, this horizon: predicted = | |
| # median + alpha*(q95-median) either mode -- only how | |
| # alpha itself is derived differs (see --blend-mode's | |
| # own help text). station_codes indexes these arrays | |
| # the same way it indexes every other real per-node | |
| # array in this loop (node_idx below), so | |
| # Q_flag_this_horizon[node_idx] lines up correctly. | |
| median_h = Q_pred_all_quantiles_real[:, h_idx, median_idx_flag] | |
| q95_h = Q_pred_all_quantiles_real[:, h_idx, q95_idx_flag] | |
| spread_h = np.clip(q95_h - median_h, 0.0, None) | |
| if args.blend_mode == "level": | |
| denom_h = train_q90_arr_by_node - train_median_arr_by_node | |
| denom_h_safe = np.where(denom_h > 0, denom_h, 1.0) | |
| alpha_h = np.clip((median_h - train_median_arr_by_node) / denom_h_safe, 0.0, 1.0) | |
| alpha_h = np.where(np.isfinite(alpha_h), alpha_h, 0.0) | |
| else: | |
| tau_arr = np.array([blend_tau_by_station.get(c, 1.0) for c in station_codes]) | |
| alpha_h = spread_h / (spread_h + tau_arr) | |
| Q_flag_this_horizon = median_h + alpha_h * spread_h | |
| else: | |
| Q_flag_this_horizon = Q_pred_real_this_horizon | |
| track_forecast_fp = args.check_forecast_fp_correlation and h in FORECAST_LEAD_TIMES | |
| if args.check_quantile_calibration: | |
| # Real, un-standardized observed value at every real | |
| # node for THIS horizon -- NaN preserved through the | |
| # arithmetic wherever no real observation exists, so | |
| # np.isnan below correctly restricts this to real | |
| # (gauge, date) pairs only, same as everywhere else in | |
| # this script that un-standardizes Y_test. | |
| Y_real_this_horizon_calib = Y_test[i][:, h_idx] * discharge_std + discharge_mean | |
| for node_idx, station_code in enumerate(station_codes): | |
| if args.check_quantile_calibration: | |
| # Deliberately NOT gated on station_thresholds | |
| # (unlike the tp/fp/fn/tn block below) -- real | |
| # calibration is a property of the predicted | |
| # quantiles themselves, independent of whether | |
| # this station happens to have a real flood-event | |
| # threshold defined at all. | |
| y_val_calib = Y_real_this_horizon_calib[node_idx] | |
| if not np.isnan(y_val_calib): | |
| station_median_calib = station_train_median.get(station_code, float("nan")) | |
| has_flow_label = not np.isnan(station_median_calib) | |
| is_low_flow_calib = has_flow_label and y_val_calib <= station_median_calib | |
| for q_idx_calib, q_calib in enumerate(QUANTILES): | |
| pred_q_calib = Q_pred_all_quantiles_real[node_idx, h_idx, q_idx_calib] | |
| calib_totals[q_calib] += 1 | |
| hit_q_calib = y_val_calib <= pred_q_calib | |
| if hit_q_calib: | |
| calib_hits[q_calib] += 1 | |
| if has_flow_label: | |
| # Direct low-flow-vs-high-flow split, | |
| # for EVERY real predicted quantile -- | |
| # generalized after a real plot | |
| # (H404021101, full test range) | |
| # visually showed q0.99 detaching from | |
| # observed during quiet periods far | |
| # more than q0.95/q0.9 did, suggesting | |
| # the over-coverage on low-flow days | |
| # isn't uniform across the upper | |
| # quantiles -- this checks that | |
| # directly instead of only at q=0.95. | |
| gap_q_calib = float(pred_q_calib - y_val_calib) | |
| flow_gap_bucket = gap_low_flow if is_low_flow_calib else gap_high_flow | |
| flow_cov_bucket = coverage_low_flow if is_low_flow_calib else coverage_high_flow | |
| flow_gap_bucket[q_calib].append(gap_q_calib) | |
| flow_cov_bucket[q_calib]["total"] += 1 | |
| flow_cov_bucket[q_calib]["hits"] += int(hit_q_calib) | |
| if 0.95 in QUANTILES and is_low_flow_calib and has_real_precip: | |
| q95_idx_calib = QUANTILES.index(0.95) | |
| pred_q95_calib = Q_pred_all_quantiles_real[node_idx, h_idx, q95_idx_calib] | |
| gap_calib = float(pred_q95_calib - y_val_calib) | |
| hit_calib = y_val_calib <= pred_q95_calib | |
| precip_elevated = (recent_precip_sum_real_per_node[node_idx] | |
| > args.calibration_precip_threshold_mm) | |
| bucket_gap = q95_gap_high_precip if precip_elevated else q95_gap_low_precip | |
| bucket_cov = q95_coverage_high_precip if precip_elevated else q95_coverage_low_precip | |
| bucket_gap.append(gap_calib) | |
| bucket_cov["total"] += 1 | |
| bucket_cov["hits"] += int(hit_calib) | |
| if precip_elevated: | |
| # Is this "elevated precipitation" reading | |
| # genuinely recent, or a stale forward- | |
| # filled value from long before the real | |
| # gap started (see --calibration-precip- | |
| # staleness-days' own docstring)? | |
| is_stale = (precip_dt_recent_real_per_node[node_idx] | |
| > args.calibration_precip_staleness_days) | |
| stale_gap_bucket = q95_gap_high_precip_stale if is_stale else q95_gap_high_precip_recent | |
| stale_cov_bucket = (q95_coverage_high_precip_stale if is_stale | |
| else q95_coverage_high_precip_recent) | |
| stale_gap_bucket.append(gap_calib) | |
| stale_cov_bucket["total"] += 1 | |
| stale_cov_bucket["hits"] += int(hit_calib) | |
| if station_code not in station_thresholds: | |
| continue # no real threshold for this station at all -- nothing to evaluate | |
| threshold = station_thresholds[station_code] | |
| predicted_flag = Q_flag_this_horizon[node_idx] >= threshold | |
| is_event = is_real_event_day(station_code, target_date, event_lookup) is not None | |
| if track_forecast_fp: | |
| has_real_forecast = forecast_missing_test[i][node_idx, h_idx] == 0.0 | |
| if is_event and predicted_flag: | |
| confusion[h]["tp"] += 1 | |
| elif is_event and not predicted_flag: | |
| confusion[h]["fn"] += 1 | |
| elif not is_event and predicted_flag: | |
| confusion[h]["fp"] += 1 | |
| if args.check_fp_precipitation and has_real_precip: | |
| fp_precip_values.append(float(recent_precip_sum_real_per_node[node_idx])) | |
| if track_forecast_fp: | |
| key = "fp_with_forecast" if has_real_forecast else "fp_without_forecast" | |
| forecast_fp_stats[h][key] += 1 | |
| if args.dump_fp_details == h: | |
| fp_details.append({ | |
| "station_code": station_code, | |
| "target_date": target_date.date().isoformat(), | |
| "predicted_value": float(Q_flag_this_horizon[node_idx]), | |
| "threshold": float(threshold), | |
| "margin": float(Q_flag_this_horizon[node_idx] - threshold), | |
| "days_to_nearest_real_event": days_to_nearest_real_event( | |
| station_code, target_date, event_lookup, | |
| ), | |
| }) | |
| else: | |
| confusion[h]["tn"] += 1 | |
| if args.check_fp_precipitation and has_real_precip: | |
| tn_precip_values.append(float(recent_precip_sum_real_per_node[node_idx])) | |
| if track_forecast_fp: | |
| key = "all_with_forecast" if has_real_forecast else "all_without_forecast" | |
| forecast_fp_stats[h][key] += 1 | |
| # --extra-model: loads every additional real checkpoint requested | |
| # and runs it over the exact same real X_test, but ONLY at | |
| # --plot-stations-horizon and ONLY for the real requested_stations | |
| # -- this feeds --plot-stations' comparison plot alone, not | |
| # flagging/calibration/spread, which stay on the primary checkpoint. | |
| # A real, separate forward pass per extra model (not free, but only | |
| # over the real stations actually being plotted, not the whole | |
| # real ~4,500-node graph's worth of diagnostics the primary model's | |
| # single pass above also computes). | |
| extra_station_predictions: Dict[str, Dict[str, Dict[str, list]]] = {} | |
| if requested_stations and extra_model_specs: | |
| for extra_name, extra_temporal_type, extra_path in extra_model_specs: | |
| if not extra_path.exists(): | |
| print(f"--extra-model {extra_name}: real checkpoint {extra_path} doesn't exist -- skipping.") | |
| continue | |
| extra_model = SpatiotemporalGNN( | |
| shared_static_dim=x_shared_static_t.shape[1], gauge_static_dim=x_gauge_static_t.shape[1], | |
| gauge_dynamic_dim=X.shape[-1], edge_dim=edge_attr_t.shape[1], | |
| n_basins=len(BASIN_FILE_NAMES), n_targets=2, quantiles=QUANTILES, | |
| temporal_type=extra_temporal_type, | |
| ) | |
| extra_model.load_state_dict(torch.load(extra_path, map_location="cpu")) | |
| extra_model.eval() | |
| print(f"--extra-model {extra_name}: loaded real checkpoint from {extra_path} " | |
| f"(temporal_type={extra_temporal_type}).") | |
| extra_station_predictions[extra_name] = { | |
| code: {"median": [], "q95": [], "quantiles": {q: [] for q in QUANTILES}} | |
| for code in requested_stations | |
| } | |
| with torch.no_grad(): | |
| for i_extra in range(len(X_test)): | |
| x_dynamic_seq_extra = torch.tensor(X_test[i_extra], dtype=torch.float32) | |
| forecast_precip_extra = torch.tensor(forecast_precip_test[i_extra], dtype=torch.float32) | |
| forecast_missing_extra = torch.tensor(forecast_missing_test[i_extra], dtype=torch.float32) | |
| pred_extra = extra_model(x_shared_static_t, x_gauge_static_t, x_dynamic_seq_extra, | |
| is_gauged_t, edge_index_t, edge_attr_t, basin_id_t, horizons_t, | |
| forecast_precip_extra, forecast_missing_extra) | |
| Q_std_extra = pred_extra[:, :, 0, :].numpy() | |
| Q_real_extra = Q_std_extra * discharge_std + discharge_mean | |
| for code_extra in requested_stations: | |
| node_idx_extra = station_codes.index(code_extra) | |
| extra_station_predictions[extra_name][code_extra]["median"].append( | |
| float(Q_real_extra[node_idx_extra, plot_horizon_idx, median_idx_ts])) | |
| extra_station_predictions[extra_name][code_extra]["q95"].append( | |
| float(Q_real_extra[node_idx_extra, plot_horizon_idx, q95_idx_ts])) | |
| for q_idx_extra, q_extra in enumerate(QUANTILES): | |
| extra_station_predictions[extra_name][code_extra]["quantiles"][q_extra].append( | |
| float(Q_real_extra[node_idx_extra, plot_horizon_idx, q_idx_extra])) | |
| extra_temporal_type_by_name = {name: ttype for name, ttype, _ in extra_model_specs} | |
| if requested_stations: | |
| train_end_ts_ts = pd.Timestamp(args.train_end) | |
| train_date_mask_ts = np.array([d <= train_end_ts_ts for d in dates]) | |
| print("\n" + "=" * 70) | |
| print(f"Real station time-series plots, horizon={args.plot_stations_horizon}d") | |
| print("=" * 70) | |
| for code in requested_stations: | |
| node_idx_ts = station_codes.index(code) | |
| # Real, per-station training-period q90 -- kept only as the | |
| # top panel's real optional reference line (matching the | |
| # real "Train q90=..." dashed line this layout was built | |
| # from), NOT as a plotted naive-baseline series anymore -- | |
| # dropped by real, explicit request: the real comparison | |
| # that matters now is model vs. model vs. real observed. | |
| real_train_values = (target_tensor[node_idx_ts, train_date_mask_ts] * discharge_std | |
| + discharge_mean) | |
| real_train_values = real_train_values[~np.isnan(real_train_values)] | |
| if len(real_train_values) == 0: | |
| print(f"Station {code}: no real training-period observations at all -- skipping.") | |
| continue | |
| naive_median = float(np.quantile(real_train_values, 0.5)) | |
| # m3/s, not L/s, for this plot specifically (both the real | |
| # station threshold and train_q90 reference lines, AND | |
| # every real series plotted against them, all converted by | |
| # the SAME real 1000.0 divisor -- L/s and m3/s differ by a | |
| # pure, linear 1000x factor, so this is a real relabeling, | |
| # not a different computation: NSE/KGE/PBIAS/coverage below | |
| # are dimensionless and come out numerically identical | |
| # either way; RMSE/MAE/Peak_RMSE come out 1000x smaller in | |
| # m3/s, correctly, not because the real error shrank). | |
| L_PER_S_TO_M3_PER_S = 1000.0 | |
| train_q90 = float(np.quantile(real_train_values, 0.90)) / L_PER_S_TO_M3_PER_S | |
| threshold = station_thresholds.get(code, float("nan")) | |
| if np.isfinite(threshold): | |
| threshold = threshold / L_PER_S_TO_M3_PER_S | |
| # Real, standard hydrology metrics for THIS model, in the | |
| # same real per-model dict shape the comparison table | |
| # expects -- keyed by --model-name so this plot already | |
| # supports a real second, third, ... model the moment | |
| # another real checkpoint's predictions are added to | |
| # model_predictions below, without changing this call site | |
| # again. | |
| # | |
| # q95, not median -- by real, explicit request: this plot's | |
| # real "prediction" series is now the model's real 0.95 | |
| # quantile, not its median, matching the same real quantile | |
| # --flag-quantile itself defaults to for real flood-flagging | |
| # elsewhere in this script (a coherent, matched choice, not | |
| # an arbitrary swap). | |
| observed_arr = np.array(station_timeseries_data[code]["observed"]) / L_PER_S_TO_M3_PER_S | |
| q95_arr_plot = np.array(station_timeseries_data[code]["q95"]) / L_PER_S_TO_M3_PER_S | |
| median_arr_plot = np.array(station_timeseries_data[code]["median"]) / L_PER_S_TO_M3_PER_S | |
| # --blend-quantiles: real, continuous interplay between the | |
| # median and 0.95 quantile instead of pure q0.95 -- | |
| # predicted = median + alpha*(q95-median), alpha = | |
| # spread/(spread+tau), a real saturating (Michaelis-Menten | |
| # style) weight: alpha->0 (pure median) when the two | |
| # quantiles nearly agree, alpha->1 (pure q95) when they're | |
| # far apart, with no hard switch anywhere in between. tau | |
| # (the spread at which alpha=0.5) defaults to this station's | |
| # own real median (q95-median) spread over its real plotted | |
| # period unless --blend-scale fixes it explicitly. | |
| # --blend-quantiles-max-horizon: keeps this plot consistent | |
| # with the confusion matrix's own real gating above -- if | |
| # the requested plot horizon sits above the real cutoff, | |
| # flagging itself already fell back to raw q0.95 there, so | |
| # this plot falls back the same way rather than silently | |
| # showing a blended series the confusion matrix isn't | |
| # actually using at this horizon. | |
| plot_blend_active = args.blend_quantiles and ( | |
| args.blend_quantiles_max_horizon is None | |
| or args.plot_stations_horizon <= args.blend_quantiles_max_horizon | |
| ) | |
| # Shared real per-model prediction/blend logic -- factored | |
| # out so EVERY model shown on this real comparison plot (the | |
| # primary checkpoint AND any --extra-model) is scored the | |
| # exact same way, an apples-to-apples real comparison rather | |
| # than the primary model getting special-cased treatment. | |
| # median_arr/upper_arr/raw_q95_arr are that ONE model's own | |
| # real predictions (m3/s); naive_median/train_q90 are this | |
| # STATION's own real training-period discharge stats, shared | |
| # across every model since they describe the real river, not | |
| # any one checkpoint. | |
| def _predict_series(median_arr, upper_arr, raw_q95_arr): | |
| if not plot_blend_active: | |
| return raw_q95_arr, "q0.95", None | |
| spread_arr = np.clip(upper_arr - median_arr, 0.0, None) | |
| if args.blend_mode == "level": | |
| naive_median_m3 = naive_median / L_PER_S_TO_M3_PER_S | |
| denom_local = train_q90 - naive_median_m3 | |
| denom_local_safe = denom_local if denom_local > 0 else 1.0 | |
| alpha_arr = np.clip((median_arr - naive_median_m3) / denom_local_safe, 0.0, 1.0) | |
| alpha_arr = np.where(np.isfinite(alpha_arr), alpha_arr, 0.0) | |
| desc = (f"blend(median,q{args.blend_upper_quantile}), mode=level, " | |
| f"train_median={naive_median_m3:.3f}, train_q90={train_q90:.3f} m3/s") | |
| else: | |
| finite_spread = spread_arr[np.isfinite(spread_arr)] | |
| if args.blend_scale is not None: | |
| tau_local = args.blend_scale | |
| else: | |
| tau_local = float(np.median(finite_spread)) if len(finite_spread) else float("nan") | |
| if not np.isfinite(tau_local) or tau_local <= 0: | |
| tau_local = 1.0 | |
| alpha_arr = spread_arr / (spread_arr + tau_local) | |
| desc = f"blend(median,q{args.blend_upper_quantile}), mode=spread, tau={tau_local:.3f} m3/s" | |
| predicted_arr = median_arr + alpha_arr * spread_arr | |
| return predicted_arr, desc, alpha_arr | |
| if args.blend_quantiles and not plot_blend_active: | |
| print(f" --blend-quantiles-max-horizon {args.blend_quantiles_max_horizon}: real plot " | |
| f"horizon {args.plot_stations_horizon} is above it, so this plot/its metrics use " | |
| f"raw q0.95 too (matching flagging's own real fallback at this horizon).") | |
| upper_arr_plot = None | |
| if plot_blend_active: | |
| # Upper anchor for the blend -- args.blend_upper_quantile | |
| # (default 0.95), pulled from the SAME real per-quantile | |
| # cache --compare-all-quantiles reads, not necessarily | |
| # q95_arr_plot (which stays fixed at q0.95 for the | |
| # non-blend default series, a separate, real, deliberate | |
| # choice unrelated to this setting). | |
| upper_arr_plot = (np.array(station_timeseries_data[code]["quantiles"][args.blend_upper_quantile]) | |
| / L_PER_S_TO_M3_PER_S) | |
| predicted_arr_plot, pred_desc, alpha_arr_plot = _predict_series( | |
| median_arr_plot, upper_arr_plot, q95_arr_plot) | |
| finite_alpha = alpha_arr_plot[np.isfinite(alpha_arr_plot)] if alpha_arr_plot is not None else None | |
| # Label format: "<temporal_type>-<model_name>" (e.g. | |
| # "gru-stgnn", "lstm-stgnn") -- real, explicit request, so | |
| # the plot's legend/table names each real model by which | |
| # real temporal encoder produced it, not by whether it's | |
| # blended (that's already in pred_desc, printed separately). | |
| pred_series_label = f"{args.temporal_type}-{args.model_name}" | |
| model_metrics_this = compute_station_model_metrics( | |
| station_timeseries_data[code]["dates"], observed_arr, predicted_arr_plot, code, | |
| event_lookup, threshold, | |
| ) | |
| model_predictions = {pred_series_label: predicted_arr_plot} | |
| model_metrics = {pred_series_label: model_metrics_this} | |
| # --extra-model: every additional real checkpoint requested | |
| # alongside the primary one (e.g. the real gru default plus | |
| # a real lstm checkpoint) -- scored/blended with the exact | |
| # same _predict_series logic above, then added to the same | |
| # real model_predictions/model_metrics dicts so | |
| # plot_station_comparison draws every real model's own line/ | |
| # scatter/table row on one real shared chart. | |
| for extra_name, extra_data in extra_station_predictions.items(): | |
| if code not in extra_data: | |
| continue | |
| extra_median_arr = np.array(extra_data[code]["median"]) / L_PER_S_TO_M3_PER_S | |
| extra_q95_arr = np.array(extra_data[code]["q95"]) / L_PER_S_TO_M3_PER_S | |
| extra_upper_arr = None | |
| if plot_blend_active: | |
| extra_upper_arr = (np.array(extra_data[code]["quantiles"][args.blend_upper_quantile]) | |
| / L_PER_S_TO_M3_PER_S) | |
| extra_predicted_arr, _extra_desc, _extra_alpha = _predict_series( | |
| extra_median_arr, extra_upper_arr, extra_q95_arr) | |
| extra_metrics_this = compute_station_model_metrics( | |
| station_timeseries_data[code]["dates"], observed_arr, extra_predicted_arr, code, | |
| event_lookup, threshold, | |
| ) | |
| extra_temporal_type_for_label = extra_temporal_type_by_name.get(extra_name, extra_name) | |
| extra_label = f"{extra_temporal_type_for_label}-{args.model_name}" | |
| model_predictions[extra_label] = extra_predicted_arr | |
| model_metrics[extra_label] = extra_metrics_this | |
| print(f" Real metrics ({extra_label} {_extra_desc or 'q0.95'} vs. real observed, m3/s, " | |
| f"n={extra_metrics_this['n_real']}): NSE={extra_metrics_this['NSE']:.3f}, " | |
| f"KGE={extra_metrics_this['KGE']:.3f}, RMSE={extra_metrics_this['RMSE']:.3f}, " | |
| f"MAE={extra_metrics_this['MAE']:.3f}, PBIAS={extra_metrics_this['PBIAS']:.1f}%, " | |
| f"Peak_RMSE={extra_metrics_this['Peak_RMSE']:.3f}, POD={extra_metrics_this['POD']:.3f}, " | |
| f"FAR={extra_metrics_this['FAR']:.3f}") | |
| # BASIN_FILE_NAMES maps id -> real display name (e.g. | |
| # 0 -> "eure"), same mapping used to load each basin's own | |
| # files in main() above -- already the real name, not an id. | |
| basin_id_ts = int(combined["basin_id"][node_idx_ts]) | |
| basin_label = BASIN_FILE_NAMES.get(basin_id_ts, None) | |
| plot_path = eval_output_dir / f"station_timeseries_{code}_h{args.plot_stations_horizon}.png" | |
| plot_station_comparison( | |
| station_timeseries_data[code]["dates"], observed_arr, | |
| model_predictions, code, args.plot_stations_horizon, plot_path, | |
| model_metrics=model_metrics, threshold=threshold, train_q90=train_q90, | |
| basin_label=basin_label, | |
| ) | |
| print(f"Station {code}: saved to {plot_path} (train q90={train_q90:.3f} m3/s, " | |
| f"real threshold={threshold:.3f} m3/s)") | |
| if plot_blend_active and len(finite_alpha): | |
| print(f" Real blend weight alpha (0=pure median, 1=pure q{args.blend_upper_quantile}), {pred_desc}: " | |
| f"mean={finite_alpha.mean():.3f}, min={finite_alpha.min():.3f}, " | |
| f"max={finite_alpha.max():.3f}") | |
| print(f" Real metrics ({args.model_name} {pred_desc} vs. real observed, m3/s, " | |
| f"n={model_metrics_this['n_real']}): NSE={model_metrics_this['NSE']:.3f}, " | |
| f"KGE={model_metrics_this['KGE']:.3f}, RMSE={model_metrics_this['RMSE']:.3f}, " | |
| f"MAE={model_metrics_this['MAE']:.3f}, PBIAS={model_metrics_this['PBIAS']:.1f}%, " | |
| f"Peak_RMSE={model_metrics_this['Peak_RMSE']:.3f}, POD={model_metrics_this['POD']:.3f}, " | |
| f"FAR={model_metrics_this['FAR']:.3f}") | |
| if args.compare_all_quantiles: | |
| # Real, direct comparison across EVERY real predicted | |
| # quantile -- not just whichever one --flag-quantile/ | |
| # --blend-quantiles happened to pick -- using the exact | |
| # real per-quantile predictions already cached above in | |
| # station_timeseries_data[code]["quantiles"], so this | |
| # needs no second real model forward pass. | |
| print(f" Real per-quantile comparison (station {code}, h={args.plot_stations_horizon}, " | |
| f"m3/s, vs. real observed, each independently compared to the real threshold " | |
| f"{threshold:.3f} -- not blended with any other real quantile):") | |
| for q_cmp in QUANTILES: | |
| q_arr_cmp = (np.array(station_timeseries_data[code]["quantiles"][q_cmp]) | |
| / L_PER_S_TO_M3_PER_S) | |
| q_metrics_cmp = compute_station_model_metrics( | |
| station_timeseries_data[code]["dates"], observed_arr, q_arr_cmp, code, | |
| event_lookup, threshold, | |
| ) | |
| print(f" q{q_cmp}: NSE={q_metrics_cmp['NSE']:.3f}, PBIAS={q_metrics_cmp['PBIAS']:.1f}%, " | |
| f"RMSE={q_metrics_cmp['RMSE']:.3f}, Peak_RMSE={q_metrics_cmp['Peak_RMSE']:.3f}, " | |
| f"POD={q_metrics_cmp['POD']:.3f}, FAR={q_metrics_cmp['FAR']:.3f}") | |
| if args.trace_fp_inputs: | |
| # Real, direct identification of the specific real dates | |
| # matching the real pattern seen visually: the real 0.95 | |
| # quantile spiking while real observed discharge stayed | |
| # flat and low. Only considers dates where the real | |
| # observation was at or below this station's own real | |
| # median -- a "spike" while ALREADY at high flow isn't | |
| # the same real pattern being investigated here. | |
| dates_arr = station_timeseries_data[code]["dates"] | |
| observed_arr_trace = np.array(station_timeseries_data[code]["observed"]) | |
| # Reuses the SAME real predicted_arr_plot computed above | |
| # (pure q0.95, or the --blend-quantiles blend) -- just | |
| # converted back from m3/s to L/s so this trace's real | |
| # numbers stay in this section's existing real units and | |
| # the false-spike ranking always matches whichever real | |
| # series was actually plotted/scored above, never a | |
| # separate, silently-different pure-q0.95 series. | |
| q95_arr_trace = predicted_arr_plot * L_PER_S_TO_M3_PER_S | |
| example_idx_arr = station_timeseries_data[code]["example_idx"] | |
| node_idx_trace = station_codes.index(code) | |
| valid_trace = ~np.isnan(observed_arr_trace) | |
| low_flow_trace = valid_trace & (observed_arr_trace <= naive_median) | |
| gap = np.where(low_flow_trace, q95_arr_trace - observed_arr_trace, -np.inf) | |
| top_n_idx = np.argsort(gap)[::-1][:args.trace_fp_inputs] | |
| print(f"\n Real top {len(top_n_idx)} false-spike date(s) for {code} ({pred_desc} vs. " | |
| f"real observed, among real low-flow days only):") | |
| discharge_value_idx, discharge_missing_idx = get_dynamic_channel_index( | |
| var_names_for_channels, "discharge", channels_per_var=3) | |
| discharge_dt_idx = discharge_value_idx + 2 | |
| if has_real_precip: | |
| precip_dt_idx = precip_value_idx + 2 | |
| for rank, idx in enumerate(top_n_idx): | |
| if gap[idx] == -np.inf: | |
| continue | |
| example_i = example_idx_arr[idx] | |
| trace_date = dates_arr[idx] | |
| print(f"\n #{rank+1}: real date={trace_date.date()}, real observed=" | |
| f"{observed_arr_trace[idx]:.1f} L/s, real predicted ({pred_desc})=" | |
| f"{q95_arr_trace[idx]:.1f} L/s (real gap={gap[idx]:.1f} L/s)") | |
| window = X_test[example_i] # [lookback, n_nodes, n_channels] | |
| real_discharge_recent = (window[-5:, node_idx_trace, discharge_value_idx] | |
| * discharge_std + discharge_mean) | |
| real_discharge_missing_recent = window[-5:, node_idx_trace, discharge_missing_idx] | |
| real_discharge_dt_recent = window[-5:, node_idx_trace, discharge_dt_idx] | |
| print(f" Real recent discharge (last 5 real lookback days): " | |
| f"{np.round(real_discharge_recent, 1).tolist()}") | |
| print(f" Real discharge missingness flag (1=missing): " | |
| f"{real_discharge_missing_recent.tolist()}") | |
| print(f" Real days-since-last-real-discharge-observation (compressed): " | |
| f"{np.round(real_discharge_dt_recent, 3).tolist()}") | |
| if has_real_precip: | |
| real_precip_standardized_recent = window[-5:, node_idx_trace, precip_value_idx] | |
| real_precip_mm_recent = real_precip_standardized_recent * precip_std + precip_mean | |
| real_precip_dt_recent = window[-5:, node_idx_trace, precip_dt_idx] | |
| print(f" Real recent precipitation (real mm, last 5 real days): " | |
| f"{np.round(real_precip_mm_recent, 2).tolist()}") | |
| print(f" Real days-since-last-real-precipitation-observation (compressed): " | |
| f"{np.round(real_precip_dt_recent, 3).tolist()}") | |
| forecast_this_example = forecast_precip_test[example_i][node_idx_trace, plot_horizon_idx] | |
| forecast_missing_this_example = forecast_missing_test[example_i][node_idx_trace, plot_horizon_idx] | |
| print(f" Real forecast precipitation for this exact horizon: " | |
| f"{float(forecast_this_example):.3f} (missing flag={float(forecast_missing_this_example)})") | |
| if args.ablate_missingness_flag and rank == 0: | |
| # Causal ablation -- does the missingness flag | |
| # alone (not the fill value) change the model's | |
| # output? Holds the filled discharge value fixed | |
| # at exactly what it already is and only flips | |
| # missing=1 -> 0 for this one node, across the | |
| # lookback -- isolating the flag's own causal | |
| # effect from the fill value's. Only run once, | |
| # on the single worst (largest-gap) false-spike | |
| # example, since this is a diagnostic probe, not | |
| # something meant to run for every example. | |
| forecast_precip_ab = torch.tensor(forecast_precip_test[example_i], dtype=torch.float32) | |
| forecast_missing_ab = torch.tensor(forecast_missing_test[example_i], dtype=torch.float32) | |
| x_dynamic_seq_orig = torch.tensor(X_test[example_i], dtype=torch.float32) | |
| q95_quantile_idx = QUANTILES.index(0.95) | |
| with torch.no_grad(): | |
| pred_orig = model(x_shared_static_t, x_gauge_static_t, x_dynamic_seq_orig, is_gauged_t, | |
| edge_index_t, edge_attr_t, basin_id_t, horizons_t, | |
| forecast_precip_ab, forecast_missing_ab) | |
| Q95_orig = (pred_orig[node_idx_trace, plot_horizon_idx, 0, q95_quantile_idx].item() | |
| * discharge_std + discharge_mean) | |
| # Ablation A: flag only (discharge_missing_idx -> 0), | |
| # value and days-since left untouched. | |
| x_flag_off = x_dynamic_seq_orig.clone() | |
| x_flag_off[:, node_idx_trace, discharge_missing_idx] = 0.0 | |
| with torch.no_grad(): | |
| pred_flag_off = model(x_shared_static_t, x_gauge_static_t, x_flag_off, is_gauged_t, | |
| edge_index_t, edge_attr_t, basin_id_t, horizons_t, | |
| forecast_precip_ab, forecast_missing_ab) | |
| Q95_flag_off = (pred_flag_off[node_idx_trace, plot_horizon_idx, 0, q95_quantile_idx].item() | |
| * discharge_std + discharge_mean) | |
| # Ablation B: flag off AND days-since reset to 0 -- | |
| # fully presenting the fill value as if it were a | |
| # genuine same-day observation. | |
| x_full_off = x_dynamic_seq_orig.clone() | |
| x_full_off[:, node_idx_trace, discharge_missing_idx] = 0.0 | |
| x_full_off[:, node_idx_trace, discharge_dt_idx] = 0.0 | |
| with torch.no_grad(): | |
| pred_full_off = model(x_shared_static_t, x_gauge_static_t, x_full_off, is_gauged_t, | |
| edge_index_t, edge_attr_t, basin_id_t, horizons_t, | |
| forecast_precip_ab, forecast_missing_ab) | |
| Q95_full_off = (pred_full_off[node_idx_trace, plot_horizon_idx, 0, q95_quantile_idx].item() | |
| * discharge_std + discharge_mean) | |
| real_observed = observed_arr_trace[idx] | |
| print(f" [ablate-missingness-flag] Causal ablation on this rank-0 example " | |
| f"(real observed={real_observed:.1f} L/s):") | |
| print(f" (1) unmodified: Q0.95={Q95_orig:.1f} L/s") | |
| print(f" (2) missingness flag zeroed: Q0.95={Q95_flag_off:.1f} L/s " | |
| f"(delta={Q95_flag_off - Q95_orig:+.1f} L/s)") | |
| print(f" (3) flag + days-since both zeroed: Q0.95={Q95_full_off:.1f} L/s " | |
| f"(delta={Q95_full_off - Q95_orig:+.1f} L/s)") | |
| # Real, direct check of whether low-flow days are caught as | |
| # well as high-flow days, or whether the peak under- | |
| # prediction problem already seen visually is actually | |
| # broader than just the peaks. Split by THIS station's own | |
| # real test-period median (not the naive baseline's | |
| # training-period median), so the two groups are drawn from | |
| # the exact same real data being evaluated here. | |
| valid_mask = ~np.isnan(observed_arr) | |
| if valid_mask.sum() >= 4: | |
| test_period_median = float(np.nanmedian(observed_arr)) | |
| low_mask = valid_mask & (observed_arr <= test_period_median) | |
| high_mask = valid_mask & (observed_arr > test_period_median) | |
| low_metrics = compute_continuous_metrics(observed_arr[low_mask], predicted_arr_plot[low_mask]) | |
| high_metrics = compute_continuous_metrics(observed_arr[high_mask], predicted_arr_plot[high_mask]) | |
| print(f" Real stratified check (split at this station's own real test-period " | |
| f"median={test_period_median:.1f}):") | |
| print(f" Low-flow days (n={low_metrics['n_real']}): " | |
| f"NSE={low_metrics['NSE']:.3f}, RMSE={low_metrics['RMSE']:.1f}, " | |
| f"MAE={low_metrics['MAE']:.1f}, PBIAS={low_metrics['PBIAS']:.1f}%") | |
| print(f" High-flow days (n={high_metrics['n_real']}): " | |
| f"NSE={high_metrics['NSE']:.3f}, RMSE={high_metrics['RMSE']:.1f}, " | |
| f"MAE={high_metrics['MAE']:.1f}, PBIAS={high_metrics['PBIAS']:.1f}%") | |
| if not np.isnan(low_metrics["PBIAS"]) and not np.isnan(high_metrics["PBIAS"]): | |
| if abs(high_metrics["PBIAS"]) > abs(low_metrics["PBIAS"]) * 2: | |
| print(f" Real bias is concentrated at high flow ({high_metrics['PBIAS']:.1f}% " | |
| f"vs {low_metrics['PBIAS']:.1f}%) -- consistent with a peak-specific " | |
| f"under-prediction problem, not a general low-flow issue.") | |
| elif abs(low_metrics["PBIAS"]) > abs(high_metrics["PBIAS"]) * 2: | |
| print(f" Real bias is concentrated at LOW flow ({low_metrics['PBIAS']:.1f}% " | |
| f"vs {high_metrics['PBIAS']:.1f}%) -- real, direct evidence low-flow " | |
| f"days are NOT being caught as well as high-flow days.") | |
| else: | |
| print(f" Real bias is comparable at both ends -- doesn't point specifically " | |
| f"at either low-flow or high-flow days as the dominant problem.") | |
| else: | |
| print(f" Too few real observed days (n={int(valid_mask.sum())}) to split into " | |
| f"low/high-flow groups meaningfully.") | |
| if args.check_quantile_spread: | |
| print("\n" + "=" * 70) | |
| print("Real (0.95 - median) discharge quantile spread by horizon (real L/s)") | |
| print("=" * 70) | |
| fp_counts_by_horizon = {h: confusion[h]["fp"] for h in HORIZONS} | |
| for h in HORIZONS: | |
| values = spread_by_horizon[h] | |
| if not values: | |
| continue | |
| arr = np.array(values) | |
| print(f"horizon={h:>3}d: mean spread={arr.mean():>10.2f}, median spread={np.median(arr):>10.2f}, " | |
| f"real fp count={fp_counts_by_horizon[h]:>5}") | |
| # A real, direct check of the actual hypothesis: does the | |
| # spread's own growth track where false positives explode | |
| # (day 5 -> day 10 in the real run that motivated this check), | |
| # rather than just eyeballing two printed columns. | |
| short_horizons = [h for h in HORIZONS if h <= 5 and spread_by_horizon[h]] | |
| long_horizons = [h for h in HORIZONS if h > 5 and spread_by_horizon[h]] | |
| if short_horizons and long_horizons: | |
| short_mean_spread = np.mean([np.mean(spread_by_horizon[h]) for h in short_horizons]) | |
| long_mean_spread = np.mean([np.mean(spread_by_horizon[h]) for h in long_horizons]) | |
| ratio = long_mean_spread / short_mean_spread if short_mean_spread > 0 else float("inf") | |
| print(f"\nMean spread, horizons <=5d: {short_mean_spread:.2f}; horizons >5d: {long_mean_spread:.2f} " | |
| f"(ratio: {ratio:.2f}x)") | |
| if ratio > 2.0: | |
| print(f"The real quantile spread grows substantially ({ratio:.2f}x) beyond day 5 -- " | |
| f"consistent with (not proof of) genuine hedging/calibration widening being the " | |
| f"real driver of the false-positive explosion at longer horizons, rather than a " | |
| f"precipitation-specific mechanism.") | |
| else: | |
| print(f"The real quantile spread does NOT grow substantially beyond day 5 ({ratio:.2f}x) " | |
| f"-- this specific hypothesis isn't well supported either; the real explanation " | |
| f"for the false-positive jump likely needs further investigation.") | |
| if args.plot_single_trace: | |
| try: | |
| station_code_req, anchor_str = args.plot_single_trace.split("=") | |
| except ValueError: | |
| print(f"--plot-single-trace '{args.plot_single_trace}' isn't in the expected " | |
| f"'station_code=anchor_date' format.") | |
| return | |
| if station_code_req not in station_codes: | |
| print(f"Station {station_code_req} isn't one of this run's real station codes.") | |
| return | |
| node_idx_req = station_codes.index(station_code_req) | |
| anchor_req = pd.Timestamp(anchor_str) | |
| matches = [i for i, a in enumerate(anchor_dates_test) if pd.Timestamp(a) == anchor_req] | |
| if not matches: | |
| print(f"Anchor date {anchor_str} isn't one of this run's real test anchor dates.") | |
| return | |
| i = matches[0] | |
| with torch.no_grad(): | |
| x_dynamic_seq = torch.tensor(X_test[i], dtype=torch.float32) | |
| forecast_precip_i = torch.tensor(forecast_precip_test[i], dtype=torch.float32) | |
| forecast_missing_i = torch.tensor(forecast_missing_test[i], dtype=torch.float32) | |
| pred = model(x_shared_static_t, x_gauge_static_t, x_dynamic_seq, is_gauged_t, | |
| edge_index_t, edge_attr_t, basin_id_t, horizons_t, | |
| forecast_precip_i, forecast_missing_i) | |
| quantile_values_standardized = pred[node_idx_req, :, 0, :].numpy() # [n_horizons, n_quantiles] | |
| quantile_values_real = quantile_values_standardized * discharge_std + discharge_mean | |
| threshold = station_thresholds.get(station_code_req, float("nan")) | |
| trace_plot_path = eval_output_dir / f"single_trace_{station_code_req}_{anchor_str}.png" | |
| plot_single_trace(HORIZONS, QUANTILES, quantile_values_real, threshold, | |
| ROUTING_HORIZON_COUNT, station_code_req, anchor_str, trace_plot_path) | |
| print(f"\nSaved single-trace plot to {trace_plot_path}") | |
| print(f"Real predicted values (real L/s), station={station_code_req}, anchor={anchor_str}:") | |
| for h_idx, h in enumerate(HORIZONS): | |
| row = ", ".join(f"q{q}={quantile_values_real[h_idx, q_idx]:.1f}" for q_idx, q in enumerate(QUANTILES)) | |
| print(f" horizon={h:>3}d: {row}") | |
| if args.dump_fp_details is not None: | |
| h = args.dump_fp_details | |
| print("\n" + "=" * 70) | |
| print(f"Real false positive details at horizon={h}d (n={len(fp_details)})") | |
| print("=" * 70) | |
| if not fp_details: | |
| print("No real false positives at this horizon -- nothing to dump.") | |
| else: | |
| fp_df = pd.DataFrame(fp_details) | |
| output_path = args.dump_fp_output or (eval_output_dir / f"fp_details_horizon_{h}.csv") | |
| fp_df.to_csv(output_path, index=False) | |
| print(f"Saved to {output_path}") | |
| print("\nReal false positive count by station:") | |
| print(fp_df["station_code"].value_counts().to_string()) | |
| n_near_miss = int((fp_df["days_to_nearest_real_event"] <= 3).sum()) | |
| n_no_event_at_all = int(np.isinf(fp_df["days_to_nearest_real_event"]).sum()) | |
| print(f"\n{n_near_miss}/{len(fp_df)} real false positives are within 3 real days of an actual " | |
| f"real event (near-misses, not random noise); {n_no_event_at_all}/{len(fp_df)} are at " | |
| f"stations with no real event anywhere in the real record at all.") | |
| print(f"\nReal date range of these false positives: " | |
| f"{fp_df['target_date'].min()} to {fp_df['target_date'].max()}") | |
| print(f"Real margin (predicted - threshold) stats: mean={fp_df['margin'].mean():.2f}, " | |
| f"median={fp_df['margin'].median():.2f}, max={fp_df['margin'].max():.2f}") | |
| if args.check_forecast_fp_correlation: | |
| print("\n" + "=" * 70) | |
| print("Real forecast availability: false positives vs. the real base rate, by horizon") | |
| print("=" * 70) | |
| if not forecast_fp_stats: | |
| print(f"No requested horizons fall in FORECAST_LEAD_TIMES {FORECAST_LEAD_TIMES} -- nothing to check.") | |
| for h, stats in forecast_fp_stats.items(): | |
| n_fp = stats["fp_with_forecast"] + stats["fp_without_forecast"] | |
| n_all = stats["all_with_forecast"] + stats["all_without_forecast"] | |
| if n_fp == 0 or n_all == 0: | |
| print(f"horizon={h:>2}d: no real false positives or no real examples at all -- nothing to compare.") | |
| continue | |
| fp_forecast_rate = stats["fp_with_forecast"] / n_fp | |
| base_forecast_rate = stats["all_with_forecast"] / n_all | |
| print(f"horizon={h:>2}d: {fp_forecast_rate*100:.1f}% of real false positives had real forecast " | |
| f"data available, vs. {base_forecast_rate*100:.1f}% base rate across all real examples " | |
| f"(n_fp={n_fp}, n_all={n_all})") | |
| # A real, direct verdict, not just printed numbers to eyeball -- | |
| # flags horizons where false positives are genuinely | |
| # disproportionately drawn from forecast-available examples, | |
| # not just nominally different by a small amount. | |
| flagged_horizons = [] | |
| for h, stats in forecast_fp_stats.items(): | |
| n_fp = stats["fp_with_forecast"] + stats["fp_without_forecast"] | |
| n_all = stats["all_with_forecast"] + stats["all_without_forecast"] | |
| if n_fp == 0 or n_all == 0: | |
| continue | |
| fp_rate = stats["fp_with_forecast"] / n_fp | |
| base_rate = stats["all_with_forecast"] / n_all | |
| if base_rate > 0 and fp_rate / base_rate > 1.5: | |
| flagged_horizons.append(h) | |
| if flagged_horizons: | |
| print(f"\nReal false positives are disproportionately drawn from forecast-available examples " | |
| f"at horizon(s) {flagged_horizons} -- real, direct evidence the forecast-conditioning " | |
| f"pathway itself is contributing to false positives there, not generic horizon-driven " | |
| f"hedging alone.") | |
| else: | |
| print(f"\nNo real horizon shows false positives disproportionately drawn from forecast-" | |
| f"available examples -- this specific hypothesis isn't well supported by this " | |
| f"real evidence either.") | |
| if args.check_fp_precipitation: | |
| if not has_real_precip: | |
| print("\n--check-fp-precipitation requested, but no real precipitation channel is present " | |
| "in this run -- nothing to check.") | |
| elif not fp_precip_values or not tn_precip_values: | |
| print("\n--check-fp-precipitation requested, but there were no real false positives or no " | |
| "real true negatives to compare -- nothing to check.") | |
| else: | |
| fp_arr, tn_arr = np.array(fp_precip_values), np.array(tn_precip_values) | |
| print("\n" + "=" * 70) | |
| print(f"Real recent precipitation (mm, summed over the last {ROUTING_HORIZON_COUNT} real " | |
| f"days), false positives vs. true negatives") | |
| print("=" * 70) | |
| print(f"False positives (n={len(fp_arr)}): mean={fp_arr.mean():.2f}mm, median={np.median(fp_arr):.2f}mm") | |
| print(f"True negatives (n={len(tn_arr)}): mean={tn_arr.mean():.2f}mm, median={np.median(tn_arr):.2f}mm") | |
| if fp_arr.mean() > tn_arr.mean(): | |
| print(f"False positives show real, higher recent precipitation on average " | |
| f"({fp_arr.mean():.2f}mm vs {tn_arr.mean():.2f}mm) -- consistent with (not proof " | |
| f"of) water_balance_loss's ET=0/no-storage gap pushing the model to over-predict " | |
| f"discharge whenever real precipitation is high.") | |
| else: | |
| print(f"False positives do NOT show higher recent precipitation than true negatives -- " | |
| f"this specific hypothesis isn't supported by this real evidence; the real " | |
| f"explanation for weak precision likely lies elsewhere.") | |
| if args.check_quantile_calibration: | |
| print("\n" + "=" * 70) | |
| print("Real quantile calibration: empirical coverage vs. nominal quantile level") | |
| print("=" * 70) | |
| print("Coverage = real fraction of real observed values at or below that quantile's real " | |
| "prediction, across every real (gauge, date, horizon) in the held-out test set. A " | |
| "well-calibrated quantile has coverage close to its own nominal level (e.g. 0.95 -> ~95%); " | |
| "coverage well ABOVE the nominal level means that quantile is systematically too " | |
| "high/wide.") | |
| for q in QUANTILES: | |
| total_q = calib_totals[q] | |
| if total_q == 0: | |
| print(f" q={q}: no real observations to check.") | |
| continue | |
| coverage_q = calib_hits[q] / total_q | |
| flag = " <-- over-wide" if coverage_q - q > 0.05 else (" <-- under-wide" if q - coverage_q > 0.05 else "") | |
| print(f" q={q:.2f}: real empirical coverage={coverage_q:.3f} (n={total_q}){flag}") | |
| print("\nReal coverage, split DIRECTLY by flow regime (low vs. high, relative to each " | |
| "station's own real training-period median) -- for EVERY real predicted quantile, not " | |
| "just 0.95, added after a real plot (H404021101, full test range) visually showed q0.99 " | |
| "detaching from real observed discharge during quiet periods far more than q0.95/q0.9 " | |
| "did:") | |
| low_high_cov: Dict[float, Dict[str, float]] = {} | |
| for q in QUANTILES: | |
| print(f" q={q:.2f}:") | |
| for label, gaps_by_q, cov_by_q in [ | |
| (" Low-flow days ", gap_low_flow, coverage_low_flow), | |
| (" High-flow days", gap_high_flow, coverage_high_flow), | |
| ]: | |
| gaps = gaps_by_q[q] | |
| cov = cov_by_q[q] | |
| if not gaps: | |
| print(f" {label}: no real examples in this bucket.") | |
| continue | |
| gaps_arr = np.array(gaps) | |
| cov_rate = cov["hits"] / cov["total"] if cov["total"] > 0 else float("nan") | |
| low_high_cov.setdefault(q, {})[label.strip()] = cov_rate | |
| print(f" {label} (n={len(gaps_arr)}): mean (q{q}-observed) gap={gaps_arr.mean():.1f} L/s, " | |
| f"median={np.median(gaps_arr):.1f} L/s, real coverage={cov_rate:.3f}") | |
| if "Low-flow days" in low_high_cov.get(q, {}) and "High-flow days" in low_high_cov.get(q, {}): | |
| low_cov = low_high_cov[q]["Low-flow days"] | |
| high_cov = low_high_cov[q]["High-flow days"] | |
| if (low_cov - q) > 0.05 and (q - high_cov) > 0.05: | |
| print(f" -> real, opposite-direction miscalibration at q={q}: over-covers on " | |
| f"low-flow days ({low_cov:.3f}) and under-covers on high-flow days " | |
| f"({high_cov:.3f}).") | |
| if 0.95 in QUANTILES and 0.99 in QUANTILES: | |
| low95, low99 = low_high_cov.get(0.95, {}).get("Low-flow days"), low_high_cov.get(0.99, {}).get("Low-flow days") | |
| if low95 is not None and low99 is not None: | |
| if (low99 - 0.99) > (low95 - 0.95) + 0.02: | |
| print(f"\nReal, direct confirmation the quiet-period over-coverage is WORSE at " | |
| f"q=0.99 ({low99:.3f}, nominal 0.99) than at q=0.95 ({low95:.3f}, nominal " | |
| f"0.95) -- consistent with the visual pattern on the H404021101 plot. The " | |
| f"upper-tail miscalibration is not uniform across quantiles; q0.99 " | |
| f"specifically may be worth excluding or down-weighting if it's used for " | |
| f"anything downstream.") | |
| else: | |
| print(f"\nq=0.99 low-flow over-coverage ({low99:.3f}) is NOT meaningfully worse " | |
| f"than q=0.95's ({low95:.3f}) once measured directly across the full test " | |
| f"set -- the visual q0.99 spikes on the H404021101 plot may be specific to " | |
| f"that station/period rather than a general pattern.") | |
| if 0.95 in QUANTILES and has_real_precip: | |
| print("\nReal low-flow-day Q0.95 behavior, split by whether real recent precipitation " | |
| f"(summed over the preceding {ROUTING_HORIZON_COUNT} real days) exceeded " | |
| f"{args.calibration_precip_threshold_mm}mm:") | |
| for label, gaps, cov in [ | |
| ("Elevated precipitation", q95_gap_high_precip, q95_coverage_high_precip), | |
| ("Quiet (below threshold)", q95_gap_low_precip, q95_coverage_low_precip), | |
| ]: | |
| if not gaps: | |
| print(f" {label}: no real low-flow examples in this bucket.") | |
| continue | |
| gaps_arr = np.array(gaps) | |
| cov_rate = cov["hits"] / cov["total"] if cov["total"] > 0 else float("nan") | |
| print(f" {label} (n={len(gaps_arr)}): mean (Q0.95-observed) gap={gaps_arr.mean():.1f} L/s, " | |
| f"median={np.median(gaps_arr):.1f} L/s, real Q0.95 coverage={cov_rate:.3f}") | |
| if q95_gap_high_precip and q95_gap_low_precip: | |
| mean_high = np.mean(q95_gap_high_precip) | |
| mean_low = np.mean(q95_gap_low_precip) | |
| if mean_high > mean_low * 1.5: | |
| print(f"Real, substantially larger overprediction gap on elevated-precipitation " | |
| f"low-flow days ({mean_high:.1f} L/s vs {mean_low:.1f} L/s) -- direct evidence " | |
| f"a recent real rain pulse is a systematic driver of the remaining false-spike " | |
| f"problem, not a coincidence of the specific examples --trace-fp-inputs surfaced.") | |
| else: | |
| print(f"Real overprediction gap is comparable regardless of recent precipitation " | |
| f"({mean_high:.1f} L/s vs {mean_low:.1f} L/s) -- this specific hypothesis isn't " | |
| f"well supported by this real evidence; the false spikes seen in " | |
| f"--trace-fp-inputs likely reflect general quantile miscalibration rather than " | |
| f"a precipitation-specific trigger.") | |
| # Staleness split of the elevated-precipitation bucket -- | |
| # is the "elevated precipitation" reading genuinely recent | |
| # rain, or a stale forward-filled value being mistaken for | |
| # one (see --calibration-precip-staleness-days' own | |
| # docstring for the two real traced examples that motivated | |
| # this, both of which sat exactly on the 365-day sentinel). | |
| n_stale = len(q95_gap_high_precip_stale) | |
| n_recent = len(q95_gap_high_precip_recent) | |
| n_elevated_total = n_stale + n_recent | |
| if n_elevated_total: | |
| print(f"\nReal staleness split of the 'elevated precipitation' bucket above (real " | |
| f"days-since-last-real-precipitation-observation at the most recent real lookback " | |
| f"day, threshold={args.calibration_precip_staleness_days:.0f} real days) -- " | |
| f"{n_stale}/{n_elevated_total} ({100*n_stale/n_elevated_total:.1f}%) of the " | |
| f"'elevated precipitation' bucket is actually a stale reading, not genuinely recent " | |
| f"rain:") | |
| for label, gaps, cov in [ | |
| ("Stale (old forward-filled reading)", q95_gap_high_precip_stale, q95_coverage_high_precip_stale), | |
| ("Recent (genuinely fresh rain)", q95_gap_high_precip_recent, q95_coverage_high_precip_recent), | |
| ]: | |
| if not gaps: | |
| print(f" {label}: no real examples in this bucket.") | |
| continue | |
| gaps_arr = np.array(gaps) | |
| cov_rate = cov["hits"] / cov["total"] if cov["total"] > 0 else float("nan") | |
| print(f" {label} (n={len(gaps_arr)}): mean (Q0.95-observed) gap={gaps_arr.mean():.1f} " | |
| f"L/s, median={np.median(gaps_arr):.1f} L/s, real Q0.95 coverage={cov_rate:.3f}") | |
| if q95_gap_high_precip_stale and q95_gap_high_precip_recent: | |
| mean_stale = np.mean(q95_gap_high_precip_stale) | |
| mean_recent = np.mean(q95_gap_high_precip_recent) | |
| if mean_stale > mean_recent * 1.2: | |
| print(f"Real, larger overprediction gap on the STALE side ({mean_stale:.1f} L/s vs " | |
| f"{mean_recent:.1f} L/s) -- direct evidence the 'precipitation-driven false " | |
| f"spike' finding is substantially a stale-forward-fill artifact on the " | |
| f"precipitation channel, not genuine recent-rain sensitivity. Because rain, " | |
| f"unlike discharge, has no real hydrological persistence, forward-filling it " | |
| f"across a long real gap is a real design mismatch (the discharge case's " | |
| f"causal-safety argument doesn't carry over) -- worth decaying the " | |
| f"precipitation fill toward a real climatological/zero prior past some gap " | |
| f"length, rather than carrying forward the last real reading indefinitely.") | |
| else: | |
| print(f"Real overprediction gap is comparable between stale and recent readings " | |
| f"({mean_stale:.1f} L/s vs {mean_recent:.1f} L/s) -- the 'elevated " | |
| f"precipitation' effect isn't mainly a staleness artifact; genuinely recent " | |
| f"rain really does drive part of this project's remaining overprediction.") | |
| print("\n" + "=" * 70) | |
| print("Flood-detection confusion matrix by lead time (real held-out test data)") | |
| print("=" * 70) | |
| rows = [] | |
| for h in HORIZONS: | |
| c = confusion[h] | |
| precision = c["tp"] / (c["tp"] + c["fp"]) if (c["tp"] + c["fp"]) > 0 else float("nan") | |
| recall = c["tp"] / (c["tp"] + c["fn"]) if (c["tp"] + c["fn"]) > 0 else float("nan") | |
| # POD (Probability of Detection) is the same real quantity as | |
| # recall, just the standard meteorological/hydrological | |
| # verification name for it -- included as its own column | |
| # alongside recall (not a replacement), so both real naming | |
| # conventions are directly available without ambiguity. | |
| pod = recall | |
| # FAR (False Alarm Ratio) = 1 - precision -- the real fraction | |
| # of flagged events that were NOT real events, a standard | |
| # complement to precision under a different real name. | |
| far = 1.0 - precision if not np.isnan(precision) else float("nan") | |
| # CSI (Critical Success Index / Threat Score) -- a real, single | |
| # metric combining both false positives AND false negatives, | |
| # unlike precision/recall which each only penalize one of the | |
| # two real error types on their own. | |
| csi_denom = c["tp"] + c["fp"] + c["fn"] | |
| csi = c["tp"] / csi_denom if csi_denom > 0 else float("nan") | |
| rows.append({"horizon_days": h, "tp": c["tp"], "fp": c["fp"], "fn": c["fn"], "tn": c["tn"], | |
| "precision": precision, "recall": recall, "POD": pod, "FAR": far, "CSI": csi}) | |
| result_df = pd.DataFrame(rows) | |
| print(result_df.to_string(index=False)) | |
| output_path = eval_output_dir / "flood_detection_evaluation.csv" | |
| result_df.to_csv(output_path, index=False) | |
| print(f"\nSaved to {output_path}") | |
| try: | |
| confusion_plot_path = eval_output_dir / "flood_detection_confusion_matrices.png" | |
| plot_confusion_matrices_by_horizon(confusion, HORIZONS, confusion_plot_path) | |
| print(f"Saved confusion matrix grid to {confusion_plot_path}") | |
| pr_plot_path = eval_output_dir / "flood_detection_precision_recall.png" | |
| plot_precision_recall_by_horizon(result_df, pr_plot_path) | |
| print(f"Saved precision/recall-vs-horizon plot to {pr_plot_path}") | |
| except Exception as e: | |
| print(f" [warning] failed to save evaluation plots: {e}") | |
| if __name__ == "__main__": | |
| main() |