zhangrenchao commited on
Commit
186a48a
·
verified ·
1 Parent(s): ae2d9f8

Publish Streamflow-LSTM reproduction

Browse files
.gitattributes CHANGED
@@ -1,35 +1,3 @@
1
- *.7z filter=lfs diff=lfs merge=lfs -text
2
- *.arrow filter=lfs diff=lfs merge=lfs -text
3
- *.bin filter=lfs diff=lfs merge=lfs -text
4
- *.bz2 filter=lfs diff=lfs merge=lfs -text
5
- *.ckpt filter=lfs diff=lfs merge=lfs -text
6
- *.ftz filter=lfs diff=lfs merge=lfs -text
7
- *.gz filter=lfs diff=lfs merge=lfs -text
8
- *.h5 filter=lfs diff=lfs merge=lfs -text
9
- *.joblib filter=lfs diff=lfs merge=lfs -text
10
- *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
- *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
- *.model filter=lfs diff=lfs merge=lfs -text
13
- *.msgpack filter=lfs diff=lfs merge=lfs -text
14
- *.npy filter=lfs diff=lfs merge=lfs -text
15
- *.npz filter=lfs diff=lfs merge=lfs -text
16
- *.onnx filter=lfs diff=lfs merge=lfs -text
17
- *.ot filter=lfs diff=lfs merge=lfs -text
18
- *.parquet filter=lfs diff=lfs merge=lfs -text
19
- *.pb filter=lfs diff=lfs merge=lfs -text
20
- *.pickle filter=lfs diff=lfs merge=lfs -text
21
- *.pkl filter=lfs diff=lfs merge=lfs -text
22
  *.pt filter=lfs diff=lfs merge=lfs -text
