OneScience-Group/Streamflow-LSTM
023
1#!/usr/bin/env python32import argparse3import json4import os5from pathlib import Path6import sys7 8import numpy as np9import torch10from torch.utils.data import DataLoader, TensorDataset11 12sys.path.insert(0, str(Path(__file__).resolve().parents[1]))13from model.streamflow_lstm import StreamflowLSTM, nse, write_json14 15parser = argparse.ArgumentParser(description="Train independent ensembles for all stream gauges")16parser.add_argument("--config", default="conf/config.yaml")17parser.add_argument("--paper", action="store_true", help="Use 50 units, 100 members, best 5")18parser.add_argument("--members", type=int, default=None)19args = parser.parse_args()20with open(args.config, encoding="utf-8") as handle:21 config = json.load(handle)22torch.set_num_threads(config["runtime"]["num_threads"])23rank = int(os.environ.get("RANK", 0))24world_size = int(os.environ.get("WORLD_SIZE", 1))25local_rank = int(os.environ.get("LOCAL_RANK", 0))26distributed = world_size > 127use_cuda = config["runtime"]["device"] == "auto" and torch.cuda.is_available()28if distributed:29 torch.distributed.init_process_group(backend="nccl" if use_cuda else "gloo")30if use_cuda:31 torch.cuda.set_device(local_rank)32device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu")33mode = config["paper_model"] if args.paper else {**config["training"], **config["model"]}34members = args.members or mode["ensemble_members_per_gauge"]35hidden, epochs = mode["hidden_size"], mode["epochs"]36with np.load(config["data"]["path"]) as data:37 train_x, train_y = data["train_x"], data["train_y"]38 val_x, val_y, gauges = data["val_x"], data["val_y"], data["gauges"].astype(str)39records, trained_members = {}, []40for gauge_index, gauge in enumerate(gauges):41 x_mean = train_x[gauge_index].mean((0, 1), keepdims=True)42 x_std = train_x[gauge_index].std((0, 1), keepdims=True) + 1e-643 y_mean, y_std = float(train_y[gauge_index].mean()), float(train_y[gauge_index].std() + 1e-6)44 x_train = torch.from_numpy((train_x[gauge_index] - x_mean) / x_std)45 y_train = torch.from_numpy((train_y[gauge_index] - y_mean) / y_std)46 x_val = torch.from_numpy((val_x[gauge_index] - x_mean) / x_std).to(device)47 loader = DataLoader(TensorDataset(x_train, y_train), batch_size=config["training"]["batch_size"], shuffle=True)48 scores = []49 for member in range(members):50 if (gauge_index * members + member) % world_size != rank:51 continue52 seed = config["seed"] + gauge_index * 1000 + member53 torch.manual_seed(seed)54 model = StreamflowLSTM(hidden_size=hidden, dropout=config["model"]["dropout"]).to(device)55 optimizer = torch.optim.Adam(model.parameters(), lr=config["training"]["learning_rate"])56 losses = []57 model.train()58 for _ in range(epochs):59 for inputs, targets in loader:60 optimizer.zero_grad(set_to_none=True)61 loss = torch.mean((model(inputs.to(device)) - targets.to(device)) ** 2)62 loss.backward()63 optimizer.step()64 losses.append(float(loss.detach()))65 model.eval()66 with torch.no_grad():67 prediction = model(x_val).cpu().numpy() * y_std + y_mean68 score = nse(prediction, val_y[gauge_index])69 payload = {"state_dict": {key: value.detach().cpu() for key, value in model.state_dict().items()},70 "hidden_size": hidden, "dropout": config["model"]["dropout"],71 "x_mean": x_mean, "x_std": x_std, "y_mean": y_mean, "y_std": y_std,72 "gauge": gauge, "member": member, "validation_nse": score}73 trained_members.append(payload)74 scores.append({"member": member, "validation_nse": score, "final_mse": losses[-1]})75 print(f"gauge={gauge} member={member} mse={losses[-1]:.6f} val_nse={score:.4f}")76 records[gauge] = scores77if distributed:78 gathered = [None] * world_size if rank == 0 else None79 torch.distributed.gather_object((trained_members, records), gathered, dst=0)80 if rank == 0:81 trained_members = [item for members_and_records in gathered for item in members_and_records[0]]82 records = {gauge: [] for gauge in gauges}83 for _, rank_records in gathered:84 for gauge, values in rank_records.items():85 records[gauge].extend(values)86if rank == 0:87 trained_members.sort(key=lambda item: (item["gauge"], item["member"]))88 for scores in records.values():89 scores.sort(key=lambda item: item["validation_nse"], reverse=True)90 checkpoint = Path(config["paths"]["checkpoint"])91 checkpoint.parent.mkdir(parents=True, exist_ok=True)92 torch.save({"format_version": config["data"]["format_version"], "gauges": gauges.tolist(),93 "members_per_gauge": members, "paper_mode": args.paper,94 "members": trained_members}, checkpoint)95 write_json(config["paths"]["training_metrics"], {"paper_mode": args.paper, "hidden_size": hidden,96 "epochs": epochs, "members_per_gauge": members, "world_size": world_size, "gauges": records})97 print(f"wrote {checkpoint}: {len(trained_members)} members")98if distributed:99 torch.distributed.destroy_process_group()100 