Download scripts/inference.py from OneScience-Group/PPNN: direct link, hf CLI and curl.
- Browser
- Download file 2.72 kB
-
https://huggingface.co/OneScience-Group/PPNN/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/PPNN/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/PPNN/resolve/main/scripts/inference.py
2.72 kB
| """Run all PPNN replicas and save the complete verification payload.""" | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import yaml | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| from model.ppnn import FORMAT_VERSION, PPNN, ensemble_features | |
| def main(): | |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) | |
| data = np.load(ROOT / config["data"]["file"]) | |
| checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location="cpu", weights_only=False) | |
| if str(data["format_version"]) != FORMAT_VERSION or checkpoint.get("format_version") != FORMAT_VERSION: | |
| raise ValueError("data/checkpoint format version mismatch") | |
| if data["ensemble"].shape[1:] != (50, 18) or len(data["station_id"]) != 537 or int(data["lead_hours"]) != 48: | |
| raise ValueError("inference requires 50 members, 18 variables, 537 stations, and 48h lead") | |
| scaling = checkpoint["scaling"] | |
| continuous = ensemble_features(torch.from_numpy(data["ensemble"]), torch.from_numpy(data["auxiliary"])) | |
| continuous = (continuous - scaling["feature_mean"]) / scaling["feature_std"] | |
| station = torch.from_numpy(data["station_index"]) | |
| mus, sigmas = [], [] | |
| for state in checkpoint["model"]: | |
| model = PPNN(int(checkpoint["model_config"]["hidden_size"]), eps=float(checkpoint["model_config"]["sigma_epsilon"])) | |
| model.load_state_dict(state) | |
| model.eval() | |
| with torch.no_grad(): | |
| mu_scaled, sigma_scaled = model(continuous, station) | |
| mus.append(mu_scaled * scaling["target_std"] + scaling["target_mean"]) | |
| sigmas.append(sigma_scaled * scaling["target_std"]) | |
| # The paper averages Gaussian parameters from random networks; this is not a mixture distribution. | |
| mu = torch.stack(mus).mean(0).numpy() | |
| sigma = torch.stack(sigmas).mean(0).numpy() | |
| target = data["target"] | |
| raw_t2m = data["ensemble"][:, :, 0] | |
| if not all(np.isfinite(x).all() for x in (mu, sigma, target, raw_t2m)) or not np.all(sigma > 0): | |
| raise RuntimeError("inference output is non-finite or has non-positive sigma") | |
| output = ROOT / config["paths"]["inference"] | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| np.savez_compressed(output, mu=mu, sigma=sigma, targets=target, raw_ensemble_t2m=raw_t2m, | |
| station_index=data["station_index"], station_id=data["station_id"], date_index=data["date_index"], | |
| dates=data["dates"], lead_hours=data["lead_hours"], format_version=np.asarray(FORMAT_VERSION)) | |
| print(f"saved={output.relative_to(ROOT)} samples={len(mu)} replicas={len(mus)} raw_ensemble={raw_t2m.shape}") | |
| if __name__ == "__main__": | |
| main() | |