Spaces:
Running on Zero
Running on Zero
Download scripts/tests/train_climate_anomaly_model.py from ageraustine/River_Network: direct link, hf CLI and curl.
- Browser
- Download file 15.3 kB
-
https://huggingface.co/spaces/ageraustine/River_Network/resolve/main/scripts/tests/train_climate_anomaly_model.py
- Command line
-
hf download hf://spaces/ageraustine/River_Network/scripts/tests/train_climate_anomaly_model.py
-
curl -L -o train_climate_anomaly_model.py https://huggingface.co/spaces/ageraustine/River_Network/resolve/main/scripts/tests/train_climate_anomaly_model.py
15.3 kB
| """ | |
| Trains one ClimateAnomalyModel per climate variable on real ERA5 grid- | |
| cell anomaly data, and evaluates each against the one baseline that | |
| actually matters: predicting a zero anomaly (i.e., trusting climatology | |
| alone). A model that can't beat zero-anomaly at a given horizon isn't | |
| adding real value at that horizon and shouldn't be trusted there. | |
| Usage: | |
| python -m scripts.train_climate_anomaly_model --data-root datasets \ | |
| --train-end 2023-12-31 --test-end 2026-12-31 | |
| """ | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import numpy as np | |
| import pandas as pd | |
| import torch | |
| import torch.nn as nn | |
| from torch.utils.data import TensorDataset, DataLoader | |
| 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 src.models.climate_anomaly_model import prepare_training_windows, ClimateAnomalyModel, gaussian_nll_loss | |
| 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 src.models.climate_anomaly_model import prepare_training_windows, ClimateAnomalyModel, gaussian_nll_loss | |
| def train_one_variable( | |
| var: str, train_anomalies: pd.DataFrame, val_anomalies: pd.DataFrame, test_anomalies: pd.DataFrame, | |
| horizons: list, lookback_days: int, max_epochs: int, patience: int, batch_size: int, lr: float, | |
| weight_decay: float, device: str, | |
| ) -> dict: | |
| print(f"--- Training {var} ---") | |
| X_train, Y_train, _, _ = prepare_training_windows(train_anomalies, lookback_days, horizons) | |
| X_val, Y_val, _, _ = prepare_training_windows(val_anomalies, lookback_days, horizons) | |
| X_test, Y_test, _, _ = prepare_training_windows(test_anomalies, lookback_days, horizons) | |
| print(f" train examples: {len(X_train)}, val examples: {len(X_val)}, test examples: {len(X_test)}") | |
| if len(X_train) == 0 or len(X_val) == 0 or len(X_test) == 0: | |
| print(f" not enough data to train {var} -- skipping") | |
| return {"variable": var, "status": "skipped_insufficient_data"} | |
| X_train_t = torch.tensor(X_train, dtype=torch.float32) | |
| Y_train_t = torch.tensor(Y_train, dtype=torch.float32) | |
| X_val_t = torch.tensor(X_val, dtype=torch.float32).to(device) | |
| Y_val_t = torch.tensor(Y_val, dtype=torch.float32).to(device) | |
| X_test_t = torch.tensor(X_test, dtype=torch.float32).to(device) | |
| Y_test_t = torch.tensor(Y_test, dtype=torch.float32).to(device) | |
| horizons_t = torch.tensor(horizons, dtype=torch.float32).to(device) | |
| model = ClimateAnomalyModel(init_tau=30.0).to(device) | |
| optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=weight_decay) | |
| loader = DataLoader(TensorDataset(X_train_t, Y_train_t), batch_size=batch_size, shuffle=True) | |
| best_val_loss = float("inf") | |
| best_state = None | |
| epochs_without_improvement = 0 | |
| for epoch in range(max_epochs): | |
| model.train() | |
| total_loss, n_batches = 0.0, 0 | |
| for xb, yb in loader: | |
| xb, yb = xb.to(device), yb.to(device) | |
| optimizer.zero_grad() | |
| mean, log_var = model(xb, horizons_t) | |
| loss = gaussian_nll_loss(mean, log_var, yb) | |
| if torch.isnan(loss): | |
| continue # this batch had zero real targets across every example -- skip, don't corrupt gradients | |
| loss.backward() | |
| optimizer.step() | |
| total_loss += loss.item() | |
| n_batches += 1 | |
| model.eval() | |
| with torch.no_grad(): | |
| val_mean, val_log_var = model(X_val_t, horizons_t) | |
| val_loss = gaussian_nll_loss(val_mean, val_log_var, Y_val_t) | |
| val_loss_value = val_loss.item() if not torch.isnan(val_loss) else float("inf") | |
| if val_loss_value < best_val_loss: | |
| best_val_loss = val_loss_value | |
| best_state = {k: v.clone() for k, v in model.state_dict().items()} | |
| epochs_without_improvement = 0 | |
| else: | |
| epochs_without_improvement += 1 | |
| if n_batches > 0 and (epoch % max(1, max_epochs // 20) == 0 or epoch == max_epochs - 1): | |
| print(f" epoch {epoch+1}/{max_epochs}: train NLL = {total_loss / n_batches:.4f}, " | |
| f"val NLL = {val_loss_value:.4f}, tau = {torch.exp(model.log_tau).item():.1f}" | |
| f"{' (best)' if epochs_without_improvement == 0 else ''}") | |
| if epochs_without_improvement >= patience: | |
| print(f" stopped early at epoch {epoch+1} -- no validation improvement for {patience} epochs") | |
| break | |
| if best_state is not None: | |
| model.load_state_dict(best_state) | |
| print(f" best validation NLL: {best_val_loss:.4f}") | |
| model.eval() | |
| with torch.no_grad(): | |
| pred_mean_test, pred_log_var_test = model(X_test_t, horizons_t) | |
| results = [] | |
| for i, h in enumerate(horizons): | |
| target_h = Y_test_t[:, i] | |
| pred_h = pred_mean_test[:, i] | |
| mask = ~torch.isnan(target_h) | |
| n_valid = int(mask.sum().item()) | |
| if n_valid == 0: | |
| continue | |
| model_rmse = torch.sqrt(((pred_h[mask] - target_h[mask]) ** 2).mean()).item() | |
| zero_anomaly_rmse = torch.sqrt((target_h[mask] ** 2).mean()).item() # predicting anomaly=0, i.e. pure climatology | |
| mean_predicted_variance = torch.exp(pred_log_var_test[:, i][mask]).mean().item() | |
| results.append({ | |
| "horizon_days": h, "n_test_examples": n_valid, | |
| "model_rmse": model_rmse, "climatology_only_rmse": zero_anomaly_rmse, | |
| "model_beats_climatology": model_rmse < zero_anomaly_rmse, | |
| "mean_predicted_std": mean_predicted_variance ** 0.5, | |
| }) | |
| skill_df = pd.DataFrame(results) | |
| print(skill_df.to_string(index=False)) | |
| print() | |
| return { | |
| "variable": var, "status": "trained", "tau_learned": torch.exp(model.log_tau).item(), | |
| "best_val_nll": best_val_loss, "skill_by_horizon": results, "model_state": model.state_dict(), | |
| } | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description="Train the climate anomaly model on real ERA5 data") | |
| 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("--train-start", type=str, default="2013-01-01") | |
| parser.add_argument("--train-end", type=str, default="2021-12-31", | |
| help="End of the actual training period. Default leaves 2022-2023 for " | |
| "validation and 2024+ for the final held-out test, so all three " | |
| "splits are genuinely distinct periods, not overlapping.") | |
| parser.add_argument("--val-end", type=str, default="2023-12-31", | |
| help="End of the validation period, used for early stopping only -- " | |
| "never seen by the optimizer directly.") | |
| parser.add_argument("--test-end", type=str, default="2026-12-31") | |
| parser.add_argument("--regimes", type=str, nargs="+", default=list(REGIME_CONFIG.keys()), | |
| help="Which resolution regimes to train (default: all of daily/weekly/monthly). " | |
| "See REGIME_CONFIG in climate_climatology.py for each regime's resolution, " | |
| "lookback, and native-unit horizon list.") | |
| 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, | |
| help="Upper bound only -- early stopping (see --patience) will " | |
| "typically stop well before this for most variables.") | |
| parser.add_argument("--patience", type=int, default=10, | |
| help="Stop training if validation NLL hasn't improved for this many epochs.") | |
| 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, | |
| help="L2 regularization strength -- pushes the model toward smaller, " | |
| "less confident outputs by default, complementing the decay gate " | |
| "and the NLL loss's own built-in incentive to admit uncertainty.") | |
| parser.add_argument("--variables", type=str, nargs="+", default=None, | |
| help="Restrict to specific variable names (e.g. --variables precip_mm temp_C). " | |
| "Default: train on every real variable found.") | |
| parser.add_argument("--output-dir", type=Path, default=None, | |
| help="Defaults to <data-root>/climate_anomaly_models") | |
| 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(f"Loading real ERA5 grid data (bbox={bbox})...") | |
| climate_dict, grid_coords = build_climate_grid_timeseries( | |
| args.data_root / "safran", bbox, (args.train_start, args.test_end) | |
| ) | |
| print(f"Loaded {len(climate_dict)} real variable(s) across {len(grid_coords)} real grid cell(s)") | |
| print() | |
| train_start, train_end, val_end, test_end = ( | |
| pd.Timestamp(args.train_start), pd.Timestamp(args.train_end), | |
| pd.Timestamp(args.val_end), pd.Timestamp(args.test_end) | |
| ) | |
| # Confirmed as a real failure mode, not a hypothetical: overriding | |
| # --train-end and --test-end while leaving --val-end at its default | |
| # silently produced an impossible test window (date > val_end AND | |
| # date <= test_end, with val_end already past test_end) -- every | |
| # variable trained "successfully" on an empty test split with zero | |
| # examples and no error at all. Failing loudly here instead of | |
| # letting that happen silently again. | |
| if not (train_start < train_end < val_end < test_end): | |
| raise ValueError( | |
| f"Date ranges must be strictly increasing: train_start < train_end < val_end < test_end.\n" | |
| f"Got: train_start={train_start.date()}, train_end={train_end.date()}, " | |
| f"val_end={val_end.date()}, test_end={test_end.date()}\n" | |
| f"If you're overriding --train-end and/or --test-end, you likely also need to pass " | |
| f"--val-end explicitly -- it doesn't automatically move with the others." | |
| ) | |
| variables = args.variables or list(climate_dict.keys()) | |
| output_dir = args.output_dir or (args.data_root / "climate_anomaly_models") | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| all_results = {} | |
| for var in variables: | |
| if var not in climate_dict: | |
| print(f"'{var}' not found among loaded variables ({list(climate_dict.keys())}), skipping") | |
| 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)] | |
| # Climatology fit ONLY on the training period, at daily resolution | |
| # regardless of which regime consumes it below -- validation must | |
| # be genuinely held out from every stage, and daily climatology is | |
| # the single source every regime's anomaly is computed relative to | |
| # (see aggregate_anomalies's docstring for why refitting per | |
| # resolution isn't needed). | |
| coeffs = fit_harmonic_climatology(train_wide, n_harmonics=args.n_harmonics, | |
| min_months_covered=args.min_months_covered) | |
| if not coeffs: | |
| print(f"'{var}': no cell had enough real data/seasonal coverage for climatology, skipping") | |
| continue | |
| train_daily_anomalies = compute_anomalies(train_wide, coeffs, args.n_harmonics) | |
| val_daily_anomalies = compute_anomalies(val_wide, coeffs, args.n_harmonics) | |
| test_daily_anomalies = compute_anomalies(test_wide, coeffs, args.n_harmonics) | |
| # One model per (variable, regime), not one model spanning all | |
| # horizons -- daily resolution is the structurally wrong target | |
| # past ~2 weeks (see REGIME_CONFIG's docstring); each regime gets | |
| # its own appropriately-aggregated anomaly series, its own native- | |
| # unit horizon list, and its own saved model/skill table. | |
| for regime_name, regime in REGIME_CONFIG.items(): | |
| if regime_name not in args.regimes: | |
| continue | |
| if regime["freq"] is None: | |
| train_anomalies = train_daily_anomalies | |
| val_anomalies = val_daily_anomalies | |
| test_anomalies = test_daily_anomalies | |
| else: | |
| train_anomalies = aggregate_anomalies(train_daily_anomalies, regime["freq"]) | |
| val_anomalies = aggregate_anomalies(val_daily_anomalies, regime["freq"]) | |
| test_anomalies = aggregate_anomalies(test_daily_anomalies, regime["freq"]) | |
| print(f"[{var} / {regime_name}] resolution={regime['freq'] or 'daily'}, " | |
| f"lookback={regime['lookback']}, native horizons={regime['horizons_native']}") | |
| result = train_one_variable( | |
| f"{var}_{regime_name}", train_anomalies, val_anomalies, test_anomalies, | |
| regime["horizons_native"], regime["lookback"], | |
| args.max_epochs, args.patience, args.batch_size, args.lr, args.weight_decay, device, | |
| ) | |
| all_results[f"{var}_{regime_name}"] = result | |
| if result["status"] == "trained": | |
| model_path = output_dir / f"{var}_{regime_name}_anomaly_model.pt" | |
| torch.save(result["model_state"], model_path) | |
| print(f"Saved {var}/{regime_name} model to {model_path}") | |
| summary = { | |
| var: {"status": r["status"], "tau_learned": r.get("tau_learned"), | |
| "best_val_nll": r.get("best_val_nll"), "skill_by_horizon": r.get("skill_by_horizon")} | |
| for var, r in all_results.items() | |
| } | |
| summary_path = output_dir / "training_summary.json" | |
| summary_path.write_text(json.dumps(summary, indent=2, default=str)) | |
| print(f"\nSaved training summary to {summary_path}") | |
| print("\n" + "=" * 70) | |
| print("OVERALL: does the anomaly model beat pure climatology, by horizon?") | |
| print("=" * 70) | |
| for var, r in all_results.items(): | |
| if r["status"] != "trained": | |
| continue | |
| beats = [row["horizon_days"] for row in r["skill_by_horizon"] if row["model_beats_climatology"]] | |
| print(f"{var}: beats climatology at horizons {beats}") | |
| if __name__ == "__main__": | |
| main() |