Team Ai
Modelpublic

OneScience-Group/Streamflow-LSTM

sourceHugging Faceapache-2.0updated 23d agoView on Hugging Face
0likes23downloads
train.py100 linesDownload Raw Back to scripts
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