Spaces:
Running on Zero
Running on Zero
| """ | |
| 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() |