River_Network / scripts /tests /cross_validate.py
ageraustine's picture
Upload folder using huggingface_hub (part 2)
f2046b4 verified
Raw
History Blame Contribute Delete
12.1 kB
"""
Rolling-window cross-validation for the climate anomaly model --
built specifically because single train/val/test runs proved unstable
against real data: restricting an otherwise-identical test window from
2024-2026 down to 2024-2025 (removing one year) flipped runoff_mm from
"beats climatology at 5 of 7 horizons" to "beats climatology nowhere at
all", and temp_C from winning at all 7 horizons to losing at 3 of them.
That's not noise to shrug off -- it means no single run's per-horizon
table is trustworthy on its own, and re-running with yet another date
range and reading whichever result looks best is the wrong fix.
This trains and evaluates across several genuinely independent
(non-overlapping test period) folds with an expanding training window,
then reports, per (variable, horizon), the FRACTION of folds where the
model actually beat climatology and the mean/std of the RMSE margin --
a "win" that holds up in 4 of 4 folds is real signal; a "win" that
shows up in 1 of 4 is very likely the kind of single-run artifact that
misled the earlier single-window runs.
Usage:
python -m scripts.cross_validate_climate_anomaly_model --data-root datasets
"""
import argparse
import json
from pathlib import Path
import numpy as np
import pandas as pd
import torch
try:
from src.graph.dynamic_features import build_climate_grid_timeseries
from src.graph.climate_climatology import fit_harmonic_climatology, compute_anomalies, aggregate_anomalies, REGIME_CONFIG
from scripts.tests.train_climate_anomaly_model import train_one_variable
except ImportError:
import sys
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from src.graph.dynamic_features import build_climate_grid_timeseries
from src.graph.climate_climatology import fit_harmonic_climatology, compute_anomalies, aggregate_anomalies, REGIME_CONFIG
from scripts.tests.train_climate_anomaly_model import train_one_variable
# Expanding-window folds, each with a genuinely NON-OVERLAPPING test
# period -- this is what actually lets results from different folds be
# compared as independent evidence, rather than four correlated re-reads
# of mostly the same years. Spans as much of the real 2013-2026 ERA5
# record as possible while keeping every fold's train/val/test properly
# separated in time (no leakage from a later fold's test period into an
# earlier fold's training set, since every fold's train window starts
# at the same point and only extends forward).
DEFAULT_FOLDS = [
{"train_start": "2013-01-01", "train_end": "2016-12-31", "val_end": "2018-12-31", "test_end": "2020-12-31"},
{"train_start": "2013-01-01", "train_end": "2018-12-31", "val_end": "2020-12-31", "test_end": "2022-12-31"},
{"train_start": "2013-01-01", "train_end": "2020-12-31", "val_end": "2022-12-31", "test_end": "2024-12-31"},
{"train_start": "2013-01-01", "train_end": "2022-12-31", "val_end": "2024-12-31", "test_end": "2026-12-31"},
]
# Val periods widened from 1 year to 2 years (24 months) vs the original
# design -- confirmed as a real, not hypothetical, problem against real
# data: the monthly regime (lookback=12 + max horizon=6 = 18 months
# minimum to produce even one training example) got "val examples: 0"
# in every fold under the original 1-year val windows, silently skipping
# that entire regime rather than crashing. 24 months clears the 18-month
# floor with real margin for a stable early-stopping signal, not just
# the bare minimum. Test windows remain genuinely non-overlapping across
# folds (2019-2020, 2021-2022, 2023-2024, 2025-2026), same principle as
# before, just shifted to make room for the wider val period.
def run_cross_validation(
climate_dict: dict, variables: list, folds: list, regime_names: list,
max_epochs: int, patience: int, batch_size: int, lr: float, weight_decay: float,
n_harmonics: int, min_months_covered: int, device: str,
) -> pd.DataFrame:
all_rows = []
for fold_idx, fold in enumerate(folds):
train_start = pd.Timestamp(fold["train_start"])
train_end = pd.Timestamp(fold["train_end"])
val_end = pd.Timestamp(fold["val_end"])
test_end = pd.Timestamp(fold["test_end"])
print("=" * 70)
print(f"FOLD {fold_idx + 1}/{len(folds)}: train [{train_start.date()}, {train_end.date()}], "
f"val (...{val_end.date()}], test (...{test_end.date()}]")
print("=" * 70)
for var in variables:
if var not in climate_dict:
continue
wide = climate_dict[var]
train_wide = wide[(wide.index >= train_start) & (wide.index <= train_end)]
val_wide = wide[(wide.index > train_end) & (wide.index <= val_end)]
test_wide = wide[(wide.index > val_end) & (wide.index <= test_end)]
coeffs = fit_harmonic_climatology(train_wide, n_harmonics=n_harmonics,
min_months_covered=min_months_covered)
if not coeffs:
print(f" {var}: insufficient climatology coverage this fold, skipping")
continue
train_daily = compute_anomalies(train_wide, coeffs, n_harmonics)
val_daily = compute_anomalies(val_wide, coeffs, n_harmonics)
test_daily = compute_anomalies(test_wide, coeffs, n_harmonics)
# One model per (variable, regime) per fold -- same reasoning
# as train_climate_anomaly_model.py: daily resolution is the
# wrong target past ~2 weeks, so each regime gets its own
# appropriately-aggregated series and its own native horizons.
for regime_name, regime in REGIME_CONFIG.items():
if regime_name not in regime_names:
continue
if regime["freq"] is None:
train_anomalies, val_anomalies, test_anomalies = train_daily, val_daily, test_daily
else:
train_anomalies = aggregate_anomalies(train_daily, regime["freq"])
val_anomalies = aggregate_anomalies(val_daily, regime["freq"])
test_anomalies = aggregate_anomalies(test_daily, regime["freq"])
result = train_one_variable(
f"{var}_{regime_name}", train_anomalies, val_anomalies, test_anomalies,
regime["horizons_native"], regime["lookback"],
max_epochs, patience, batch_size, lr, weight_decay, device,
)
if result["status"] != "trained":
continue
for row in result["skill_by_horizon"]:
all_rows.append({
"fold": fold_idx + 1, "variable": var, "regime": regime_name,
"horizon_native": row["horizon_days"],
"model_rmse": row["model_rmse"], "climatology_only_rmse": row["climatology_only_rmse"],
"model_beats_climatology": row["model_beats_climatology"],
"rmse_margin_pct": 100 * (row["climatology_only_rmse"] - row["model_rmse"]) / row["climatology_only_rmse"],
})
return pd.DataFrame(all_rows)
def summarize_cross_validation(results: pd.DataFrame) -> pd.DataFrame:
"""
The actual answer this whole script exists to produce: per
(variable, regime, horizon), across all folds, what fraction
actually beat climatology, and what the margin looked like on
average -- this is what should be trusted, not any single fold's
own table. Grouped by regime as well as horizon now -- "horizon 3"
means 3 days in the daily regime, 3 weeks in the weekly regime, and
3 months in the monthly regime, so conflating them would silently
average together numbers that don't refer to the same real horizon.
"""
summary = results.groupby(["variable", "regime", "horizon_native"]).agg(
n_folds=("fold", "count"),
n_folds_beat_climatology=("model_beats_climatology", "sum"),
mean_rmse_margin_pct=("rmse_margin_pct", "mean"),
std_rmse_margin_pct=("rmse_margin_pct", "std"),
).reset_index()
summary["fraction_folds_beat"] = summary["n_folds_beat_climatology"] / summary["n_folds"]
summary["verdict"] = summary["fraction_folds_beat"].apply(
lambda f: "CONSISTENT WIN" if f == 1.0 else ("CONSISTENT LOSS" if f == 0.0 else "UNSTABLE")
)
return summary.sort_values(["variable", "regime", "horizon_native"])
def main() -> None:
parser = argparse.ArgumentParser(description="Rolling-window cross-validation for the climate anomaly model")
parser.add_argument("--data-root", type=Path, default=Path("datasets"))
parser.add_argument("--stations-path", type=Path, default=None)
parser.add_argument("--bbox-margin-deg", type=float, default=0.3)
parser.add_argument("--regimes", type=str, nargs="+", default=list(REGIME_CONFIG.keys()),
help="Which resolution regimes to cross-validate (default: all).")
parser.add_argument("--n-harmonics", type=int, default=2)
parser.add_argument("--min-months-covered", type=int, default=8)
parser.add_argument("--max-epochs", type=int, default=200)
parser.add_argument("--patience", type=int, default=10)
parser.add_argument("--batch-size", type=int, default=64)
parser.add_argument("--lr", type=float, default=1e-3)
parser.add_argument("--weight-decay", type=float, default=1e-4)
parser.add_argument("--variables", type=str, nargs="+", default=None)
parser.add_argument("--output-dir", type=Path, default=None)
args = parser.parse_args()
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Using device: {device}")
stations_path = args.stations_path or (args.data_root / "station_elevations.csv")
nodes_df = pd.read_csv(stations_path)
m = args.bbox_margin_deg
bbox = (nodes_df["longitude"].min() - m, nodes_df["latitude"].min() - m,
nodes_df["longitude"].max() + m, nodes_df["latitude"].max() + m)
print("Loading real ERA5 grid data once, reused across every fold...")
climate_dict, grid_coords = build_climate_grid_timeseries(
args.data_root / "safran", bbox, ("2013-01-01", "2026-12-31")
)
print(f"Loaded {len(climate_dict)} real variable(s) across {len(grid_coords)} real grid cell(s)\n")
variables = args.variables or list(climate_dict.keys())
output_dir = args.output_dir or (args.data_root / "climate_anomaly_models" / "cross_validation")
output_dir.mkdir(parents=True, exist_ok=True)
results = run_cross_validation(
climate_dict, variables, DEFAULT_FOLDS, args.regimes,
args.max_epochs, args.patience, args.batch_size, args.lr, args.weight_decay,
args.n_harmonics, args.min_months_covered, device,
)
results_path = output_dir / "raw_fold_results.csv"
results.to_csv(results_path, index=False)
print(f"\nSaved raw per-fold results to {results_path}")
summary = summarize_cross_validation(results)
summary_path = output_dir / "cross_validation_summary.csv"
summary.to_csv(summary_path, index=False)
print("\n" + "=" * 70)
print("CROSS-VALIDATION SUMMARY -- this is the trustworthy answer, not any single fold")
print("=" * 70)
print(summary.to_string(index=False))
print(f"\nSaved to {summary_path}")
print("\nConsistent wins (real, repeatable skill) by variable and regime:")
for var in summary["variable"].unique():
for regime in summary[summary["variable"] == var]["regime"].unique():
var_regime_summary = summary[(summary["variable"] == var) & (summary["regime"] == regime)]
consistent = var_regime_summary[var_regime_summary["verdict"] == "CONSISTENT WIN"]["horizon_native"].tolist()
unstable = var_regime_summary[var_regime_summary["verdict"] == "UNSTABLE"]["horizon_native"].tolist()
print(f" {var} / {regime}: consistent wins at {consistent}, unstable (don't trust either way) at {unstable}")
if __name__ == "__main__":
main()