#!/usr/bin/env python3 import argparse import json import os from pathlib import Path import sys import numpy as np import torch from torch.utils.data import DataLoader, TensorDataset sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from model.streamflow_lstm import StreamflowLSTM, nse, write_json parser = argparse.ArgumentParser(description="Train independent ensembles for all stream gauges") parser.add_argument("--config", default="conf/config.yaml") parser.add_argument("--paper", action="store_true", help="Use 50 units, 100 members, best 5") parser.add_argument("--members", type=int, default=None) args = parser.parse_args() with open(args.config, encoding="utf-8") as handle: config = json.load(handle) torch.set_num_threads(config["runtime"]["num_threads"]) rank = int(os.environ.get("RANK", 0)) world_size = int(os.environ.get("WORLD_SIZE", 1)) local_rank = int(os.environ.get("LOCAL_RANK", 0)) distributed = world_size > 1 use_cuda = config["runtime"]["device"] == "auto" and torch.cuda.is_available() if distributed: torch.distributed.init_process_group(backend="nccl" if use_cuda else "gloo") if use_cuda: torch.cuda.set_device(local_rank) device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu") mode = config["paper_model"] if args.paper else {**config["training"], **config["model"]} members = args.members or mode["ensemble_members_per_gauge"] hidden, epochs = mode["hidden_size"], mode["epochs"] with np.load(config["data"]["path"]) as data: train_x, train_y = data["train_x"], data["train_y"] val_x, val_y, gauges = data["val_x"], data["val_y"], data["gauges"].astype(str) records, trained_members = {}, [] for gauge_index, gauge in enumerate(gauges): x_mean = train_x[gauge_index].mean((0, 1), keepdims=True) x_std = train_x[gauge_index].std((0, 1), keepdims=True) + 1e-6 y_mean, y_std = float(train_y[gauge_index].mean()), float(train_y[gauge_index].std() + 1e-6) x_train = torch.from_numpy((train_x[gauge_index] - x_mean) / x_std) y_train = torch.from_numpy((train_y[gauge_index] - y_mean) / y_std) x_val = torch.from_numpy((val_x[gauge_index] - x_mean) / x_std).to(device) loader = DataLoader(TensorDataset(x_train, y_train), batch_size=config["training"]["batch_size"], shuffle=True) scores = [] for member in range(members): if (gauge_index * members + member) % world_size != rank: continue seed = config["seed"] + gauge_index * 1000 + member torch.manual_seed(seed) model = StreamflowLSTM(hidden_size=hidden, dropout=config["model"]["dropout"]).to(device) optimizer = torch.optim.Adam(model.parameters(), lr=config["training"]["learning_rate"]) losses = [] model.train() for _ in range(epochs): for inputs, targets in loader: optimizer.zero_grad(set_to_none=True) loss = torch.mean((model(inputs.to(device)) - targets.to(device)) ** 2) loss.backward() optimizer.step() losses.append(float(loss.detach())) model.eval() with torch.no_grad(): prediction = model(x_val).cpu().numpy() * y_std + y_mean score = nse(prediction, val_y[gauge_index]) payload = {"state_dict": {key: value.detach().cpu() for key, value in model.state_dict().items()}, "hidden_size": hidden, "dropout": config["model"]["dropout"], "x_mean": x_mean, "x_std": x_std, "y_mean": y_mean, "y_std": y_std, "gauge": gauge, "member": member, "validation_nse": score} trained_members.append(payload) scores.append({"member": member, "validation_nse": score, "final_mse": losses[-1]}) print(f"gauge={gauge} member={member} mse={losses[-1]:.6f} val_nse={score:.4f}") records[gauge] = scores if distributed: gathered = [None] * world_size if rank == 0 else None torch.distributed.gather_object((trained_members, records), gathered, dst=0) if rank == 0: trained_members = [item for members_and_records in gathered for item in members_and_records[0]] records = {gauge: [] for gauge in gauges} for _, rank_records in gathered: for gauge, values in rank_records.items(): records[gauge].extend(values) if rank == 0: trained_members.sort(key=lambda item: (item["gauge"], item["member"])) for scores in records.values(): scores.sort(key=lambda item: item["validation_nse"], reverse=True) checkpoint = Path(config["paths"]["checkpoint"]) checkpoint.parent.mkdir(parents=True, exist_ok=True) torch.save({"format_version": config["data"]["format_version"], "gauges": gauges.tolist(), "members_per_gauge": members, "paper_mode": args.paper, "members": trained_members}, checkpoint) write_json(config["paths"]["training_metrics"], {"paper_mode": args.paper, "hidden_size": hidden, "epochs": epochs, "members_per_gauge": members, "world_size": world_size, "gauges": records}) print(f"wrote {checkpoint}: {len(trained_members)} members") if distributed: torch.distributed.destroy_process_group()