File size: 2,826 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 | #!/usr/bin/env python3
import argparse
import json
from pathlib import Path
import sys
import numpy as np
import torch
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from model.streamflow_lstm import load_member
parser = argparse.ArgumentParser(description="Generate 40-step streamflow forecasts")
parser.add_argument("--config", default="conf/config.yaml")
parser.add_argument("--paper", action="store_true")
args = parser.parse_args()
with open(args.config, encoding="utf-8") as handle:
config = json.load(handle)
device = torch.device("cuda" if config["runtime"]["device"] == "auto" and torch.cuda.is_available() else "cpu")
best_count = config["paper_model" if args.paper else "training"]["best_members"]
with np.load(config["data"]["path"]) as data:
forecast_x, target = data["forecast_x"], data["forecast_y"]
persistence, glofas = data["persistence"], data["glofas"]
gauges, lead_hours = data["gauges"].astype(str), data["lead_hours"]
prediction = np.empty_like(target)
selected = {}
checkpoint_path = Path(config["paths"]["checkpoint"])
checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False)
if checkpoint.get("format_version") != config["data"]["format_version"]:
raise ValueError(f"checkpoint format does not match {config['data']['format_version']}")
if checkpoint.get("gauges") != gauges.tolist():
raise ValueError("checkpoint gauge order does not match input data")
for gauge_index, gauge in enumerate(gauges):
payloads = [item for item in checkpoint["members"] if item["gauge"] == gauge]
payloads.sort(key=lambda item: item["validation_nse"], reverse=True)
chosen = payloads[:best_count]
if not chosen:
raise FileNotFoundError(f"no members for {gauge} in {checkpoint_path}; run scripts/train.py first")
x = forecast_x[gauge_index].reshape(-1, 28, 23)
member_predictions = []
for payload in chosen:
model, payload = load_member(payload, device)
normalized = torch.from_numpy((x - payload["x_mean"]) / payload["x_std"]).to(device)
with torch.no_grad():
values = model(normalized).cpu().numpy() * payload["y_std"] + payload["y_mean"]
member_predictions.append(np.maximum(0, values.reshape(target.shape[1:])))
prediction[gauge_index] = np.mean(member_predictions, axis=0)
selected[gauge] = np.array([item["member"] for item in chosen], dtype=np.int64)
path = Path(config["paths"]["predictions"])
path.parent.mkdir(parents=True, exist_ok=True)
np.savez_compressed(path, prediction=prediction, target=target, persistence=persistence,
glofas=glofas, gauges=gauges, lead_hours=lead_hours,
selected_json=np.array(json.dumps({key: value.tolist() for key, value in selected.items()})))
print(f"wrote {path}: {prediction.shape}")
|