Download scripts/inference.py from OneScience-Group/Streamflow-LSTM: direct link, hf CLI and curl.
- Browser
- Download file 2.83 kB
-
https://huggingface.co/OneScience-Group/Streamflow-LSTM/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/Streamflow-LSTM/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/Streamflow-LSTM/resolve/main/scripts/inference.py
2.83 kB
| #!/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}") | |