OneScience-Group/Streamflow-LSTM
023
1#!/usr/bin/env python32import argparse3import json4from pathlib import Path5import sys6 7import numpy as np8import torch9 10sys.path.insert(0, str(Path(__file__).resolve().parents[1]))11from model.streamflow_lstm import load_member12 13parser = argparse.ArgumentParser(description="Generate 40-step streamflow forecasts")14parser.add_argument("--config", default="conf/config.yaml")15parser.add_argument("--paper", action="store_true")16args = parser.parse_args()17with open(args.config, encoding="utf-8") as handle:18 config = json.load(handle)19device = torch.device("cuda" if config["runtime"]["device"] == "auto" and torch.cuda.is_available() else "cpu")20best_count = config["paper_model" if args.paper else "training"]["best_members"]21with 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"]25prediction = np.empty_like(target)26selected = {}27checkpoint_path = Path(config["paths"]["checkpoint"])28checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False)29if checkpoint.get("format_version") != config["data"]["format_version"]:30 raise ValueError(f"checkpoint format does not match {config['data']['format_version']}")31if checkpoint.get("gauges") != gauges.tolist():32 raise ValueError("checkpoint gauge order does not match input data")33for 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)49path = Path(config["paths"]["predictions"])50path.parent.mkdir(parents=True, exist_ok=True)51np.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()})))54print(f"wrote {path}: {prediction.shape}")55 