Spaces:
Running on Zero
Running on Zero
File size: 21,497 Bytes
a74054f 4bb7968 a74054f 4bb7968 a74054f 4bb7968 a74054f 4bb7968 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 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 | """
Test / validation script for src/graph/build_graph.py and
src/graph/node_features.py β runs the full pipeline against real data
under --data-root and checks the result is actually sane (no NaN/Inf,
correct shapes, no accidental cross-basin edges, targets kept separate
from inputs, etc.) rather than just "it didn't crash".
Usage:
python test_build_graph.py --data-root datasets
Expected layout under DATASETS_DIR (any missing piece is skipped, same
as generate_plots.py):
station_elevations.csv
idpr.csv
ades/groundwater_levels_watershed.csv
ades/groundwater_stations.csv
safran/era5_*.nc
hydrometric/discharge_observations.csv
hydrometric/waterlevel_observations.csv
centerlines/eure_centerline.csv
centerlines/risle_centerline.csv
"""
import argparse
import sys
from pathlib import Path
import numpy as np
try:
from .graph.node_features import build_node_features
from .graph.build_graph import build_surface_edges, build_pyg_graph, build_pyg_graphs_per_basin
from .data.river_centerline import load_centerline, snap_gauges_to_centerline
from .graph.physics_losses import build_confluence_index, build_braid_index
except ImportError:
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from src.graph.node_features import build_node_features
from src.graph.build_graph import build_surface_edges, build_pyg_graph, build_pyg_graphs_per_basin
from src.data.river_centerline import load_centerline, snap_gauges_to_centerline
from src.graph.physics_losses import build_confluence_index, build_braid_index
import pandas as pd
def _to_numpy(t):
"""
Convert a torch tensor (or the plain-numpy stub used when this
script is tested without torch installed) to a numpy array.
NOTE: `hasattr(t, "data")` is NOT a safe way to distinguish "real
torch tensor" from "something else" β every real torch tensor has
a `.data` attribute (it returns a Tensor, not a numpy array), so
that check silently does nothing on a real install and lets a live
Tensor leak through. `.detach()` is a better discriminator here:
real tensors have it, plain numpy/stub objects don't.
"""
if hasattr(t, "detach"):
return t.detach().cpu().numpy()
if hasattr(t, "numpy"):
return t.numpy()
return np.asarray(t)
def _check(label: str, condition, detail: str = "") -> bool:
"""Print PASS/FAIL and return a genuine Python bool. `condition` may be
a plain bool, numpy bool_, or (if a torch boolean expression leaked
through) a tensor β bool() forces it to a real Python bool either way,
so a stray tensor can never contaminate the running `all_ok` total."""
condition = bool(condition)
status = "PASS" if condition else "FAIL"
print(f" [{status}] {label}" + (f" β {detail}" if detail and not condition else ""))
return condition
def run_checks(data_root: Path, basin_file_names=None) -> bool:
basin_file_names = basin_file_names or {0: "eure", 1: "risle"}
all_ok = True
print("=" * 60)
print("STEP 1: build_node_features")
print("=" * 60)
enriched, feat_report = build_node_features(
station_elevations_path=data_root / "station_elevations.csv",
idpr_path=data_root / "idpr.csv",
ades_path=data_root / "ades",
safran_path=data_root / "safran",
hydrometric_path=data_root / "hydrometric",
)
print(feat_report)
all_ok &= _check("base station table is non-empty", len(enriched) > 0)
all_ok &= _check("no duplicate station_code", enriched["station_code"].is_unique)
print()
print("=" * 60)
print("STEP 2: build_surface_edges")
print("=" * 60)
base_nodes, edges_df, graph_report = build_surface_edges(
enriched, centerline_dir=data_root / "centerlines", basin_file_names=basin_file_names,
)
print(graph_report)
all_ok &= _check("nodes returned == nodes fed in (no silent drops)",
len(base_nodes) == len(enriched),
f"{len(base_nodes)} vs {len(enriched)}")
all_ok &= _check("every edge endpoint exists in nodes",
set(edges_df["source"]).union(edges_df["target"]).issubset(set(base_nodes["station_code"])))
cross_basin = edges_df.merge(
base_nodes[["station_code", "basin_id"]].rename(columns={"station_code": "source", "basin_id": "src_basin"}),
on="source",
).merge(
base_nodes[["station_code", "basin_id"]].rename(columns={"station_code": "target", "basin_id": "tgt_basin"}),
on="target",
)
all_ok &= _check("no cross-basin edges (surface connectivity only, as intended)",
(cross_basin["src_basin"] == cross_basin["tgt_basin"]).all())
print()
print("=" * 60)
print("STEP 3: merge enriched features back onto graph-ordered nodes")
print("=" * 60)
full_nodes = base_nodes.merge(
enriched.drop(columns=["basin_id", "latitude", "longitude", "elevation_m"]),
on="station_code", how="left",
)
all_ok &= _check("merge preserved row count", len(full_nodes) == len(base_nodes))
print()
print("=" * 60)
print("STEP 4: build_pyg_graph (directed)")
print("=" * 60)
data = build_pyg_graph(full_nodes, edges_df)
print(data)
print("feature_names:", data.feature_names)
print("target_names:", getattr(data, "target_names", None))
x_np = _to_numpy(data.x)
all_ok &= _check("x has no NaN", not np.isnan(x_np).any())
all_ok &= _check("x has no Inf", not np.isinf(x_np).any())
all_ok &= _check("x row count matches node count", x_np.shape[0] == len(full_nodes))
all_ok &= _check("target_* columns excluded from feature_names",
not any(c.startswith("target_") for c in data.feature_names))
# Confirm standardization actually happened: for any feature that isn't
# constant (std > 0 in the original data) and isn't a __was_missing flag,
# its standardized column should have ~0 mean and ~1 std. A feature that
# silently stayed in raw units (e.g. climate_solar_Wm2 in the hundreds
# sitting next to a z-scored elevation near 0) would dominate distance-
# based GNN layers purely by scale, not by actual signal.
non_flag_idx = [i for i, name in enumerate(data.feature_names) if not name.endswith("__was_missing")]
means = x_np[:, non_flag_idx].mean(axis=0)
# ddof=1 (sample std) to match pandas' .std() default, which is what
# build_pyg_graph actually standardizes with β comparing against numpy's
# default ddof=0 (population std) would show a spurious mismatch for
# small station counts (sqrt((n-1)/n) off from 1.0), not a real bug.
stds = x_np[:, non_flag_idx].std(axis=0, ddof=1)
non_constant = stds > 1e-6
all_ok &= _check(
"standardized features have ~0 mean / ~1 std (feature normalization is active)",
bool(np.allclose(means[non_constant], 0, atol=1e-3) and np.allclose(stds[non_constant], 1, atol=1e-3)),
f"mean range [{means[non_constant].min():.3f}, {means[non_constant].max():.3f}], "
f"std range [{stds[non_constant].min():.3f}, {stds[non_constant].max():.3f}]"
if non_constant.any() else "no non-constant features to check",
)
ei_np = _to_numpy(data.edge_index)
all_ok &= _check("edge_index shape is (2, n_edges)",
ei_np.shape == (2, len(edges_df)), f"got {ei_np.shape}")
all_ok &= _check("edge_index values are valid node indices",
bool((ei_np >= 0).all()) and bool((ei_np < len(full_nodes)).all()))
if hasattr(data, "y"):
y_np = _to_numpy(data.y)
all_ok &= _check("y row count matches node count", y_np.shape[0] == len(full_nodes))
all_ok &= _check("y column count matches target_names", y_np.shape[1] == len(data.target_names))
print()
print("=" * 60)
print("STEP 5: build_pyg_graph (bidirectional)")
print("=" * 60)
data_bidir = build_pyg_graph(full_nodes, edges_df, bidirectional=True)
ei_bidir = _to_numpy(data_bidir.edge_index)
all_ok &= _check("bidirectional edge_index is exactly 2x directed",
ei_bidir.shape[1] == 2 * len(edges_df), f"got {ei_bidir.shape[1]}")
print()
print("=" * 60)
print("STEP 6: build_pyg_graphs_per_basin (Eure and Risle as SEPARATE graphs)")
print("=" * 60)
graphs = build_pyg_graphs_per_basin(full_nodes, edges_df)
all_ok &= _check("exactly one graph per basin present in nodes_df",
set(graphs.keys()) == set(full_nodes["basin_id"].unique()),
f"got basins {sorted(graphs.keys())}")
all_ok &= _check("returned more than one graph (not accidentally merged)", len(graphs) > 1)
total_nodes_across_graphs = 0
for basin_id, g in graphs.items():
print(f" basin {basin_id}: {g}")
g_x = _to_numpy(g.x)
g_ei = _to_numpy(g.edge_index)
total_nodes_across_graphs += g_x.shape[0]
all_ok &= _check(f" basin {basin_id}: basin_id excluded from its own feature set",
"basin_id" not in g.feature_names)
all_ok &= _check(f" basin {basin_id}: edge_index uses local 0..n-1 indexing",
bool((g_ei >= 0).all()) and bool((g_ei < g_x.shape[0]).all()),
f"n_nodes={g_x.shape[0]}, edge_index range [{g_ei.min()}, {g_ei.max()}]")
all_ok &= _check(f" basin {basin_id}: x has no NaN", not np.isnan(g_x).any())
all_ok &= _check("node counts across per-basin graphs sum to the combined total",
total_nodes_across_graphs == len(full_nodes),
f"{total_nodes_across_graphs} vs {len(full_nodes)}")
print()
print("=" * 60)
print("STEP 7: known_losing_reaches flagging (smoke test, using first edge)")
print("=" * 60)
if len(edges_df) > 0:
first_edge = (edges_df.iloc[0]["source"], edges_df.iloc[0]["target"])
_, flagged_edges, flagged_report = build_surface_edges(
enriched, centerline_dir=data_root / "centerlines", basin_file_names=basin_file_names,
known_losing_reaches=[first_edge],
)
print(flagged_report)
flagged_row = flagged_edges[
(flagged_edges["source"] == first_edge[0]) & (flagged_edges["target"] == first_edge[1])
]
all_ok &= _check("flagged edge has verified_continuous=False",
bool((~flagged_row["verified_continuous"]).iloc[0]))
else:
print(" (skipped: no edges to test with)")
print()
print("=" * 60)
print("STEP 8: centerline accuracy sanity check (snap_distance_km)")
print("=" * 60)
# A real BD TOPO/BD TOPAGE-sourced centerline should put gauges within
# ~sub-km of the mapped line. A stale/approximate centerline (e.g. one
# still digitized from a road-map screenshot, ~4km georeferencing
# error) shows up here as a much larger mean snap distance -- this
# check exists specifically to catch "one basin got upgraded to real
# geometry and the other was silently left on the old approximate
# version" without anyone having to eyeball it.
MEAN_SNAP_DISTANCE_WARN_KM = 2.0
centerline_dir = data_root / "centerlines"
if not centerline_dir.exists():
print(" (skipped: no centerlines/ directory found)")
else:
for basin_id, file_key in basin_file_names.items():
csv_path = centerline_dir / f"{file_key}_centerline.csv"
if not csv_path.exists():
print(f" basin {basin_id} ({file_key}): no centerline file, skipped")
continue
cl = load_centerline(csv_path)
b_nodes = full_nodes[full_nodes["basin_id"] == basin_id]
if b_nodes.empty:
continue
snapped = snap_gauges_to_centerline(cl, b_nodes)
mean_dist = snapped["snap_distance_km"].mean()
max_dist = snapped["snap_distance_km"].max()
print(f" basin {basin_id} ({file_key}): {len(cl)} centerline points, "
f"mean snap_distance={mean_dist:.3f} km, max={max_dist:.3f} km")
all_ok &= _check(
f" basin {basin_id} ({file_key}): mean snap distance < {MEAN_SNAP_DISTANCE_WARN_KM} km "
f"(a much larger value usually means this basin is still on an approximate/stale centerline)",
mean_dist < MEAN_SNAP_DISTANCE_WARN_KM,
f"mean={mean_dist:.3f} km",
)
print()
print("=" * 60)
print("RESULT:", "ALL CHECKS PASSED" if all_ok else "SOME CHECKS FAILED β see above")
print("=" * 60)
return all_ok
def run_reach_graph_checks(data_root: Path, basin_file_names=None) -> bool:
"""
Validates the reach graph pipeline (scripts/build_reach_graphs.py ->
scripts/enrich_reach_graph.py -> build_pyg_graph), separately from
run_checks' coverage of the original single-chain pipeline -- the
two build nodes_df/edges_df in genuinely different ways (real
branching topology + splits/rejoins vs. a single ordered chain), so
keeping their checks in separate functions keeps each readable
rather than threading conditionals through one long function.
Specifically regression-tests the three bugs found and fixed in the
previous session: structural columns (is_gauged/is_confluence/etc.)
leaking into model features, targets attaching to non-gauge nodes,
and edge_attr breaking on the real edges schema's extra string
columns (toponym, cleabs).
Gracefully returns True (not a failure) if the reach graph hasn't
been built/enriched yet -- this is meant to add coverage once that
pipeline is in use, not force it to exist.
"""
basin_file_names = basin_file_names or {0: "eure", 1: "risle"}
all_ok = True
graph_dir = data_root / "reach_graph"
print()
print("=" * 60)
print("REACH GRAPH CHECKS")
print("=" * 60)
if not graph_dir.exists():
print(f" No {graph_dir} found -- skipping (run scripts/build_reach_graphs.py "
f"and scripts/enrich_reach_graph.py first for this coverage). Not a failure.")
return True
any_basin_found = False
for basin_id, file_key in basin_file_names.items():
enriched_path = graph_dir / f"{file_key}_nodes_enriched.csv"
edges_path = graph_dir / f"{file_key}_edges.csv"
if not enriched_path.exists() or not edges_path.exists():
print(f" basin {basin_id} ({file_key}): missing enriched nodes/edges CSV, skipped")
continue
any_basin_found = True
nodes_df = pd.read_csv(enriched_path)
edges_df = pd.read_csv(edges_path)
print()
print(f"--- basin {basin_id} ({file_key}): {len(nodes_df)} nodes, {len(edges_df)} edges ---")
structural_cols = ["is_gauged", "is_confluence", "is_split_point", "is_rejoin_point"]
for col in structural_cols:
all_ok &= _check(f" '{col}' column present on nodes_df", col in nodes_df.columns)
n_gauged = int(nodes_df["is_gauged"].sum()) if "is_gauged" in nodes_df.columns else 0
n_confluence = int(nodes_df["is_confluence"].sum()) if "is_confluence" in nodes_df.columns else 0
n_split = int(nodes_df["is_split_point"].sum()) if "is_split_point" in nodes_df.columns else 0
n_rejoin = int(nodes_df["is_rejoin_point"].sum()) if "is_rejoin_point" in nodes_df.columns else 0
print(f" {n_gauged} gauged, {n_confluence} confluences, {n_split} splits, {n_rejoin} rejoins")
data = build_pyg_graph(nodes_df, edges_df)
x_np = _to_numpy(data.x)
all_ok &= _check(" x has no NaN", not np.isnan(x_np).any())
all_ok &= _check(" x has no Inf", not np.isinf(x_np).any())
all_ok &= _check(" x row count matches node count", x_np.shape[0] == len(nodes_df),
f"{x_np.shape[0]} vs {len(nodes_df)}")
# REGRESSION: structural columns must never leak into model features
leaked = [f for f in data.feature_names
if f in ("is_gauged", "is_confluence", "is_split_point", "is_rejoin_point",
"snap_distance_km", "braid_id")]
all_ok &= _check(" no structural columns leaked into feature_names", not leaked,
f"leaked: {leaked}")
# REGRESSION: structural columns still accessible as their own Data attributes
for col in structural_cols:
all_ok &= _check(f" data.{col} attribute present", hasattr(data, col))
# REGRESSION: targets only ever attach to gauged nodes, never confluences/virtual nodes
if hasattr(data, "y") and hasattr(data, "is_gauged"):
y_np = _to_numpy(data.y)
is_gauged_np = _to_numpy(data.is_gauged).astype(bool)
has_real_target = ~np.isnan(y_np).all(axis=1)
mislabeled = has_real_target & ~is_gauged_np
all_ok &= _check(" no target values on non-gauged nodes", not mislabeled.any(),
f"{int(mislabeled.sum())} mislabeled node(s)")
all_ok &= _check(" target coverage matches gauge count or less",
int(has_real_target.sum()) <= n_gauged,
f"{int(has_real_target.sum())} with targets vs {n_gauged} gauged")
# REGRESSION: edge_attr stays 3 columns despite extra edge metadata (toponym, cleabs)
edge_attr_np = _to_numpy(data.edge_attr)
all_ok &= _check(" edge_attr has exactly 3 columns despite extra edge metadata",
edge_attr_np.shape[1] == 3, f"got {edge_attr_np.shape[1]} columns")
# Physics loss index builders (physics_losses.py) sanity-checked against this same data
conf_idx = build_confluence_index(nodes_df, edges_df)
braid_idx = build_braid_index(nodes_df)
all_ok &= _check(" confluence_index count matches is_confluence sum",
len(conf_idx) == n_confluence, f"{len(conf_idx)} vs {n_confluence}")
if conf_idx:
max_idx = max(max(upstream) for _, upstream in conf_idx)
all_ok &= _check(" confluence_index indices within node bounds",
0 <= max_idx < len(nodes_df))
min_upstream = min(len(upstream) for _, upstream in conf_idx)
all_ok &= _check(" every confluence has >= 2 upstream branches", min_upstream >= 2)
if braid_idx:
max_braid_idx = max(max(pair) for pair in braid_idx)
all_ok &= _check(" braid_index indices within node bounds",
0 <= max_braid_idx < len(nodes_df))
all_ok &= _check(" braid_index count matches is_rejoin_point sum",
len(braid_idx) == n_rejoin, f"{len(braid_idx)} vs {n_rejoin}")
n_climate = int(nodes_df["climate_precip_mm"].notna().sum()) if "climate_precip_mm" in nodes_df.columns else 0
n_catchment = int(nodes_df["catchment_area_km2"].notna().sum()) if "catchment_area_km2" in nodes_df.columns else 0
n_idpr = int(nodes_df["idpr_value"].notna().sum()) if "idpr_value" in nodes_df.columns else 0
print(f" coverage: climate {n_climate}/{len(nodes_df)}, idpr {n_idpr}/{len(nodes_df)}, "
f"catchment_area (Hub'Eau, gauges only) {n_catchment}/{len(nodes_df)}")
all_ok &= _check(" idpr present on the enriched table (was missing in a real run once -- "
"check the idpr_path passed to enrich_reach_graph.py if this fails)",
"idpr_value" in nodes_df.columns)
if "cumulative_catchment_area_km2" in nodes_df.columns:
n_cumulative = int(nodes_df["cumulative_catchment_area_km2"].notna().sum())
print(f" cumulative_catchment_area_km2 (BD TOPO, graph-wide) coverage: "
f"{n_cumulative}/{len(nodes_df)}")
all_ok &= _check(" cumulative_catchment_area_km2 covers meaningfully more than "
"the Hub'Eau-only catchment_area_km2",
n_cumulative > n_catchment,
f"{n_cumulative} vs {n_catchment} -- if this fails, "
f"scripts/compute_cumulative_catchment.py likely needs a re-run")
else:
print(" cumulative_catchment_area_km2 not present -- run "
"scripts/compute_cumulative_catchment.py for graph-wide catchment coverage "
"(known gap otherwise: only the ~8 real gauges Hub'Eau publishes it for).")
if not any_basin_found:
print(" No enriched reach graph files found for any basin -- run "
"scripts/build_reach_graphs.py then scripts/enrich_reach_graph.py first.")
return True
print()
print("=" * 60)
print("REACH GRAPH RESULT:", "ALL CHECKS PASSED" if all_ok else "SOME CHECKS FAILED β see above")
print("=" * 60)
return all_ok
def main() -> None:
parser = argparse.ArgumentParser(description="Test build_graph.py / node_features.py against real data")
parser.add_argument("--data-root", type=Path, default=Path("datasets"))
args = parser.parse_args()
ok_original = run_checks(args.data_root)
ok_reach_graph = run_reach_graph_checks(args.data_root)
sys.exit(0 if (ok_original and ok_reach_graph) else 1)
if __name__ == "__main__":
main() |