River_Network / src /graph /subgraph_selection.py
ageraustine's picture
Upload folder using huggingface_hub (part 2)
f2046b4 verified
Raw History Blame Contribute Delete
7.49 kB
"""
Selects a small subgraph -- real gauges plus a capped number of real
confluences/braid nodes -- for fast training iteration, without losing
real connectivity the way a naive node filter would.
THE REAL PROBLEM THIS SOLVES: most edges between two gauges (or a gauge
and a confluence) pass through several virtual infill nodes. Simply
dropping every non-selected node and keeping only edges where BOTH
endpoints survive would leave most selected nodes disconnected from
each other -- their only real path ran through the nodes just removed.
That would break both message passing (no path for information to
flow) and the physics losses (routing_consistency_loss needs a real
lag between connected nodes; confluence_mass_balance_loss needs real
upstream branches).
THE FIX: for each dropped chain of virtual nodes between two kept
nodes, collapse it into a single new edge, summing distance_km and
elevation_drop_m along the real path, and ANDing verified_continuous
(if any real segment along the collapsed path was flagged
non-continuous -- e.g. the bétoire karst stretch -- the collapsed edge
correctly inherits that, rather than silently losing the flag).
"""
from typing import List, Optional, Set, Tuple
import numpy as np
import pandas as pd
def select_subgraph_nodes(
nodes_df: pd.DataFrame, max_nodes: int = 100,
include_confluences: bool = True, max_braid_nodes: Optional[int] = None,
) -> Set[str]:
"""
Real gauges first (never dropped, whatever the budget), then real
confluences (small in number system-wide -- 28 across both basins
in the last real reach-graph build -- so "include all" is a
reasonable default, not likely to blow the budget on its own), then
braid (split/rejoin) nodes filling whatever budget remains.
Returns the set of station_codes to keep.
"""
gauge_codes = set(nodes_df.loc[nodes_df["is_gauged"], "station_code"])
if len(gauge_codes) > max_nodes:
raise ValueError(f"{len(gauge_codes)} real gauge(s) alone exceed max_nodes={max_nodes} -- "
f"raise max_nodes, gauges are never dropped.")
selected = set(gauge_codes)
remaining_budget = max_nodes - len(selected)
if include_confluences and remaining_budget > 0 and "is_confluence" in nodes_df.columns:
confluence_codes = nodes_df.loc[nodes_df["is_confluence"], "station_code"].tolist()
take = confluence_codes[:remaining_budget]
selected.update(take)
remaining_budget -= len(take)
if remaining_budget > 0 and "is_split_point" in nodes_df.columns and "is_rejoin_point" in nodes_df.columns:
braid_codes = nodes_df.loc[
nodes_df["is_split_point"] | nodes_df["is_rejoin_point"], "station_code"
].tolist()
n_braid_take = min(remaining_budget, max_braid_nodes) if max_braid_nodes else remaining_budget
selected.update(braid_codes[:n_braid_take])
return selected
def _build_adjacency(edges_df: pd.DataFrame) -> dict:
"""station_code -> list of (neighbor_code, distance_km, elevation_drop_m, verified_continuous, width_m)"""
adj = {}
has_width = "width_m" in edges_df.columns
for _, e in edges_df.iterrows():
width = e["width_m"] if has_width else np.nan
adj.setdefault(e["source"], []).append(
(e["target"], e["distance_km"], e["elevation_drop_m"], bool(e.get("verified_continuous", True)), width)
)
return adj
def _collapse_path_from(
start: str, adj: dict, selected: Set[str], max_hops: int = 500,
) -> List[Tuple[str, float, float, bool, float]]:
"""
Walks forward from `start` along real edges until hitting the next
selected node (or a dead end), summing distance/elevation along the
way. Returns a list of (reached_selected_node, total_distance_km,
total_elevation_drop_m, all_verified_continuous, avg_width_m) -- a
list, not a single result, because a node can have multiple
outgoing edges (a real split point), each leading to a different
eventual selected node.
Width is handled differently from distance/elevation deliberately:
it is NOT additive -- a collapsed edge spanning several real
tronçons has a genuine distance-weighted AVERAGE width, not a sum
(summing would make a long collapsed edge look absurdly wide).
Segments with no real width data (width_m is NaN -- see
compute_edge_width.py, which leaves it NaN rather than fabricating
a value when no real BD TOPO surface polygon matched) are excluded
from the average, weighted by their own real distance; the
collapsed edge's width only falls back to NaN if NONE of its real
segments had width data at all, the same "skip missing, don't
fabricate" principle used in compute_edge_width.py's own
aggregation.
"""
results = []
# (node, dist, elev, cont, hops, width_weighted_sum, width_weight_total)
stack = [(start, 0.0, 0.0, True, 0, 0.0, 0.0)]
visited_this_walk = set()
while stack:
node, dist, elev, cont, hops, w_sum, w_weight = stack.pop()
if hops > max_hops:
continue # defensive -- avoid infinite loops on a real but unexpected cycle
for neighbor, seg_dist, seg_elev, seg_cont, seg_width in adj.get(node, []):
new_dist, new_elev, new_cont = dist + seg_dist, elev + seg_elev, cont and seg_cont
if not np.isnan(seg_width):
new_w_sum, new_w_weight = w_sum + seg_width * seg_dist, w_weight + seg_dist
else:
new_w_sum, new_w_weight = w_sum, w_weight
if neighbor in selected:
avg_width = new_w_sum / new_w_weight if new_w_weight > 0 else float("nan")
results.append((neighbor, new_dist, new_elev, new_cont, avg_width))
elif (neighbor, hops + 1) not in visited_this_walk:
visited_this_walk.add((neighbor, hops + 1))
stack.append((neighbor, new_dist, new_elev, new_cont, hops + 1, new_w_sum, new_w_weight))
return results
def build_collapsed_subgraph(
nodes_df: pd.DataFrame, edges_df: pd.DataFrame, max_nodes: int = 100,
include_confluences: bool = True, max_braid_nodes: Optional[int] = None,
) -> Tuple[pd.DataFrame, pd.DataFrame]:
"""
The real deliverable: a small nodes_df/edges_df pair, real
connectivity preserved via collapsed edges, ready to feed directly
into combine_basins the same way the full reach graph is.
"""
selected = select_subgraph_nodes(nodes_df, max_nodes, include_confluences, max_braid_nodes)
small_nodes_df = nodes_df[nodes_df["station_code"].isin(selected)].reset_index(drop=True)
adj = _build_adjacency(edges_df)
collapsed_rows = []
seen_pairs = set() # avoid duplicate (source, target) edges from multiple equivalent paths
for source in selected:
for target, dist, elev, cont, width in _collapse_path_from(source, adj, selected):
if source == target:
continue # a path that loops back to its own start isn't a real edge
key = (source, target)
if key in seen_pairs:
continue
seen_pairs.add(key)
collapsed_rows.append({
"source": source, "target": target,
"distance_km": dist, "elevation_drop_m": elev, "verified_continuous": cont,
"width_m": width,
})
small_edges_df = pd.DataFrame(collapsed_rows)
return small_nodes_df, small_edges_df