File size: 9,408 Bytes
a74054f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
"""
Build the reach-based river network graph for both basins from real BD TOPO
data, and report on every step so problems are visible immediately instead
of surfacing later as a silent wrong answer downstream.

Usage:
    python -m scripts.build_reach_graphs --data-root datasets
"""
import argparse
import json
import sys
from pathlib import Path

import pandas as pd

try:
    from src.graph.build_reach_graph import (
        check_flow_direction_coverage, load_troncons_for_basin, build_node_link_digraph,
        best_component_for_stations, snap_gauges_to_reach_graph, insert_virtual_nodes,
        build_reach_graph_tables, summarize_confluences_by_tributary, TRIBUTARY_NAMES, normalize_toponym,
    )
    from src.data.loaders.station_elevations import StationElevationsLoader
    from src.data.river_graph import assign_basin_id
except ImportError:
    sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
    from src.graph.build_reach_graph import (
        check_flow_direction_coverage, load_troncons_for_basin, build_node_link_digraph,
        best_component_for_stations, snap_gauges_to_reach_graph, insert_virtual_nodes,
        build_reach_graph_tables, summarize_confluences_by_tributary, TRIBUTARY_NAMES, normalize_toponym,
    )
    from src.data.loaders.station_elevations import StationElevationsLoader
    from src.data.river_graph import assign_basin_id

BASIN_NAMES = {0: "La Eure", 1: "La Risle"}


def run_for_basin(
    troncon_geojson: dict,
    stations_df: pd.DataFrame,
    basin_id: int,
    virtual_node_spacing_km: float,
    anchor_radius_km: float,
    component_search_radius_km: float,
) -> "tuple[pd.DataFrame, pd.DataFrame] | None":
    name = BASIN_NAMES[basin_id]
    print("=" * 70)
    print(f"{name}")
    print("=" * 70)

    basin_stations = stations_df[stations_df["basin_id"] == basin_id].copy()

    matched = load_troncons_for_basin(
        troncon_geojson, basin_id=basin_id,
        anchor_stations=basin_stations, max_distance_from_anchor_km=anchor_radius_km,
    )
    print(f"Matched {len(matched)} tronçons by name AND within {anchor_radius_km} km of a "
          f"real gauge (name-only matching pulled in geographically unrelated tronçons "
          f"in earlier runs -- see build_reach_graph.py's load_troncons_for_basin docstring)")
    if not matched:
        print("  Nothing matched -- try a larger --anchor-radius-km, or check "
              "TRIBUTARY_NAMES against this file's actual toponym values.")
        return None

    print()
    print("Flow direction field coverage (sens_de_l_ecoulement), all matched tronçons:")
    coverage = check_flow_direction_coverage(matched)
    for value, count in sorted(coverage.items(), key=lambda x: -x[1]):
        print(f"  {value!r}: {count}")
    print("  If this doesn't show recognizable direction values, "
          "build_node_link_digraph's sens_downstream_values/sens_upstream_values "
          "need updating before its elevation fallback is the only thing working.")

    print()
    G, report = build_node_link_digraph(matched)
    print(report)

    n_known_tributaries = len(TRIBUTARY_NAMES.get(basin_id, [])) - 1  # exclude the main river itself
    print(f"  ({n_known_tributaries} known named tributaries for this basin -- confluence count "
          f"should be in this neighborhood, not far beyond it, if a tributary joins the main "
          f"river roughly once)")
    tributary_summary = summarize_confluences_by_tributary(G)
    if not tributary_summary.empty:
        counts = tributary_summary["joining_river"].value_counts()
        print(f"  Confluences by joining river (post name-normalization): {dict(counts)}")
        if len(counts) > n_known_tributaries + 3:  # some slack for real sub-tributary branching
            print(f"  NOTE: more distinct joining-river names ({len(counts)}) than known "
                  f"tributaries ({n_known_tributaries}) -- worth checking whether some of "
                  f"these are actually the same river under yet another name variant "
                  f"normalize_toponym doesn't catch, or genuinely unlisted tributaries.")

    if report.disconnected_components > 1:
        print(f"  {report.disconnected_components} component(s) found. Selecting by which "
              f"one actually contains our real gauges, not by raw size -- the two aren't "
              f"guaranteed to agree, and in earlier runs they didn't.")
    G, gauge_counts = best_component_for_stations(G, basin_stations, search_radius_km=component_search_radius_km)
    if len(gauge_counts) > 1:
        sorted_counts = sorted(gauge_counts.values(), reverse=True)
        print(f"  Gauges captured per component: {sorted_counts[:5]}{'...' if len(sorted_counts) > 5 else ''}")
        if sorted_counts[0] < len(basin_stations):
            print(f"  WARNING: the winning component only captures {sorted_counts[0]} of "
                  f"{len(basin_stations)} gauges -- this basin's real network may still be "
                  f"fragmented across more than one component even after the anchor filter.")

    print()
    print(f"Snapping {len(basin_stations)} real gauge(s) onto the reach graph...")
    snapped = snap_gauges_to_reach_graph(G, basin_stations)
    print(snapped[["station_code", "reach_u", "reach_v", "reach_fraction", "snap_distance_km"]]
          .sort_values("snap_distance_km", ascending=False).to_string(index=False))
    unsnapped = snapped["reach_u"].isna().sum()
    if unsnapped:
        print(f"  WARNING: {unsnapped} station(s) got no nearest edge at all -- "
              f"they're likely outside this basin's matched tronçon set.")
    large_snap = (snapped["snap_distance_km"] > 2.0).sum()
    if large_snap:
        print(f"  NOTE: {large_snap} station(s) snapped more than 2 km from any "
              f"reach -- worth checking those aren't actually on an unmatched tributary.")

    print()
    G_virtual, n_inserted = insert_virtual_nodes(G, spacing_km=virtual_node_spacing_km)
    print(f"Inserted {n_inserted} virtual node(s) at ~{virtual_node_spacing_km} km spacing "
          f"({G.number_of_nodes()} real nodes -> {G_virtual.number_of_nodes()} total)")

    print()
    nodes_df, edges_df = build_reach_graph_tables(G_virtual, snapped, basin_id=basin_id)
    n_confluences = int(nodes_df["is_confluence"].sum())
    print(f"Final table: {len(nodes_df)} nodes ({nodes_df['is_gauged'].sum()} gauged, "
          f"{n_confluences} confluences, {len(nodes_df) - nodes_df['is_gauged'].sum() - n_confluences} "
          f"virtual/other), {len(edges_df)} edges")

    return nodes_df, edges_df


