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