23
- *.pth filter=lfs diff=lfs merge=lfs -text
24
- *.rar filter=lfs diff=lfs merge=lfs -text
25
- *.safetensors filter=lfs diff=lfs merge=lfs -text
26
- saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
- *.tar.* filter=lfs diff=lfs merge=lfs -text
28
- *.tar filter=lfs diff=lfs merge=lfs -text
29
- *.tflite filter=lfs diff=lfs merge=lfs -text
30
- *.tgz filter=lfs diff=lfs merge=lfs -text
31
- *.wasm filter=lfs diff=lfs merge=lfs -text
32
- *.xz filter=lfs diff=lfs merge=lfs -text
33
- *.zip filter=lfs diff=lfs merge=lfs -text
34
- *.zst filter=lfs diff=lfs merge=lfs -text
35
- *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  *.pt filter=lfs diff=lfs merge=lfs -text
2
+ *.npz filter=lfs diff=lfs merge=lfs -text
3
+ *.png filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
 
 
 
 
conf/config.yaml ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "seed": 2026,
3
+ "data": {
4
+ "format_version": "streamflow-lstm-v1",
5
+ "path": "data/streamflow_fake.npz",
6
+ "gauges": ["BNDN5", "ARWN8", "TCCC1", "CARO2", "ESSC2", "NFDC1", "LABW4", "CLNK1", "TRAC2", "NFSW4"],
7
+ "input_variables": 23,
8
+ "history_steps": 28,
9
+ "step_hours": 6,
10
+ "forecast_steps": 40,
11
+ "train_samples": 48,
12
+ "validation_samples": 12,
13
+ "forecast_cases": 12
14
+ },
15
+ "model": {
16
+ "hidden_size": 16,
17
+ "num_layers": 3,
18
+ "dropout": 0.1
19
+ },
20
+ "paper_model": {
21
+ "hidden_size": 50,
22
+ "num_layers": 3,
23
+ "activations": ["relu", "relu", "tanh"],
24
+ "ensemble_members_per_gauge": 100,
25
+ "best_members": 5,
26
+ "epochs": 10
27
+ },
28
+ "training": {
29
+ "ensemble_members_per_gauge": 2,
30
+ "best_members": 2,
31
+ "epochs": 10,
32
+ "batch_size": 16,
33
+ "learning_rate": 0.001,
34
+ "loss": "MSE"
35
+ },
36
+ "runtime": {"device": "auto", "num_threads": 2},
37
+ "paths": {
38
+ "checkpoint": "result/checkpoints/streamflow_lstm.pt",
39
+ "training_metrics": "result/training/metrics.json",
40
+ "predictions": "result/output/predictions.npz",
41
+ "evaluation_metrics": "result/evaluation/metrics.json",
42
+ "comparison": "result/evaluation/comparison.png"
43
+ }
44
+ }
config.json ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model_name": "Streamflow-LSTM",
3
+ "model_type": "gauge_specific_stacked_lstm",
4
+ "architectures": ["StreamflowLSTM"],
5
+ "framework": "PyTorch",
6
+ "torch_dtype": "float32",
7
+ "transformers_version": "4.44.0",
8
+ "domain": "hydrology",
9
+ "task": "six-hourly-streamflow-forecasting",
10
+ "license": "apache-2.0",
11
+ "paper": {
12
+ "title": "Using a long short-term memory (LSTM) neural network to boost river streamflow forecasts over the western United States",
13
+ "doi": "10.5194/hess-26-5449-2022"
14
+ },
15
+ "implementation": {
16
+ "entry_point": "model/streamflow_lstm.py",
17
+ "scope": "core-method reduced-ensemble engineering reproduction",
18
+ "train_script": "scripts/train.py",
19
+ "inference_script": "scripts/inference.py",
20
+ "evaluation_script": "scripts/result.py",
21
+ "synthetic_data_script": "scripts/fake_data.py"
22
+ },
23
+ "architecture": {
24
+ "input_shape": ["B", 28, 23],
25
+ "output_shape": ["B"],
26
+ "paper_hidden_sizes": [50, 50, 50],
27
+ "activations": ["ReLU", "ReLU", "tanh"],
28
+ "dense_outputs": 1
29
+ },
30
+ "data": {
31
+ "format_version": "streamflow-lstm-v1",
32
+ "gauges": 10,
33
+ "input_shape": ["B", 28, 23],
34
+ "forecast_steps": 40,
35
+ "step_hours": 6,
36
+ "unit": "m3 s-1",
37
+ "synthetic": true
38
+ },
39
+ "checkpoint": {
40
+ "path": "result/checkpoints/streamflow_lstm.pt",
41
+ "format": "single-file multi-gauge multi-member PyTorch checkpoint",
42
+ "official_weights": false
43
+ },
44
+ "configuration_sources": [
45
+ "conf/config.yaml",
46
+ "model/streamflow_lstm.py",
47
+ "scripts/fake_data.py",
48
+ "scripts/train.py",
49
+ "scripts/inference.py",
50
+ "scripts/result.py"
51
+ ]
52
+ }
model/streamflow_lstm.py ADDED
@@ -0,0 +1,75 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ from pathlib import Path
3
+
4
+ import numpy as np
5
+ import torch
6
+ from torch import nn
7
+
8
+
9
+ def load_config(path):
10
+ with open(path, encoding="utf-8") as handle:
11
+ return json.load(handle)
12
+
13
+
14
+ def write_json(path, value):
15
+ path = Path(path)
16
+ path.parent.mkdir(parents=True, exist_ok=True)
17
+ with open(path, "w", encoding="utf-8") as handle:
18
+ json.dump(value, handle, indent=2, sort_keys=True)
19
+
20
+
21
+ class StreamflowLSTM(nn.Module):
22
+ """Paper stack with explicit inter-layer activation semantics."""
23
+
24
+ def __init__(self, input_size=23, hidden_size=50, dropout=0.1):
25
+ super().__init__()
26
+ self.layers = nn.ModuleList([
27
+ nn.LSTM(input_size, hidden_size, batch_first=True),
28
+ nn.LSTM(hidden_size, hidden_size, batch_first=True),
29
+ nn.LSTM(hidden_size, hidden_size, batch_first=True),
30
+ ])
31
+ self.dropout = nn.Dropout(dropout)
32
+ self.dense = nn.Linear(hidden_size, 1)
33
+
34
+ def forward(self, inputs):
35
+ if inputs.ndim != 3 or inputs.shape[-2:] != (28, 23):
36
+ raise ValueError(f"expected [B,28,23], got {tuple(inputs.shape)}")
37
+ values, _ = self.layers[0](inputs)
38
+ values = self.dropout(torch.relu(values))
39
+ values, _ = self.layers[1](values)
40
+ values = self.dropout(torch.relu(values))
41
+ values, _ = self.layers[2](values)
42
+ values = torch.tanh(values[:, -1])
43
+ return self.dense(values).squeeze(-1)
44
+
45
+
46
+ def nse(prediction, target):
47
+ prediction, target = np.asarray(prediction), np.asarray(target)
48
+ denominator = np.sum((target - target.mean()) ** 2)
49
+ return float(1.0 - np.sum((prediction - target) ** 2) / max(denominator, 1e-12))
50
+
51
+
52
+ def metrics(prediction, target):
53
+ prediction = np.asarray(prediction, dtype=np.float64)
54
+ target = np.asarray(target, dtype=np.float64)
55
+ correlation = float(np.corrcoef(prediction, target)[0, 1]) if prediction.size > 1 else 0.0
56
+ if not np.isfinite(correlation):
57
+ correlation = 0.0
58
+ alpha = float(prediction.std() / max(target.std(), 1e-12))
59
+ beta = float(prediction.mean() / max(target.mean(), 1e-12))
60
+ kge = float(1.0 - np.sqrt((correlation - 1) ** 2 + (alpha - 1) ** 2 + (beta - 1) ** 2))
61
+ return {
62
+ "kge": kge,
63
+ "kge_r": correlation,
64
+ "kge_alpha": alpha,
65
+ "kge_beta": beta,
66
+ "nse": nse(prediction, target),
67
+ "rmse_m3_s": float(np.sqrt(np.mean((prediction - target) ** 2))),
68
+ }
69
+
70
+
71
+ def load_member(payload, device):
72
+ model = StreamflowLSTM(hidden_size=payload["hidden_size"], dropout=payload["dropout"]).to(device)
73
+ model.load_state_dict(payload["state_dict"])
74
+ model.eval()
75
+ return model, payload
scripts/fake_data.py ADDED
@@ -0,0 +1,69 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ import argparse
3
+ import json
4
+ from pathlib import Path
5
+
6
+ import numpy as np
7
+
8
+ parser = argparse.ArgumentParser(description="Generate deterministic hydrometeorological smoke-test data")
9
+ parser.add_argument("--config", default="conf/config.yaml")
10
+ args = parser.parse_args()
11
+ with open(args.config, encoding="utf-8") as handle:
12
+ config = json.load(handle)
13
+
14
+ rng = np.random.default_rng(config["seed"])
15
+ data = config["data"]
16
+ gauges, train_n, val_n = len(data["gauges"]), data["train_samples"], data["validation_samples"]
17
+ cases, leads, history, variables = data["forecast_cases"], data["forecast_steps"], 28, 23
18
+ means = np.array([0.37, 8.24, 20.5, 3.28, 4.47, 49.0, 26.0, 0.14, 1.19, 8.55], dtype=np.float32)
19
+
20
+
21
+ def forcing(gauge, count, phase_offset=0.0):
22
+ phase = rng.uniform(0, 2 * np.pi, (count, 1)) + phase_offset
23
+ t = np.arange(history, dtype=np.float32)[None] / 4.0
24
+ x = rng.normal(0, 0.35, (count, history, variables)).astype(np.float32)
25
+ temperature = 9 + 12 * np.sin(2 * np.pi * (t / 365 + phase / (2 * np.pi))) - gauge * 0.35
26
+ precipitation = rng.gamma(1.2, 0.8, (count, history)) * (0.5 + gauge / 18)
27
+ x[:, :, 2] = temperature
28
+ x[:, :, 5] = precipitation
29
+ x[:, :, 6] = np.maximum(0, precipitation * 0.55 + rng.normal(0, 0.1, precipitation.shape))
30
+ x[:, :, 10:14] = np.clip(0.45 + np.cumsum(precipitation[:, :, None] * 0.004, axis=1), 0, 1)
31
+ response = means[gauge] + means[gauge] * (0.08 * precipitation[:, -8:].mean(1) + 0.025 * x[:, -1, 6])
32
+ response += rng.normal(0, max(float(means[gauge]) * 0.03, 0.02), count)
33
+ raw = np.maximum(0, response * (1.12 - gauge * 0.012) + rng.normal(0, max(float(means[gauge]) * 0.05, 0.03), count))
34
+ x[:, :, 20], x[:, :, 21], x[:, :, 22] = raw[:, None], (0.92 * raw)[:, None], means[gauge]
35
+ return x, np.maximum(0, response).astype(np.float32)
36
+
37
+
38
+ train_x = np.empty((gauges, train_n, history, variables), np.float32)
39
+ train_y = np.empty((gauges, train_n), np.float32)
40
+ val_x = np.empty((gauges, val_n, history, variables), np.float32)
41
+ val_y = np.empty((gauges, val_n), np.float32)
42
+ forecast_x = np.empty((gauges, cases, leads, history, variables), np.float32)
43
+ forecast_y = np.empty((gauges, cases, leads), np.float32)
44
+ persistence = np.empty_like(forecast_y)
45
+ glofas = np.empty_like(forecast_y)
46
+ for g in range(gauges):
47
+ train_x[g], train_y[g] = forcing(g, train_n)
48
+ val_x[g], val_y[g] = forcing(g, val_n, 0.3)
49
+ for case in range(cases):
50
+ base_x, base_y = forcing(g, 1, case / 7)
51
+ initial = float(base_y[0])
52
+ persistence[g, case] = initial
53
+ state = initial
54
+ for lead in range(leads):
55
+ window, target = forcing(g, 1, case / 7 + lead / 80)
56
+ # Forecast meteorology evolves, while observed flow is frozen at issue time.
57
+ window[:, :, 20] = initial
58
+ window[:, :, 21] *= 1.0 + 0.003 * lead
59
+ state = 0.75 * state + 0.25 * float(target[0])
60
+ forecast_x[g, case, lead] = window[0]
61
+ forecast_y[g, case, lead] = state
62
+ glofas[g, case, lead] = max(0, state * (1.10 - 0.002 * lead) + rng.normal(0, max(float(means[g]) * 0.06, 0.03)))
63
+
64
+ path = Path(data["path"])
65
+ path.parent.mkdir(parents=True, exist_ok=True)
66
+ np.savez_compressed(path, train_x=train_x, train_y=train_y, val_x=val_x, val_y=val_y,
67
+ forecast_x=forecast_x, forecast_y=forecast_y, persistence=persistence,
68
+ glofas=glofas, gauges=np.array(data["gauges"]), lead_hours=np.arange(1, leads + 1) * 6)
69
+ print(f"wrote {path}: train={train_x.shape}, forecast={forecast_x.shape}")
scripts/inference.py ADDED
@@ -0,0 +1,54 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ import argparse
3
+ import json
4
+ from pathlib import Path
5
+ import sys
6
+
7
+ import numpy as np
8
+ import torch
9
+
10
+ sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
11
+ from model.streamflow_lstm import load_member
12
+
13
+ parser = argparse.ArgumentParser(description="Generate 40-step streamflow forecasts")
14
+ parser.add_argument("--config", default="conf/config.yaml")
15
+ parser.add_argument("--paper", action="store_true")
16
+ args = parser.parse_args()
17
+ with open(args.config, encoding="utf-8") as handle:
18
+ config = json.load(handle)
19
+ device = torch.device("cuda" if config["runtime"]["device"] == "auto" and torch.cuda.is_available() else "cpu")
20
+ best_count = config["paper_model" if args.paper else "training"]["best_members"]
21
+ with np.load(config["data"]["path"]) as data:
22
+ forecast_x, target = data["forecast_x"], data["forecast_y"]
23
+ persistence, glofas = data["persistence"], data["glofas"]
24
+ gauges, lead_hours = data["gauges"].astype(str), data["lead_hours"]
25
+ prediction = np.empty_like(target)
26
+ selected = {}
27
+ checkpoint_path = Path(config["paths"]["checkpoint"])
28
+ checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False)
29
+ if checkpoint.get("format_version") != config["data"]["format_version"]:
30
+ raise ValueError(f"checkpoint format does not match {config['data']['format_version']}")
31
+ if checkpoint.get("gauges") != gauges.tolist():
32
+ raise ValueError("checkpoint gauge order does not match input data")
33
+ for gauge_index, gauge in enumerate(gauges):
34
+ payloads = [item for item in checkpoint["members"] if item["gauge"] == gauge]
35
+ payloads.sort(key=lambda item: item["validation_nse"], reverse=True)
36
+ chosen = payloads[:best_count]
37
+ if not chosen:
38
+ raise FileNotFoundError(f"no members for {gauge} in {checkpoint_path}; run scripts/train.py first")
39
+ x = forecast_x[gauge_index].reshape(-1, 28, 23)
40
+ member_predictions = []
41
+ for payload in chosen:
42
+ model, payload = load_member(payload, device)
43
+ normalized = torch.from_numpy((x - payload["x_mean"]) / payload["x_std"]).to(device)
44
+ with torch.no_grad():
45
+ values = model(normalized).cpu().numpy() * payload["y_std"] + payload["y_mean"]
46
+ member_predictions.append(np.maximum(0, values.reshape(target.shape[1:])))
47
+ prediction[gauge_index] = np.mean(member_predictions, axis=0)
48
+ selected[gauge] = np.array([item["member"] for item in chosen], dtype=np.int64)
49
+ path = Path(config["paths"]["predictions"])
50
+ path.parent.mkdir(parents=True, exist_ok=True)
51
+ np.savez_compressed(path, prediction=prediction, target=target, persistence=persistence,
52
+ glofas=glofas, gauges=gauges, lead_hours=lead_hours,
53
+ selected_json=np.array(json.dumps({key: value.tolist() for key, value in selected.items()})))
54
+ print(f"wrote {path}: {prediction.shape}")
scripts/result.py ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ import argparse
3
+ import json
4
+ from pathlib import Path
5
+ import sys
6
+
7
+ import matplotlib.pyplot as plt
8
+ import numpy as np
9
+
10
+ sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
11
+ from model.streamflow_lstm import metrics, write_json
12
+
13
+ parser = argparse.ArgumentParser(description="Evaluate streamflow forecasts and baselines")
14
+ parser.add_argument("--config", default="conf/config.yaml")
15
+ args = parser.parse_args()
16
+ with open(args.config, encoding="utf-8") as handle:
17
+ config = json.load(handle)
18
+ with np.load(config["paths"]["predictions"]) as data:
19
+ methods = {name: data[name] for name in ("prediction", "persistence", "glofas")}
20
+ target, gauges, lead_hours = data["target"], data["gauges"].astype(str), data["lead_hours"]
21
+ lead_days = {"2_day": 7, "5_day": 19, "8_day": 31}
22
+ report = {"lead_definition": "zero-based indices 7/19/31 equal 48/120/192 hours", "gauges": {}}
23
+ for gauge_index, gauge in enumerate(gauges):
24
+ report["gauges"][gauge] = {}
25
+ for method, values in methods.items():
26
+ report["gauges"][gauge][method] = {
27
+ label: metrics(values[gauge_index, :, index], target[gauge_index, :, index])
28
+ for label, index in lead_days.items()
29
+ }
30
+ numbers = [value for gauge in report["gauges"].values() for method in gauge.values()
31
+ for lead in method.values() for value in lead.values()]
32
+ if not np.all(np.isfinite(numbers)):
33
+ raise FloatingPointError("evaluation metrics are not finite")
34
+ write_json(config["paths"]["evaluation_metrics"], report)
35
+
36
+ colors = {"prediction": "#146c94", "persistence": "#d17a22", "glofas": "#6a8e3a"}
37
+ figure, axes = plt.subplots(2, 1, figsize=(11, 8), constrained_layout=True)
38
+ for method, values in methods.items():
39
+ rmse = [np.mean([metrics(values[g, :, lead], target[g, :, lead])["rmse_m3_s"] for g in range(len(gauges))])
40
+ for lead in range(len(lead_hours))]
41
+ axes[0].plot(lead_hours / 24, rmse, label=method, color=colors[method], linewidth=2)
42
+ axes[0].set(title="Mean gauge RMSE across forecast lead", xlabel="Lead (days)", ylabel="RMSE (m3 s-1)")
43
+ axes[0].grid(alpha=0.25)
44
+ axes[0].legend()
45
+ x = np.arange(len(gauges))
46
+ width = 0.25
47
+ for offset, (method, values) in enumerate(methods.items()):
48
+ kge = [metrics(values[g, :, 19], target[g, :, 19])["kge"] for g in range(len(gauges))]
49
+ axes[1].bar(x + (offset - 1) * width, kge, width, label=method, color=colors[method])
50
+ axes[1].axhline(1 - np.sqrt(2), color="black", linestyle="--", linewidth=1, label="KGE skill threshold")
51
+ axes[1].set(title="Gauge KGE at 5-day lead", xlabel="Gauge", ylabel="KGE", xticks=x, xticklabels=gauges)
52
+ axes[1].legend(ncol=4, fontsize=8)
53
+ comparison = Path(config["paths"]["comparison"])
54
+ comparison.parent.mkdir(parents=True, exist_ok=True)
55
+ figure.savefig(comparison, dpi=150)
56
+ plt.close(figure)
57
+ print(config["paths"]["evaluation_metrics"], comparison)
scripts/train.py ADDED
@@ -0,0 +1,99 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ import argparse
3
+ import json
4
+ import os
5
+ from pathlib import Path
6
+ import sys
7
+
8
+ import numpy as np
9
+ import torch
10
+ from torch.utils.data import DataLoader, TensorDataset
11
+
12
+ sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
13
+ from model.streamflow_lstm import StreamflowLSTM, nse, write_json
14
+
15
+ parser = argparse.ArgumentParser(description="Train independent ensembles for all stream gauges")
16
+ parser.add_argument("--config", default="conf/config.yaml")
17
+ parser.add_argument("--paper", action="store_true", help="Use 50 units, 100 members, best 5")
18
+ parser.add_argument("--members", type=int, default=None)
19
+ args = parser.parse_args()
20
+ with open(args.config, encoding="utf-8") as handle:
21
+ config = json.load(handle)
22
+ torch.set_num_threads(config["runtime"]["num_threads"])
23
+ rank = int(os.environ.get("RANK", 0))
24
+ world_size = int(os.environ.get("WORLD_SIZE", 1))
25
+ local_rank = int(os.environ.get("LOCAL_RANK", 0))
26
+ distributed = world_size > 1
27
+ use_cuda = config["runtime"]["device"] == "auto" and torch.cuda.is_available()
28
+ if distributed:
29
+ torch.distributed.init_process_group(backend="nccl" if use_cuda else "gloo")
30
+ if use_cuda:
31
+ torch.cuda.set_device(local_rank)
32
+ device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu")
33
+ mode = config["paper_model"] if args.paper else {**config["training"], **config["model"]}
34
+ members = args.members or mode["ensemble_members_per_gauge"]
35
+ hidden, epochs = mode["hidden_size"], mode["epochs"]
36
+ with 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)
39
+ records, trained_members = {}, []
40
+ for 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-6
43
+ 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
+ continue
52
+ seed = config["seed"] + gauge_index * 1000 + member
53
+ 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_mean
68
+ 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] = scores
77
+ if distributed:
78
+ gathered = [None] * world_size if rank == 0 else None
79
+ 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)
86
+ if 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")
98
+ if distributed:
99
+ torch.distributed.destroy_process_group()
weight/.gitkeep ADDED
File without changes