""" 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 /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()