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