Spaces:
Sleeping
Sleeping
Download src/graph/subgraph_selection.py from ageraustine/River_Network: direct link, hf CLI and curl.
- Browser
- Download file 7.49 kB
-
https://huggingface.co/spaces/ageraustine/River_Network/resolve/main/src/graph/subgraph_selection.py
- Command line
-
hf download hf://spaces/ageraustine/River_Network/src/graph/subgraph_selection.py
-
curl -L -o subgraph_selection.py https://huggingface.co/spaces/ageraustine/River_Network/resolve/main/src/graph/subgraph_selection.py
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 |