River_Network / scripts /tests /train_climate_anomaly_model.py
ageraustine's picture
Upload folder using huggingface_hub (part 2)
f2046b4 verified
Raw History Blame Contribute Delete
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()