def main() -> None:
    parser = argparse.ArgumentParser(description="Build reach-based river network graphs")
    parser.add_argument("--data-root", type=Path, default=Path("datasets"))
    parser.add_argument("--virtual-node-spacing-km", type=float, default=5.0)
    parser.add_argument("--anchor-radius-km", type=float, default=20.0,
                         help="Max distance from a real gauge for a name-matched tronçon to be kept")
    parser.add_argument("--component-search-radius-km", type=float, default=5.0,
                         help="How close a gauge must be to a component to count as 'in' it")
    parser.add_argument("--output-dir", type=Path, default=None,
                         help="Defaults to <data-root>/reach_graph")
    args = parser.parse_args()

    troncon_path = args.data_root / "bdtopo_hydro" / "troncon_hydrographique.geojson"
    if not troncon_path.exists():
        print(f"Missing {troncon_path} -- run scripts/download_bdtopo_hydro.py first.")
        sys.exit(1)
    troncon_geojson = json.loads(troncon_path.read_text())

    stations_df = StationElevationsLoader(data_path=args.data_root / "station_elevations.csv").load()
    stations_df["basin_id"] = stations_df["station_code"].apply(assign_basin_id)
    unmatched = stations_df["basin_id"].isna().sum()
    if unmatched:
        print(f"NOTE: {unmatched} station(s) matched no basin prefix and are excluded entirely.\n")
    stations_df = stations_df.dropna(subset=["basin_id"])
    stations_df["basin_id"] = stations_df["basin_id"].astype(int)

    output_dir = args.output_dir or (args.data_root / "reach_graph")
    output_dir.mkdir(parents=True, exist_ok=True)

    all_nodes, all_edges = [], []
    for basin_id, file_key in [(0, "eure"), (1, "risle")]:
        result = run_for_basin(troncon_geojson, stations_df, basin_id, args.virtual_node_spacing_km,
                                args.anchor_radius_km, args.component_search_radius_km)
        print()
        if result is None:
            continue
        nodes_df, edges_df = result
        nodes_df.to_csv(output_dir / f"{file_key}_nodes.csv", index=False)
        edges_df.to_csv(output_dir / f"{file_key}_edges.csv", index=False)
        print(f"Saved {file_key}_nodes.csv / {file_key}_edges.csv to {output_dir}")
        print()
        all_nodes.append(nodes_df)
        all_edges.append(edges_df)

    if all_nodes:
        print("=" * 70)
        print("SUMMARY")
        print("=" * 70)
        combined_nodes = pd.concat(all_nodes, ignore_index=True)
        combined_edges = pd.concat(all_edges, ignore_index=True)
        print(f"Total: {len(combined_nodes)} nodes, {len(combined_edges)} edges across both basins")
        print(f"Gauged: {combined_nodes['is_gauged'].sum()}")


if __name__ == "__main__":
    main()