River_Network / scripts /build_dynamic_tensors.py
ageraustine's picture
Upload folder using huggingface_hub (part 2)
4bb7968 verified
Raw
History Blame Contribute Delete
5.4 kB
"""
Builds genuine [n_nodes, T] dynamic tensors (discharge, groundwater,
climate) for a basin's reach graph and feeds them directly into
physics_losses.py's routing_consistency_loss -- the actual integration
point dynamic_features.py exists for. Without this script, that
function has real inputs it could consume but nothing actually
producing them.
Usage:
python -m scripts.build_dynamic_tensors --data-root datasets --basin risle
"""
import argparse
from pathlib import Path
import numpy as np
import pandas as pd
try:
from src.graph.dynamic_features import (
build_discharge_timeseries, build_groundwater_timeseries,
build_climate_timeseries, assemble_dynamic_tensor,
)
from src.graph.physics_losses import build_routing_index, routing_consistency_loss
except ImportError:
import sys
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from src.graph.dynamic_features import (
build_discharge_timeseries, build_groundwater_timeseries,
build_climate_timeseries, assemble_dynamic_tensor,
)
from src.graph.physics_losses import build_routing_index, routing_consistency_loss
BASIN_FILE_NAMES = {0: "eure", 1: "risle"}
def run_for_basin(data_root: Path, basin_id: int, file_key: str, date_range, skip_climate: bool) -> None:
graph_dir = data_root / "reach_graph"
nodes_path = graph_dir / f"{file_key}_nodes_enriched.csv"
edges_path = graph_dir / f"{file_key}_edges.csv"
if not nodes_path.exists() or not edges_path.exists():
print(f"{file_key}: missing enriched nodes/edges -- run build_reach_graphs.py + "
f"enrich_reach_graph.py first. Skipping.")
return
nodes_df = pd.read_csv(nodes_path)
edges_df = pd.read_csv(edges_path)
print(f"--- {file_key}: {len(nodes_df)} nodes, {len(edges_df)} edges ---")
discharge_wide = build_discharge_timeseries(nodes_df, data_root / "hydrometric", date_range)
Q, dates = assemble_dynamic_tensor(nodes_df, discharge_wide)
n_real = int((~np.isnan(Q)).sum())
print(f"discharge tensor: {Q.shape}, {n_real}/{Q.size} real (non-NaN) values "
f"({100*n_real/Q.size:.3f}% coverage -- expect this to be tiny, only real "
f"gauges with real observations ever have a value here)")
level_wide, depth_wide = build_groundwater_timeseries(nodes_df, data_root / "ades", date_range)
level_tensor, _ = assemble_dynamic_tensor(nodes_df, level_wide)
depth_tensor, _ = assemble_dynamic_tensor(nodes_df, depth_wide)
print(f"groundwater level tensor: {level_tensor.shape}, "
f"{int((~np.isnan(level_tensor)).sum())} real values")
climate_tensors = {}
if not skip_climate and (data_root / "safran").exists():
try:
climate_dict = build_climate_timeseries(nodes_df, data_root / "safran", date_range)
for var, wide in climate_dict.items():
tensor, _ = assemble_dynamic_tensor(nodes_df, wide)
climate_tensors[var] = tensor
print(f"climate variables: {list(climate_tensors.keys())}")
except Exception as e:
print(f"climate skipped (error: {e})")
out_dir = graph_dir / "dynamic"
out_dir.mkdir(exist_ok=True)
save_kwargs = {
"discharge": Q, "groundwater_level": level_tensor, "groundwater_depth": depth_tensor,
"dates": np.array([str(d) for d in dates]), "station_codes": nodes_df["station_code"].values,
}
save_kwargs.update({f"climate_{k}": v for k, v in climate_tensors.items()})
out_path = out_dir / f"{file_key}_dynamic.npz"
np.savez(out_path, **save_kwargs)
print(f"saved to {out_path}")
# The actual integration: routing_consistency_loss needs Q + routing_index together.
routing_index = build_routing_index(nodes_df, edges_df, timestep_hours=24.0)
print(f"routing_index: {len(routing_index)} edge(s) with a usable lag at this timestep")
loss = routing_consistency_loss(Q, routing_index)
print(f"routing_consistency_loss on the REAL discharge tensor: {loss}")
if np.isnan(loss):
print(" ^ NaN. This is exactly the thing worth checking before assuming this loss "
"is usable: real Q is almost entirely NaN (only real gauges with real "
"observations have values), and if routing_consistency_loss doesn't mask "
"NaN out of its residuals before averaging, one missing value anywhere "
"poisons the entire loss to NaN. See whether this fired.")
print()
def main() -> None:
parser = argparse.ArgumentParser(description="Build dynamic tensors and wire into physics_losses.py")
parser.add_argument("--data-root", type=Path, default=Path("datasets"))
parser.add_argument("--basin", choices=["eure", "risle", "both"], default="both")
parser.add_argument("--start-date", type=str, default="2013-01-01")
parser.add_argument("--end-date", type=str, default="2026-12-31")
parser.add_argument("--skip-climate", action="store_true")
args = parser.parse_args()
date_range = (args.start_date, args.end_date)
basins = BASIN_FILE_NAMES.items() if args.basin == "both" else \
[(k, v) for k, v in BASIN_FILE_NAMES.items() if v == args.basin]
for basin_id, file_key in basins:
run_for_basin(args.data_root, basin_id, file_key, date_range, args.skip_climate)
if __name__ == "__main__":
main()