File size: 7,491 Bytes
f2046b4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
"""
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