Spaces:
Running on Zero
Running on Zero
Download src/graph/physics_losses.py from ageraustine/River_Network: direct link, hf CLI and curl.
- Browser
- Download file 54.2 kB
-
https://huggingface.co/spaces/ageraustine/River_Network/resolve/main/src/graph/physics_losses.py
- Command line
-
hf download hf://spaces/ageraustine/River_Network/src/graph/physics_losses.py
-
curl -L -o physics_losses.py https://huggingface.co/spaces/ageraustine/River_Network/resolve/main/src/graph/physics_losses.py
54.2 kB
| """ | |
| Physics-informed loss terms for streamflow prediction on the reach graph. | |
| Four constraints, each tied to real structure the graph now actually | |
| has (see src/graph/build_reach_graph.py): | |
| 1. confluence_mass_balance_loss -- new mass genuinely enters (is_confluence) | |
| 2. split_rejoin_conservation_loss -- no new mass, paired via braid_id | |
| 3. routing_consistency_loss -- travel-time lag from distance_km/elevation_drop_m | |
| 4. water_balance_loss -- P - ET - Q - deltaS ~= 0 per node | |
| All four apply graph-wide, not just at the 27 gauged nodes -- that's the | |
| actual mechanism by which sparse labels generalize to ~4,500 ungauged | |
| nodes, not an incidental detail. | |
| No model exists yet, so these are standalone functions operating on | |
| whatever Q tensor a model eventually produces. They use only operations | |
| that behave identically on a torch.Tensor or a plain numpy array | |
| (indexing, elementwise arithmetic, sum, mean) so the same code path is | |
| testable now with numpy and will work unchanged with real torch tensors | |
| and real gradients once a model exists -- verified by testing this | |
| module against numpy inputs directly. | |
| """ | |
| from typing import Dict, List, Optional, Tuple | |
| import numpy as np | |
| import pandas as pd | |
| try: | |
| import torch | |
| _HAS_TORCH = True | |
| except ImportError: | |
| _HAS_TORCH = False | |
| def _mse(residual): | |
| """ | |
| Mean squared residual, NaN-masked -- works identically on numpy or | |
| torch. Real ground-truth Q is naturally, heavily NaN (only real | |
| gauges with real observations ever have a value; confirmed against | |
| real data: ~93.5% NaN for the reach graph's discharge tensor) -- | |
| without masking, ANY single NaN anywhere in the residual poisons | |
| the entire mean to NaN, which isn't a rare edge case for this data, | |
| it's the normal shape of it. Returns NaN only if truly nothing | |
| usable exists (every entry NaN), which is a real "no data" signal | |
| worth surfacing, not silently averaging to 0 and implying perfect | |
| physics satisfaction when there was actually no evidence either way. | |
| """ | |
| if _HAS_TORCH and isinstance(residual, torch.Tensor): | |
| mask = ~torch.isnan(residual) | |
| if not mask.any(): | |
| return residual.sum() * float("nan") | |
| return (residual[mask] ** 2).mean() | |
| mask = ~np.isnan(residual) | |
| if not mask.any(): | |
| return np.nan | |
| return (residual[mask] ** 2).mean() | |
| # --------------------------------------------------------------------------- | |
| # 1. Confluence mass balance: Q_confluence ~= sum(Q_upstream_branches) | |
| # --------------------------------------------------------------------------- | |
| def build_confluence_index(nodes_df: pd.DataFrame, edges_df: pd.DataFrame) -> List[Tuple[int, List[int]]]: | |
| """ | |
| Precompute, once per graph (not per training step), which node | |
| indices feed into each real confluence. Returns | |
| [(confluence_idx, [upstream_idx, ...]), ...] using positional | |
| indices into nodes_df (0..n-1), matching how a model's output | |
| tensor would be laid out. | |
| Splitting this out from the loss function itself means the loss can | |
| just do array indexing every step -- the graph structure doesn't | |
| change between training steps, so there's no reason to recompute | |
| which nodes are involved in each confluence on every call. | |
| """ | |
| code_to_idx = {code: i for i, code in enumerate(nodes_df["station_code"])} | |
| confluence_codes = set(nodes_df[nodes_df["is_confluence"]]["station_code"]) | |
| pairs = [] | |
| for conf_code in confluence_codes: | |
| upstream = edges_df[edges_df["target"] == conf_code]["source"].tolist() | |
| upstream_idx = [code_to_idx[u] for u in upstream if u in code_to_idx] | |
| if len(upstream_idx) >= 2: | |
| pairs.append((code_to_idx[conf_code], upstream_idx)) | |
| return pairs | |
| def confluence_mass_balance_loss(Q, confluence_index: List[Tuple[int, List[int]]]): | |
| """ | |
| For each real confluence, predicted discharge there should | |
| approximately equal the sum of its upstream branches' predicted | |
| discharge -- new mass genuinely enters at a confluence (an | |
| independent tributary catchment), so this is a straightforward sum, | |
| unlike the split/rejoin case below. | |
| Ignores travel time between the branches and the confluence (an | |
| instantaneous-mass approximation) -- see routing_consistency_loss | |
| for the piece that accounts for lag separately. | |
| Args: | |
| Q: predicted discharge, shape [n_nodes] (one timestep) or | |
| [n_nodes, T] (multiple timesteps, this loss applies per | |
| timestep the same way). | |
| confluence_index: from build_confluence_index. | |
| Returns: | |
| Scalar loss (0.0, on the same array type as Q, if no confluences). | |
| """ | |
| if not confluence_index: | |
| return Q.sum() * 0.0 # zero, but keeps dtype/type consistent (torch-safe) | |
| residuals = [] | |
| for conf_idx, upstream_idx in confluence_index: | |
| upstream_sum = Q[upstream_idx[0]] | |
| for idx in upstream_idx[1:]: | |
| upstream_sum = upstream_sum + Q[idx] | |
| residuals.append(Q[conf_idx] - upstream_sum) | |
| if _HAS_TORCH and isinstance(Q, torch.Tensor): | |
| residual_stack = torch.stack(residuals) | |
| else: | |
| residual_stack = np.stack(residuals) | |
| return _mse(residual_stack) | |
| # --------------------------------------------------------------------------- | |
| # 2. Split/rejoin conservation: Q_split ~= Q_rejoin (no new mass between them) | |
| # --------------------------------------------------------------------------- | |
| def build_braid_index(nodes_df: pd.DataFrame) -> List[Tuple[int, int]]: | |
| """ | |
| Precompute (split_idx, rejoin_idx) pairs from the saved braid_id | |
| column (see build_reach_graph.py's pair_splits_and_rejoins). Same | |
| precompute-once rationale as build_confluence_index. | |
| """ | |
| code_to_idx = {code: i for i, code in enumerate(nodes_df["station_code"])} | |
| pairs = [] | |
| for _, row in nodes_df[nodes_df["braid_id"].notna()].iterrows(): | |
| rejoin_code, split_code = row["station_code"], row["braid_id"] | |
| if rejoin_code in code_to_idx and split_code in code_to_idx: | |
| pairs.append((code_to_idx[split_code], code_to_idx[rejoin_code])) | |
| return pairs | |
| def split_rejoin_conservation_loss(Q, braid_index: List[Tuple[int, int]]): | |
| """ | |
| For each matched split/rejoin pair, predicted discharge should be | |
| approximately equal at both ends -- the same water dividing into | |
| parallel channels and recombining adds no new mass, unlike a real | |
| confluence (see confluence_mass_balance_loss). This is a genuinely | |
| different physical constraint, not a weaker version of the same one: | |
| a model that learned "sum inflows" generically would get this wrong, | |
| since a rejoin's two branches together should equal the SPLIT's | |
| single value, not add something new on top of it. | |
| Args: | |
| Q: predicted discharge, shape [n_nodes] or [n_nodes, T]. | |
| braid_index: from build_braid_index. | |
| Returns: | |
| Scalar loss (0.0 if no braids in this graph). | |
| """ | |
| if not braid_index: | |
| return Q.sum() * 0.0 | |
| split_idx = [s for s, _ in braid_index] | |
| rejoin_idx = [r for _, r in braid_index] | |
| residual = Q[split_idx] - Q[rejoin_idx] | |
| return _mse(residual) | |
| # --------------------------------------------------------------------------- | |
| # 3. Routing: travel-time lag from real channel distance and slope | |
| # --------------------------------------------------------------------------- | |
| def estimate_travel_time_hours( | |
| distance_km, | |
| elevation_drop_m, | |
| min_velocity_ms: float = 0.1, | |
| max_velocity_ms: float = 3.0, | |
| velocity_coefficient: float = 1.0, | |
| ) -> float: | |
| """ | |
| Rough channel-flow velocity from slope, in the spirit of Manning's | |
| equation's slope dependence (v ~ sqrt(slope)) without the channel | |
| geometry/roughness terms Manning's actually needs, which we don't | |
| have real data for -- explicitly an approximation, not a full | |
| hydraulic solve. Slope = elevation_drop_m / (distance_km * 1000). | |
| Clamped to [min_velocity_ms, max_velocity_ms] since a near-zero or | |
| negative slope (a virtually flat reach, or a data artifact) would | |
| otherwise give a nonsensical near-infinite or negative travel time. | |
| Returns: | |
| Travel time in hours for water to traverse this edge. | |
| """ | |
| distance_m = distance_km * 1000.0 | |
| slope = np.clip(elevation_drop_m / np.maximum(distance_m, 1.0), 1e-6, None) | |
| velocity = np.clip(velocity_coefficient * np.sqrt(slope) * 10.0, min_velocity_ms, max_velocity_ms) | |
| return distance_m / velocity / 3600.0 | |
| def build_routing_index( | |
| nodes_df: pd.DataFrame, edges_df: pd.DataFrame, timestep_hours: float = 24.0, | |
| ) -> List[Tuple[int, int, int]]: | |
| """ | |
| Precompute (upstream_idx, downstream_idx, lag_timesteps) for every | |
| edge, rounding each edge's estimated travel time to the nearest | |
| whole timestep -- e.g. a 30-hour travel time at a 24-hour (daily) | |
| timestep rounds to a 1-step lag. An edge whose travel time rounds to | |
| 0 is still included (same-timestep routing, lag=0). | |
| """ | |
| code_to_idx = {code: i for i, code in enumerate(nodes_df["station_code"])} | |
| pairs = [] | |
| for _, e in edges_df.iterrows(): | |
| if e["source"] not in code_to_idx or e["target"] not in code_to_idx: | |
| continue | |
| drop = e["elevation_drop_m"] if pd.notna(e["elevation_drop_m"]) else 0.1 | |
| hours = estimate_travel_time_hours(e["distance_km"], max(drop, 0.1)) | |
| lag = int(round(hours / timestep_hours)) | |
| pairs.append((code_to_idx[e["source"]], code_to_idx[e["target"]], lag)) | |
| return pairs | |
| def routing_consistency_loss(Q, routing_index: List[Tuple[int, int, int]]): | |
| """ | |
| Q at a downstream node at time t should approximately equal Q at its | |
| upstream node at time (t - lag), lag coming from real distance and | |
| slope (build_routing_index) -- not just "conserve mass at the same | |
| instant," which routing_consistency_loss's siblings above assume as | |
| a simplification. This is the piece that makes that simplification | |
| less necessary over time: a well-trained model satisfying this loss | |
| is learning the actual travel-time behavior of each reach. | |
| Args: | |
| Q: predicted discharge, shape [n_nodes, T] -- REQUIRES a time | |
| dimension, unlike the other three losses, since travel-time | |
| lag is meaningless for a single instant. | |
| routing_index: from build_routing_index. | |
| Returns: | |
| Scalar loss (0.0 if no edges have a usable lag within Q's time range). | |
| """ | |
| T = Q.shape[1] | |
| residuals = [] | |
| for up_idx, down_idx, lag in routing_index: | |
| if lag >= T: | |
| continue # this edge's travel time exceeds the whole prediction window | |
| if lag == 0: | |
| residuals.append(Q[down_idx, :] - Q[up_idx, :]) | |
| else: | |
| residuals.append(Q[down_idx, lag:] - Q[up_idx, :-lag]) | |
| if not residuals: | |
| return Q.sum() * 0.0 | |
| if _HAS_TORCH and isinstance(Q, torch.Tensor): | |
| residual_cat = torch.cat(residuals) | |
| else: | |
| residual_cat = np.concatenate(residuals) | |
| return _mse(residual_cat) | |
| def _soft_dtw_pair_numpy(x: np.ndarray, y: np.ndarray, gamma: float) -> float: | |
| """Reference implementation for the numpy diagnostic path and for | |
| validating the torch path below.""" | |
| n, m = len(x), len(y) | |
| D = (x[:, None] - y[None, :]) ** 2 | |
| NEG_INF = 1e8 | |
| R = np.zeros((n + 1, m + 1)) | |
| R[0, 1:] = NEG_INF | |
| R[1:, 0] = NEG_INF | |
| for i in range(1, n + 1): | |
| for j in range(1, m + 1): | |
| vals = np.array([-R[i - 1, j] / gamma, -R[i - 1, j - 1] / gamma, -R[i, j - 1] / gamma]) | |
| max_val = vals.max() | |
| softmin = -gamma * (max_val + np.log(np.sum(np.exp(vals - max_val)))) | |
| R[i, j] = D[i - 1, j - 1] + softmin | |
| return float(R[n, m]) | |
| def _soft_dtw_pair_torch(x: "torch.Tensor", y: "torch.Tensor", gamma: float) -> "torch.Tensor": | |
| """ | |
| Differentiable soft-DTW between two 1-D sequences (Cuturi & Blondel, | |
| 2017). Builds the DP table as a list-of-lists of fresh scalar | |
| tensors rather than in-place indexed writes into a preallocated | |
| tensor, since in-place writes into a tensor tracked by autograd can | |
| silently break or corrupt gradients. A pure-Python double loop is | |
| fine here given these sequences are short (ROUTING_HORIZON_COUNT | |
| days); a vectorized anti-diagonal version isn't needed at this | |
| length. | |
| """ | |
| n, m = x.shape[0], y.shape[0] | |
| NEG_INF = torch.tensor(1e8, dtype=x.dtype, device=x.device) | |
| ZERO = torch.zeros((), dtype=x.dtype, device=x.device) | |
| D = (x.unsqueeze(1) - y.unsqueeze(0)) ** 2 # [n, m] | |
| R = [[None] * (m + 1) for _ in range(n + 1)] | |
| R[0][0] = ZERO | |
| for j in range(1, m + 1): | |
| R[0][j] = NEG_INF | |
| for i in range(1, n + 1): | |
| R[i][0] = NEG_INF | |
| for i in range(1, n + 1): | |
| for j in range(1, m + 1): | |
| r0, r1, r2 = R[i - 1][j], R[i - 1][j - 1], R[i][j - 1] | |
| stacked = torch.stack([-r0 / gamma, -r1 / gamma, -r2 / gamma]) | |
| softmin = -gamma * torch.logsumexp(stacked, dim=0) | |
| R[i][j] = D[i - 1, j - 1] + softmin | |
| return R[n][m] | |
| def routing_soft_dtw_loss(Q, routing_index: List[Tuple[int, int, int]], gamma: float = 0.001): | |
| """ | |
| Complements routing_consistency_loss rather than replacing it: that | |
| loss assumes a fixed, distance/slope-derived lag is exactly right; | |
| this one compares the predicted discharge shape at an upstream and | |
| downstream node over the same contiguous window via soft-DTW's | |
| elastic time-alignment, tolerant of the lag being slightly off. | |
| Keeps the fixed-lag term rather than dropping it, since letting | |
| timing be entirely learned with no physical anchor risks the same | |
| non-convergence Manning's roughness has shown on real edges. | |
| Args: | |
| Q: predicted discharge, [n_nodes, T]. Always the model's own | |
| dense prediction, never NaN, so no missingness masking is | |
| needed here (unlike the sparse-ground-truth losses above). | |
| routing_index: from build_routing_index, same (up_idx, down_idx, | |
| lag) triples routing_consistency_loss uses. | |
| gamma: soft-DTW smoothing -- smaller tracks literal (hard) DTW | |
| more closely, larger gives a softer, more averaged alignment. | |
| Also directly controls a real, confirmed failure mode: soft- | |
| DTW is NOT guaranteed non-negative except in the gamma->0 | |
| limit (where it converges to true, hard DTW). Two real | |
| rounds of validation went into this default, not one -- | |
| the first (gamma=1.0 -> 0.1) was tested against independent | |
| random sequences and looked safe (0/200 negative), but real | |
| routing edges connect upstream/downstream nodes on the same | |
| river, which are highly correlated with small real | |
| differences, not independent -- a genuinely different data | |
| shape that gamma=0.1 turned out NOT to be safe against | |
| (121-301/500 negative depending on scale, confirmed | |
| directly once real training logs showed this loss reading | |
| exactly 0 -- floored by the defensive clamp -- on every | |
| single real example). gamma=0.001 showed 0/500 negative | |
| against BOTH the realistic correlated case and the original | |
| independent case, with a comfortable positive margin in | |
| both. Validate any change to this default against a | |
| correlated-sequence test, not just an independent one -- | |
| that's the actual gap that caused this to need fixing twice. | |
| Returns: | |
| Scalar loss (0.0 if no edges are usable), averaged across edges, | |
| floored at 0 as a second layer of protection beyond gamma's own | |
| effect -- soft-DTW's negative-value risk shrinks with gamma but | |
| isn't provably eliminated at any finite value, and a loss term | |
| that can go negative is a real problem for gradient-based | |
| training regardless of how rarely it happens. | |
| """ | |
| T = Q.shape[1] | |
| use_torch = _HAS_TORCH and isinstance(Q, torch.Tensor) | |
| pair_fn = _soft_dtw_pair_torch if use_torch else _soft_dtw_pair_numpy | |
| distances = [] | |
| for up_idx, down_idx, lag in routing_index: | |
| if T < 2: | |
| continue | |
| x = Q[up_idx, :] | |
| y = Q[down_idx, :] | |
| distances.append(pair_fn(x, y, gamma)) | |
| if not distances: | |
| return Q.sum() * 0.0 if use_torch else 0.0 | |
| if use_torch: | |
| return torch.clamp(torch.stack(distances).mean(), min=0.0) | |
| return float(max(0.0, np.mean(distances))) | |
| # --------------------------------------------------------------------------- | |
| # 4. Water balance: P - ET - Q - deltaS ~= 0, per node | |
| # --------------------------------------------------------------------------- | |
| def water_balance_loss( | |
| Q_Ls, | |
| precip_mm, | |
| evap_mm, | |
| catchment_area_km2, | |
| period_days: float = 365.0, | |
| delta_storage_m3: Optional[object] = None, | |
| ): | |
| """ | |
| Precipitation minus evapotranspiration minus discharge minus storage | |
| change should balance to ~0, in volume terms, over the given period. | |
| UNIT CONVERSION (the easy part to get subtly wrong, and confirmed as | |
| a REAL bug found while implementing Manning's equation, not just a | |
| hypothetical risk): this function's discharge argument was | |
| previously named/documented as Q_m3s and used directly in m^3/s | |
| math with no conversion -- but this project's real discharge data | |
| is confirmed to actually be in L/s despite column names like | |
| "target_discharge_m3s" implying otherwise (checked directly: real | |
| observed values of 7,095-26,063 for a ~500km2 catchment are | |
| physically absurd as m^3/s -- Amazon-scale flooding -- but exactly | |
| plausible as L/s, 7.1-26.1 m^3/s). Renamed to Q_Ls and converted | |
| internally (/1000) specifically so callers don't have to remember | |
| to convert before calling, which is exactly the kind of thing that | |
| would have silently corrupted this loss by 1000x the first time it | |
| was actually wired into training. | |
| 1 mm of depth over 1 km^2 is 1000 m^3 (1 km^2 = 1e6 m^2, 1 mm = | |
| 1e-3 m, 1e6 * 1e-3 = 1e3). P and ET (mm, over the period) get | |
| converted to m^3 via catchment_area_km2 before comparing against Q, | |
| which is converted from a rate (L/s -> m^3/s) to a volume by | |
| multiplying by the period length in seconds. | |
| delta_storage_m3 defaults to zero (a steady-state approximation) -- | |
| we have no direct storage measurement (soil moisture, groundwater | |
| volume change) in this project's data, only groundwater LEVEL at | |
| sparse wells, which isn't the same thing as a basin-wide storage | |
| volume. Treating deltaS as strictly zero is a real, named | |
| approximation, not a hidden one -- pass a nonzero delta_storage_m3 | |
| if a proxy for it becomes available later (e.g. derived from | |
| groundwater level trend where well coverage allows it). | |
| Args: | |
| Q_Ls: predicted discharge in L/s (this project's real native | |
| unit), shape [n_nodes] (period-average rate). | |
| precip_mm, evap_mm: node features, already available. | |
| catchment_area_km2: from catchment.py -- NaN for ungauged/unknown | |
| catchments, in which case that node is excluded from this | |
| loss entirely (silently including it with a wrong/zero area | |
| would corrupt the term, not just add noise). | |
| period_days: length of the period P/ET/Q are aggregated over. | |
| delta_storage_m3: optional storage change; zero-array default. | |
| Returns: | |
| Scalar loss, computed only over nodes with a real catchment area. | |
| """ | |
| valid_area = ~np.isnan(catchment_area_km2) if not _HAS_TORCH or not isinstance(catchment_area_km2, torch.Tensor) \ | |
| else ~torch.isnan(catchment_area_km2) | |
| # precip_mm has its OWN, separate real NaN pattern (nodes with | |
| # insufficient real precipitation coverage in the lookback window -- | |
| # see recent_precip_sum's completeness check in the training | |
| # script), independent of catchment_area_km2's. Both need excluding | |
| # from `valid`, and both need safe-filling below -- fixing only one | |
| # would leave the exact same 0*NaN=NaN gradient corruption for | |
| # whichever one was missed. | |
| valid_precip = ~np.isnan(precip_mm) if not _HAS_TORCH or not isinstance(precip_mm, torch.Tensor) \ | |
| else ~torch.isnan(precip_mm) | |
| valid = valid_area & valid_precip | |
| # Real, distinct bug from the smoothness fix above -- confirmed | |
| # directly via gradient measurement AFTER that fix was already in | |
| # place, so this is a second, separate issue, not evidence the | |
| # first fix was wrong. Q_median (the learnable input) feeds | |
| # Q_volume_m3 for EVERY node, including the ~84% with real NaN | |
| # catchment_area_km2 -- so P_volume_m3, reference_volume, and | |
| # relative_residual are all genuinely NaN at those positions too, | |
| # even though residual[valid] correctly EXCLUDES them from the | |
| # final loss VALUE. That exclusion does not protect the GRADIENT: | |
| # PyTorch's chain rule still computes a LOCAL derivative at the | |
| # excluded positions, and even though its outer contribution is | |
| # multiplied by zero (from masking), 0 * NaN = NaN in IEEE754 | |
| # floating point, not 0 -- confirmed directly, not assumed. That | |
| # NaN then corrupts the gradient for the whole Q_median tensor. | |
| # Fixed the same way missing data is handled everywhere else in | |
| # this project: NaN is replaced with a safe, finite placeholder | |
| # BEFORE any arithmetic touches it, and the real `valid` mask is | |
| # applied only at the very end, for loss selection -- so no NaN | |
| # ever enters the computational graph in the first place, at any | |
| # position, masked or not. | |
| safe_catchment_area_km2 = torch.where(valid, catchment_area_km2, torch.ones_like(catchment_area_km2)) \ | |
| if _HAS_TORCH and isinstance(catchment_area_km2, torch.Tensor) \ | |
| else np.where(valid, catchment_area_km2, 1.0) | |
| safe_precip_mm = torch.where(valid, precip_mm, torch.zeros_like(precip_mm)) \ | |
| if _HAS_TORCH and isinstance(precip_mm, torch.Tensor) \ | |
| else np.where(valid, precip_mm, 0.0) | |
| Q_m3s = Q_Ls / 1000.0 | |
| period_seconds = period_days * 86400.0 | |
| Q_volume_m3 = Q_m3s * period_seconds | |
| P_volume_m3 = safe_precip_mm * safe_catchment_area_km2 * 1000.0 | |
| ET_volume_m3 = evap_mm * safe_catchment_area_km2 * 1000.0 | |
| dS = delta_storage_m3 if delta_storage_m3 is not None else (Q_m3s * 0.0) | |
| residual = P_volume_m3 - ET_volume_m3 - Q_volume_m3 - dS | |
| # Normalized to a RELATIVE (dimensionless) imbalance, not an | |
| # absolute m^3 residual -- confirmed as a real, necessary fix, not | |
| # a refinement: for a real, large catchment (5935 km2, confirmed | |
| # against real data), even a genuinely period-matched residual | |
| # reached ~1e15 when squared, purely because absolute volumes scale | |
| # with catchment area regardless of how well the model performs. | |
| # Dividing by a reference volume (the larger of P_volume/Q_volume, | |
| # so a real near-zero-flow period doesn't produce a division by a | |
| # near-zero denominator) gives a fractional imbalance instead -- | |
| # directly comparable in scale to this project's other physics | |
| # losses, rather than one term structurally dominating every other | |
| # by orders of magnitude. | |
| # Confirmed as a real, precise, isolated bug -- measured directly | |
| # via per-term gradient checks against real training data, not | |
| # inferred from code review: this term's FORWARD value was always | |
| # finite and reasonable, but its GRADIENT was non-finite in every | |
| # single real example checked (25/25), while every other physics | |
| # term's gradient was finite in all 25. torch.maximum(torch.abs(a), | |
| # torch.abs(b)) has real non-smooth points in its gradient -- at a | |
| # genuine tie (a==b) or at a==0/b==0 -- and real data evidently hit | |
| # one of those points for some of this run's 16 real catchment-area | |
| # nodes. Replaced with a smooth L2-norm-style denominator, which has | |
| # no non-smooth points anywhere (its gradient is well-defined for | |
| # every real input, not just almost every input). | |
| reference_volume = np.sqrt(P_volume_m3**2 + Q_volume_m3**2 + 1e-6) if not _HAS_TORCH or not isinstance(Q_Ls, torch.Tensor) \ | |
| else torch.sqrt(P_volume_m3**2 + Q_volume_m3**2 + 1e-6) | |
| relative_residual = residual / reference_volume | |
| residual_valid = relative_residual[valid] | |
| if (residual_valid.shape[0] if hasattr(residual_valid, "shape") else len(residual_valid)) == 0: | |
| return Q_Ls.sum() * 0.0 | |
| return _mse(residual_valid) | |
| def manning_consistency_loss( | |
| Q_Ls, H_mm, edge_index, edge_attr, manning_n_logit, log_bank_height_m, gate_sharpness: float = 3.0, | |
| ): | |
| """ | |
| Manning's equation as an explicit cross-check between the model's | |
| two outputs: Q = (1/n) * A * R^(2/3) * S^(1/2). Unlike every other | |
| physics term here, which constrains discharge across space (a | |
| confluence, a routing lag) or against external mass-balance | |
| inputs, this is the first constraint tying discharge and water | |
| level TOGETHER -- nothing previously enforced the model's two | |
| predictions were physically consistent with each other at all. | |
| Manning's n (channel roughness) has no real measured value in this | |
| project's data -- learned as a single scalar (log-parameterized so | |
| n = exp(log_n) is always positive, the same trick already proven | |
| for the climate model's tau), initialized near a real, physically | |
| reasonable value for a natural stream (n around 0.03) rather than | |
| an arbitrary starting point. | |
| BANK-HEIGHT RELAXATION -- the actual reason this function exists in | |
| its current form, not just Manning's equation on its own. In-channel | |
| Manning's equation (a single rectangular cross-section) stops being | |
| physically true the moment water goes overbank -- flow spreads onto | |
| a wider, rougher floodplain with completely different geometry. A | |
| full compound-channel model would need real floodplain width and | |
| roughness, neither of which exist in this project's data (no DEM, | |
| no floodplain survey) -- three more ungrounded parameters on top of | |
| an already-uncertain n. Instead: a smooth sigmoid gate multiplies | |
| the loss weight, keyed to a single learned bank_height_m threshold | |
| (log-parameterized for positivity, same pattern as n). Well below | |
| the threshold, gate is ~1 (Manning's is enforced normally). Well | |
| above it, gate decays toward 0 (the model stops being penalized for | |
| deviating from in-channel physics exactly where that physics | |
| genuinely no longer applies). This directly targets the real, | |
| named gap from the physics-discussion: previously nothing in this | |
| model's physics had anything to say about flood conditions | |
| specifically -- the exact regime flood prediction cares about most. | |
| gate_sharpness controls how quickly the transition happens around | |
| the threshold -- fixed rather than learned, deliberately: bank_ | |
| height and n are already two learned scalars trained against a | |
| genuinely small real dataset (100-node subgraph, sparse real | |
| supervision); a third learned parameter controlling the SHAPE of | |
| the relaxation risks being underconstrained rather than adding real | |
| value. A fixed, reasonable sharpness (transition occurs over roughly | |
| +/-1m around the threshold at the default value) is more honest | |
| given the data available than pretending this shape is learnable | |
| from what little real high-water data exists. | |
| UNITS: this project's real discharge is confirmed L/s, water level | |
| is mm (per node_features.py's target_waterlevel_mm convention) -- | |
| both converted internally to Manning's native SI units (m^3/s, m) | |
| rather than assuming callers remember to convert, the same | |
| discipline just applied to water_balance_loss above. | |
| Channel geometry: rectangular-channel approximation (real | |
| trapezoidal/natural cross-sections aren't available -- see | |
| compute_edge_width.py's own docstring on this same limitation), | |
| using each edge's real width_m and slope (elevation_drop_m / | |
| distance_km), and the water level PREDICTED at the edge's | |
| downstream node as the depth proxy. | |
| Args: | |
| Q_Ls: [n_nodes] predicted discharge, L/s. | |
| H_mm: [n_nodes] predicted water level, mm. | |
| edge_index: [2, n_edges]. | |
| edge_attr: [n_edges, 5] -- (distance_km, elevation_drop_m, | |
| verified_continuous, width_m, width_missing_flag), matching | |
| combine_basins' real construction. | |
| manning_n_logit: scalar learnable parameter (nn.Parameter in the | |
| real model; a plain float/array here works identically for | |
| testing without torch). | |
| log_bank_height_m: scalar learnable parameter, same convention. | |
| gate_sharpness: fixed (not learned -- see above), controls how | |
| quickly the relaxation transitions around the threshold. | |
| Returns: | |
| Scalar loss, computed only over edges with real (non-missing) | |
| width data and verified_continuous=True -- the bétoire karst | |
| stretch, where the channel doesn't behave as a normal | |
| conveyance, is deliberately excluded, same as it's excluded | |
| from routing. Each included edge's contribution is additionally | |
| weighted by the bank-height gate before averaging. | |
| """ | |
| target = edge_index[1] | |
| distance_km, elevation_drop_m = edge_attr[:, 0], edge_attr[:, 1] | |
| verified_continuous, width_m, width_missing = edge_attr[:, 2], edge_attr[:, 3], edge_attr[:, 4] | |
| is_np = not _HAS_TORCH or not isinstance(Q_Ls, torch.Tensor) | |
| exp_fn = np.exp if is_np else torch.exp | |
| sqrt_fn = np.sqrt if is_np else torch.sqrt | |
| sigmoid_fn = (lambda x: 1.0 / (1.0 + np.exp(-x))) if is_np else torch.sigmoid | |
| clamp_fn = (lambda x, lo: np.clip(x, lo, None)) if is_np else (lambda x, lo: torch.clamp(x, min=lo)) | |
| clamp_range_fn = (lambda x, lo, hi: np.clip(x, lo, hi)) if is_np else (lambda x, lo, hi: torch.clamp(x, min=lo, max=hi)) | |
| slope = elevation_drop_m / (distance_km * 1000.0 + 1e-9) | |
| slope = clamp_fn(slope, 1e-6) # Manning's eq assumes positive slope; guard real data noise (near-flat/negative segments) | |
| depth_m = H_mm[target] / 1000.0 | |
| # Real, not hypothetical, failure mode: an untrained model's very | |
| # first H_pred values are essentially random, unstandardized | |
| # numbers with no realistic-range guarantee -- and Manning's | |
| # equation scales roughly with depth^(5/3), so a wild early guess | |
| # can overflow to inf in float32 before the model has learned | |
| # anything. That inf then poisons every later forward pass once it | |
| # corrupts a weight via backward(). Upper-bounded at 15m -- a | |
| # physically generous ceiling for these rivers even during a major | |
| # flood, not a realistic operating value -- specifically to stop | |
| # the explosion at its source, not just catch it downstream. | |
| depth_m = clamp_range_fn(depth_m, 1e-3, 15.0) | |
| area = width_m * depth_m | |
| wetted_perimeter = width_m + 2.0 * depth_m | |
| hydraulic_radius = area / (wetted_perimeter + 1e-9) | |
| # Same class of steep-local-derivative risk as n below, via a | |
| # different operation: d(x^(2/3))/dx = (2/3)*x^(-1/3) also diverges | |
| # as x->0, and depth_m's existing lower clamp (1e-3) is small enough | |
| # that hydraulic_radius could still land in this dangerous region | |
| # for some real (width, depth) combinations. Bounded away from zero | |
| # for the same reason n is bounded below -- prevents this specific | |
| # operation's gradient from being able to explode, independent of | |
| # whatever else does or doesn't turn out to have caused a given | |
| # real failure. | |
| hydraulic_radius = clamp_fn(hydraulic_radius, 0.01) | |
| # n bounded to a physically real range for natural channels via a | |
| # sigmoid reparameterization, not a hard clamp on an exponentiated | |
| # value -- replaces an earlier log-parameterization + torch.clamp | |
| # that had a confirmed, real failure mode: clamp's gradient is | |
| # exactly zero outside its range (verified directly via finite | |
| # differences), so once exp(the raw parameter) drifted past 0.3 or | |
| # below 0.01, Manning's loss could never pull it back -- confirmed | |
| # as the real cause of n landing at 0.38-1.0 across several actual | |
| # training runs and never converging. n = 0.01 + 0.29*sigmoid(raw) | |
| # is always in [0.01, 0.3] by construction, and its gradient, while | |
| # it does shrink far from the center, never hits exactly zero the | |
| # way a hard clamp's does -- so the optimizer always retains at | |
| # least some signal to correct course. | |
| n = 0.01 + 0.29 * sigmoid_fn(manning_n_logit) | |
| Q_manning_m3s = (1.0 / n) * area * hydraulic_radius ** (2.0 / 3.0) * sqrt_fn(slope) | |
| Q_manning_Ls = Q_manning_m3s * 1000.0 | |
| Q_model_Ls = Q_Ls[target] | |
| residual = Q_manning_Ls - Q_model_Ls | |
| # Bank-height relaxation gate: ~1 well below the learned threshold | |
| # (Manning's enforced normally), decaying toward ~0 well above it | |
| # (the model isn't penalized for deviating from in-channel physics | |
| # once conditions are genuinely overbank). See the module docstring | |
| # above for why this is a fixed-sharpness sigmoid rather than a | |
| # full compound-channel model. bank_height_m also bounded (0.1m to | |
| # 20m) for the same reason as n above -- the sigmoid gate itself | |
| # saturates naturally so this is lower-risk than n was, but bounding | |
| # it costs nothing and closes the same class of failure mode. | |
| bank_height_m = clamp_range_fn(exp_fn(log_bank_height_m), 0.1, 20.0) | |
| gate = sigmoid_fn(-gate_sharpness * (depth_m - bank_height_m)) | |
| valid = (width_missing < 0.5) & (verified_continuous > 0.5) | |
| residual_valid = residual[valid] | |
| gate_valid = gate[valid] | |
| n_valid = residual_valid.shape[0] if hasattr(residual_valid, "shape") else len(residual_valid) | |
| if n_valid == 0: | |
| return Q_Ls.sum() * 0.0 | |
| # Gate-weighted mean squared residual, NaN-safe (same principle as | |
| # _mse -- a genuinely missing value must not poison every other | |
| # edge's real contribution), but weighted rather than uniform: an | |
| # edge that's gated toward 0 (deep overbank water) should barely | |
| # count toward this loss, not be averaged in as if Manning's | |
| # equation were still fully binding there. | |
| # | |
| # Divides by the COUNT of real edges, not the SUM of gate weights -- | |
| # confirmed as a real bug during testing, not a hypothetical one: | |
| # dividing by the weight sum makes this a weighted MEAN, and for a | |
| # single edge (or edges all near the same gate value), gate*r^2 / | |
| # gate collapses back to just r^2 regardless of how small gate is -- | |
| # the relaxation completely cancels out. Dividing by count instead | |
| # means an edge with gate~0 genuinely contributes ~0 to the total, | |
| # which is the actual, intended behavior: an overbank edge's | |
| # contribution should shrink in absolute terms, not just lose | |
| # relative standing within an average. | |
| # | |
| # isfinite, not isnan, here and below -- confirmed as a real, not | |
| # hypothetical, gap: an untrained model produced an actually- | |
| # infinite prediction early in training, and isnan(inf) is False, | |
| # so it sailed straight past a nan-only check into a real training | |
| # run's validation loss (observed directly: "val supervised loss = | |
| # inf", not nan). depth_m is now upper-bounded above, but Q_model_Ls | |
| # itself (the model's own, equally unstandardized-random early | |
| # discharge output) could independently cause the same failure, so | |
| # this checks both explicitly rather than trusting the depth clamp | |
| # alone to be the only place this could go wrong. | |
| is_finite = np.isfinite(residual_valid) if is_np else torch.isfinite(residual_valid) | |
| has_any = is_finite.any() if hasattr(is_finite, "any") else any(is_finite) | |
| if not has_any: | |
| return residual_valid.sum() * float("nan") if not is_np else np.nan | |
| weighted_sq = gate_valid[is_finite] * (residual_valid[is_finite] ** 2) | |
| n_finite = is_finite.sum() if hasattr(is_finite, "sum") else sum(is_finite) | |
| # Final defense-in-depth ceiling -- even a finite-but-astronomically- | |
| # large residual (e.g. a partially-trained model still producing | |
| # unrealistic values) could overflow when squared and summed; capping | |
| # each term before summing means this function can never itself | |
| # return inf, regardless of what produced the residual. | |
| weighted_sq = np.minimum(weighted_sq, 1e12) if is_np else torch.clamp(weighted_sq, max=1e12) | |
| return weighted_sq.sum() / n_finite | |
| # --------------------------------------------------------------------------- | |
| # Combined loss | |
| # --------------------------------------------------------------------------- | |
| def physics_informed_loss( | |
| Q_supervised_pred, Q_supervised_true, gauged_mask, | |
| Q_full, confluence_index, braid_index, | |
| weights: Optional[Dict[str, float]] = None, | |
| Q_timeseries=None, routing_index=None, | |
| precip_mm=None, evap_mm=None, catchment_area_km2=None, | |
| ) -> Dict[str, float]: | |
| """ | |
| Combines the supervised loss (masked to gauged nodes) with all | |
| physics terms that have the inputs to compute (routing and water | |
| balance are optional -- they need a time dimension / climate data | |
| respectively, which not every training step may have on hand). | |
| Returns a dict of every individual term plus 'total', rather than | |
| just the summed scalar -- so it's possible to see which physics | |
| term is actually driving the loss during training, not just that | |
| "the loss" went up or down. | |
| """ | |
| weights = weights or {"confluence": 1.0, "split_rejoin": 1.0, "routing": 1.0, "water_balance": 1.0} | |
| supervised_residual = (Q_supervised_pred - Q_supervised_true)[gauged_mask] | |
| losses = {"supervised": _mse(supervised_residual)} | |
| losses["confluence"] = confluence_mass_balance_loss(Q_full, confluence_index) | |
| losses["split_rejoin"] = split_rejoin_conservation_loss(Q_full, braid_index) | |
| if Q_timeseries is not None and routing_index is not None: | |
| losses["routing"] = routing_consistency_loss(Q_timeseries, routing_index) | |
| if precip_mm is not None and evap_mm is not None and catchment_area_km2 is not None: | |
| losses["water_balance"] = water_balance_loss(Q_full, precip_mm, evap_mm, catchment_area_km2) | |
| total = losses["supervised"] | |
| for name, w_key in [("confluence", "confluence"), ("split_rejoin", "split_rejoin"), | |
| ("routing", "routing"), ("water_balance", "water_balance")]: | |
| if name in losses: | |
| total = total + weights.get(w_key, 1.0) * losses[name] | |
| losses["total"] = total | |
| return losses | |
| # --------------------------------------------------------------------------- | |
| # 6. NSE (Nash-Sutcliffe Efficiency): variance-tracking, a genuinely | |
| # different failure mode from pinball loss's calibration objective | |
| # --------------------------------------------------------------------------- | |
| def compute_per_node_historical_mean(target_tensor: np.ndarray) -> np.ndarray: | |
| """ | |
| Per-node mean of whatever real (non-NaN) observations exist in | |
| target_tensor. Meant to be computed once from the training-period | |
| tensor (matching how discharge_mean/discharge_std are computed | |
| elsewhere: fit on train, never on val/test), not recomputed | |
| per-batch -- a within-batch mean would be noisy for exactly the | |
| nodes that need this most. | |
| Nodes with zero real observations anywhere get NaN, not a | |
| fabricated 0; nse_loss's own masking excludes them naturally. | |
| """ | |
| import warnings | |
| with warnings.catch_warnings(): | |
| warnings.filterwarnings("ignore", message="Mean of empty slice") | |
| means = np.nanmean(target_tensor, axis=1) | |
| return means.astype(np.float32) | |
| def compute_per_node_historical_median(target_tensor: np.ndarray) -> np.ndarray: | |
| """ | |
| Per-node MEDIAN of whatever real (non-NaN) observations exist in | |
| target_tensor -- deliberately separate from | |
| compute_per_node_historical_mean, not a replacement for it. | |
| NSE's formula (1 - sum((obs-pred)^2)/sum((obs-mean)^2)) mathematically | |
| requires the real MEAN specifically -- swapping mean for median | |
| there would silently stop computing NSE at all. This function | |
| exists for a genuinely different real use: a robust, real, per- | |
| node FILL VALUE for prepare_graph_training_windows' missing-value | |
| handling. | |
| Real, confirmed reason median, not mean, is the right choice there: | |
| a real, direct trace showed the per-node MEAN fill value for one | |
| real, low-flow station (H605641401) was itself ~2023 L/s -- still | |
| roughly 30x that station's real low-flow values (real ~57-70 L/s) | |
| -- because real discharge is right-skewed (rare, real high-flow | |
| events pull the real mean well above what a "typical" day actually | |
| looks like). The real median is far more robust to that real skew, | |
| landing much closer to a station's real, typical day. | |
| Nodes with zero real observations anywhere get NaN, not a | |
| fabricated 0 -- the caller (prepare_graph_training_windows) falls | |
| back to the original, global 0.0 fill for such nodes, same as | |
| compute_per_node_historical_mean's own NaN nodes already do. | |
| """ | |
| import warnings | |
| with warnings.catch_warnings(): | |
| warnings.filterwarnings("ignore", message="All-NaN slice encountered") | |
| medians = np.nanmedian(target_tensor, axis=1) | |
| return medians.astype(np.float32) | |
| def compute_per_node_historical_std(target_tensor: np.ndarray, epsilon: float = 1e-3) -> np.ndarray: | |
| """ | |
| Per-node std of whatever real (non-NaN) observations exist in | |
| target_tensor -- same real, training-period-only discipline as | |
| compute_per_node_historical_mean. Built specifically for real, | |
| per-station-relative quantile loss normalization: real, direct | |
| evidence (check_discharge_quantile_distribution.py) showed this | |
| project's 6 real gauges have WILDLY different real discharge | |
| scales and skew (Q50/Q01 ranging from 1.4x at one real station to | |
| literally infinite -- real zero -- at another), so a single, | |
| GLOBAL discharge_std treats a real, proportionally large error at | |
| a small station and a real, proportionally tiny error at a large | |
| station as if they carried the same real weight in the loss, which | |
| they do not. | |
| A real, small floor (epsilon) prevents division by a near-zero std | |
| for a station with real, genuinely low variability (e.g. the two | |
| Eure stations found to be nearly flat at low flow) -- without this, | |
| such a station's ALREADY small real errors would be divided by an | |
| even smaller real denominator, exploding into a disproportionate, | |
| destabilizing loss contribution -- the same class of numerical | |
| risk KGE's real bias-ratio term already showed us directly. | |
| Nodes with zero real observations, or only one, get NaN (std is | |
| undefined from a single point) -- callers must handle this the | |
| same way nse_loss already handles a NaN per_node_mean. | |
| """ | |
| import warnings | |
| with warnings.catch_warnings(): | |
| warnings.filterwarnings("ignore", message="Degrees of freedom <= 0 for slice") | |
| stds = np.nanstd(target_tensor, axis=1) | |
| stds = np.where(stds < epsilon, epsilon, stds) | |
| n_real_per_node = np.sum(~np.isnan(target_tensor), axis=1) | |
| stds = np.where(n_real_per_node < 2, np.nan, stds) | |
| return stds.astype(np.float32) | |
| def nse_loss(Y_true, Y_pred, per_node_mean, epsilon: float = 1e-6): | |
| """ | |
| 1 - NSE, framed as a loss to minimize. A different signal from | |
| pinball/quantile loss: pinball loss checks whether predicted | |
| quantiles are well-calibrated; NSE checks whether the median | |
| prediction tracks each station's own day-to-day variability, or | |
| degenerates toward a flat, near-mean prediction regardless of | |
| quantile calibration elsewhere. | |
| Normalized by each station's own historical variance (via | |
| per_node_mean), not a shared global one -- a small, calm stream and | |
| a large, variable one need different absolute error tolerances. | |
| Applies to the median prediction only, matching every other loss | |
| in this file. | |
| Args: | |
| Y_true: sparse ground truth, [n_nodes, ...], NaN-masked as | |
| elsewhere in this file. | |
| Y_pred: median prediction, same shape as Y_true. | |
| per_node_mean: [n_nodes], from compute_per_node_historical_mean, | |
| broadcast against Y_true/Y_pred's leading node dimension. | |
| epsilon: prevents division by zero for a node whose historical | |
| variance is near-zero. | |
| Returns: | |
| Scalar loss (NaN if no real ground truth exists at all). | |
| """ | |
| is_torch = _HAS_TORCH and isinstance(Y_true, torch.Tensor) | |
| if is_torch: | |
| mean_broadcast = per_node_mean if per_node_mean.dim() == Y_true.dim() else \ | |
| per_node_mean.view(-1, *([1] * (Y_true.dim() - 1))) | |
| mask = ~torch.isnan(Y_true) | |
| if not mask.any(): | |
| return Y_true.sum() * float("nan") | |
| numerator = ((Y_pred - Y_true)[mask] ** 2).sum() | |
| denominator = ((Y_true - mean_broadcast).expand_as(Y_true)[mask] ** 2).sum() + epsilon | |
| return numerator / denominator | |
| mean_broadcast = per_node_mean if per_node_mean.ndim == Y_true.ndim else \ | |
| per_node_mean.reshape(-1, *([1] * (Y_true.ndim - 1))) | |
| mask = ~np.isnan(Y_true) | |
| if not mask.any(): | |
| return np.nan | |
| numerator = ((Y_pred - Y_true)[mask] ** 2).sum() | |
| denominator = ((Y_true - np.broadcast_to(mean_broadcast, Y_true.shape))[mask] ** 2).sum() + epsilon | |
| return float(numerator / denominator) | |
| def kge_loss_from_pooled_pairs(obs, pred, epsilon: float = 1e-6, max_loss: float = 20.0): | |
| """ | |
| 1 - KGE (Kling-Gupta Efficiency), framed as a loss to minimize -- | |
| a real, temporally-meaningful replacement for nse_loss, not a | |
| cross-station approximation. KGE's correlation term (r) needs | |
| multiple real, PAIRED (observed, predicted) points for the SAME | |
| node to mean anything -- unlike nse_loss's formula, which only | |
| needs a node's own pre-computed historical mean and works fine | |
| within a single, sparse per-example training step, a real | |
| correlation genuinely requires several real points. This function | |
| expects the CALLER to have already gathered those real pairs for | |
| one node, pooled across multiple real anchor dates within a batch | |
| (see the real pooling logic in train_spatiotemporal_gnn.py's batch | |
| loop) -- it does not gather or mask NaN itself, unlike every other | |
| loss in this file, because by the time this function is called, | |
| only real, already-matched pairs should remain. | |
| KGE = 1 - sqrt((r-1)^2 + (alpha-1)^2 + (beta-1)^2), where: | |
| r = real Pearson correlation between obs and pred (timing) | |
| alpha = std(pred) / std(obs) (variability ratio) | |
| beta = mean(pred) / mean(obs) (bias ratio) | |
| CRITICAL, CONFIRMED REQUIREMENT: obs/pred MUST be in real, | |
| UN-STANDARDIZED physical units (real L/s, real cm) here, NOT | |
| standardized (z-scored) values -- confirmed directly as a real, | |
| structural bug, not a rare edge case: on realistic standardized- | |
| space-like pooled samples (mean~0, std~1, exactly what a small, | |
| sparse real pool of z-scored values looks like), beta = pred_mean | |
| / obs_mean exploded to a worst-case loss of 8,260 across 1000 | |
| trials, with 4% exceeding loss=10 -- versus a worst case of 4.07, | |
| zero trials over 10, on realistic REAL, physical-unit discharge | |
| values. Standardization deliberately centers data at mean~0, which | |
| is exactly the condition that makes a mean-ratio term numerically | |
| unstable; real, physical units don't have that problem, since a | |
| real station's real discharge is never centered near zero. The | |
| caller (train_spatiotemporal_gnn.py) un-standardizes before pooling | |
| for exactly this reason -- do not call this on standardized values. | |
| Args: | |
| obs, pred: 1-D, same length, already real (no NaN), in REAL, | |
| UN-STANDARDIZED physical units -- every real (observed, | |
| predicted) pair for ONE node, pooled across however many | |
| real anchor dates/horizons contributed a real observation | |
| within the current batch. | |
| epsilon: prevents division by zero when a node's real pooled | |
| std or mean is at or near zero. | |
| max_loss: a real, measured defensive ceiling, not a guess -- | |
| even in real, physical units, a small pooled sample | |
| combined with a genuinely bad prediction can still produce | |
| a real, large loss (measured worst case 28.92 across 3000 | |
| trials including deliberately bad predictions, with only | |
| 0.5% exceeding 20) -- this is a second, independent layer | |
| of protection on top of using real units, not a substitute | |
| for it. | |
| Returns: | |
| Scalar loss, clamped to [0, max_loss]. Callers should require | |
| at least ~3-5 real pooled points before calling this (checked | |
| by the caller, not here) -- a correlation from 2 points is not | |
| real evidence of timing skill, it is guaranteed to be +-1 or | |
| undefined. | |
| """ | |
| is_torch = _HAS_TORCH and isinstance(obs, torch.Tensor) | |
| if is_torch: | |
| obs_mean, pred_mean = obs.mean(), pred.mean() | |
| obs_std, pred_std = obs.std(unbiased=False), pred.std(unbiased=False) | |
| cov = ((obs - obs_mean) * (pred - pred_mean)).mean() | |
| r = cov / (obs_std * pred_std + epsilon) | |
| alpha = pred_std / (obs_std + epsilon) | |
| beta = pred_mean / (obs_mean + epsilon) | |
| kge = 1.0 - torch.sqrt((r - 1) ** 2 + (alpha - 1) ** 2 + (beta - 1) ** 2) | |
| return torch.clamp(1.0 - kge, min=0.0, max=max_loss) | |
| obs_mean, pred_mean = obs.mean(), pred.mean() | |
| obs_std, pred_std = obs.std(), pred.std() | |
| cov = ((obs - obs_mean) * (pred - pred_mean)).mean() | |
| r = cov / (obs_std * pred_std + epsilon) | |
| alpha = pred_std / (obs_std + epsilon) | |
| beta = pred_mean / (obs_mean + epsilon) | |
| kge = 1.0 - np.sqrt((r - 1) ** 2 + (alpha - 1) ** 2 + (beta - 1) ** 2) | |
| return float(np.clip(1.0 - kge, 0.0, max_loss)) | |
| # --------------------------------------------------------------------------- | |
| # 7. Gated spatial smoothness (Dirichlet energy / graph Laplacian) -- | |
| # only where continuity is actually expected | |
| # --------------------------------------------------------------------------- | |
| def spatial_smoothness_loss( | |
| Q, edge_index, distance_km, verified_continuous, epsilon: float = 0.1, | |
| ): | |
| """ | |
| Dirichlet-energy / graph-Laplacian smoothness term -- predicted | |
| values at two connected, hydrologically-continuous nodes should | |
| vary smoothly along a reach, not erratically. | |
| Gated on verified_continuous, not applied uniformly graph-wide: La | |
| Risle's bétoire stretch has 3 confirmed karst edges where water | |
| genuinely disappears and can resurge elsewhere, so downstream | |
| values there should differ sharply from upstream ones. Applying | |
| smoothness across those edges would penalize the model for | |
| correctly representing that. They're excluded entirely, not | |
| down-weighted, since even a partial penalty there is still wrong. | |
| Caveat: verified_continuous currently reflects only what's been | |
| confirmed by station-naming evidence (3 edges). A separate cavity- | |
| proximity cross-check found other stations sitting just as close to | |
| real natural cavities without confirmed karst status -- this gate | |
| protects against what's confirmed discontinuous, not everything | |
| that might be. | |
| Weighted by inverse distance (1 / (distance_km + epsilon)) -- closer | |
| nodes are expected to behave more similarly than distant ones, so | |
| smoothness is enforced more strongly at short range. | |
| Args: | |
| Q: predicted discharge, [n_nodes, ...] -- always the model's | |
| own dense prediction, never NaN, so no masking is needed. | |
| edge_index: [2, n_edges] (upstream, downstream) node index pairs. | |
| distance_km: [n_edges] distance per edge. | |
| verified_continuous: [n_edges] boolean/0-1 flag -- edges where | |
| this is False are excluded entirely. | |
| epsilon: prevents an unboundedly large weight for a very small | |
| distance. | |
| Returns: | |
| Scalar loss (0.0 if no continuous edges are available). | |
| """ | |
| is_torch = _HAS_TORCH and isinstance(Q, torch.Tensor) | |
| if is_torch: | |
| continuous_mask = verified_continuous.bool() if torch.is_tensor(verified_continuous) else \ | |
| torch.tensor(verified_continuous, dtype=torch.bool, device=Q.device) | |
| if not continuous_mask.any(): | |
| return Q.sum() * 0.0 | |
| up_idx = edge_index[0][continuous_mask] | |
| down_idx = edge_index[1][continuous_mask] | |
| dist = distance_km[continuous_mask] if torch.is_tensor(distance_km) else \ | |
| torch.tensor(np.asarray(distance_km)[continuous_mask.cpu().numpy()], dtype=Q.dtype, device=Q.device) | |
| weights = 1.0 / (dist + epsilon) | |
| diffs_sq = (Q[up_idx] - Q[down_idx]) ** 2 | |
| # weights broadcasts against any extra trailing dims (e.g. a | |
| # horizon axis) the same way distance itself has none of those | |
| # dims -- one real weight per edge, applied identically across | |
| # whatever else Q carries per node. | |
| while weights.dim() < diffs_sq.dim(): | |
| weights = weights.unsqueeze(-1) | |
| weighted = weights * diffs_sq | |
| return weighted.sum() / (weights.sum() * diffs_sq.shape[1:].numel() if diffs_sq.dim() > 1 else weights.sum()) | |
| continuous_mask = np.asarray(verified_continuous).astype(bool) | |
| if not continuous_mask.any(): | |
| return 0.0 | |
| up_idx = np.asarray(edge_index[0])[continuous_mask] | |
| down_idx = np.asarray(edge_index[1])[continuous_mask] | |
| dist = np.asarray(distance_km)[continuous_mask] | |
| weights = 1.0 / (dist + epsilon) | |
| diffs_sq = (Q[up_idx] - Q[down_idx]) ** 2 | |
| weights_b = weights if diffs_sq.ndim == 1 else weights.reshape(-1, *([1] * (diffs_sq.ndim - 1))) | |
| weighted = weights_b * diffs_sq | |
| return float(weighted.sum() / (np.broadcast_to(weights_b, diffs_sq.shape).sum())) |