River_Network / src /graph /physics_losses.py
ageraustine's picture
Upload folder using huggingface_hub (part 2)
04d6dda verified
Raw History Blame Contribute Delete
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()))