File size: 5,199 Bytes
186a48a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
#!/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()