zhangrenchao's picture
Publish Streamflow-LSTM reproduction
186a48a verified
Raw
History Blame Contribute Delete
5.2 kB
#!/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()