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