Publish Streamflow-LSTM reproduction
Browse files- .gitattributes +2 -34
- conf/config.yaml +44 -0
- config.json +52 -0
- model/streamflow_lstm.py +75 -0
- scripts/fake_data.py +69 -0
- scripts/inference.py +54 -0
- scripts/result.py +57 -0
- scripts/train.py +99 -0
- weight/.gitkeep +0 -0
.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 |
-
*.
|
| 24 |
-
*.
|
| 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
|