Download scripts/inference.py from OneScience-Group/MetNet-2: direct link, hf CLI and curl.
- Browser
- Download file 2.27 kB
-
https://huggingface.co/OneScience-Group/MetNet-2/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/MetNet-2/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/MetNet-2/resolve/main/scripts/inference.py
2.27 kB
| #!/usr/bin/env python3 | |
| import argparse | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| from model.metnet_2 import CLASS_RATES, ProceduralField, build_model, load_checkpoint, load_config | |
| parser = argparse.ArgumentParser(description="Run selected-window or streamed full-domain inference") | |
| parser.add_argument("--config", default="conf/config.yaml") | |
| parser.add_argument("--lead", type=int, default=None) | |
| parser.add_argument("--full", action="store_true") | |
| parser.add_argument("--cdf", action="store_true") | |
| args = parser.parse_args() | |
| config = load_config(args.config) | |
| torch.set_num_threads(config["runtime"]["num_threads"]) | |
| device = torch.device("cuda" if config["runtime"]["device"] == "auto" and torch.cuda.is_available() | |
| else "cpu" if config["runtime"]["device"] == "auto" else config["runtime"]["device"]) | |
| model = build_model(config).to(device) | |
| load_checkpoint(config["paths"]["checkpoint"], model) | |
| field, lead = ProceduralField(2001), args.lead or config["inference"]["lead_minutes"] | |
| if args.full: | |
| output = Path(config["paths"]["predictions"]).with_suffix(".npy") | |
| print(model.assemble_full(field, lead, output, config["data"]["window"], config["data"]["halo"], | |
| config["training"]["class_chunk"], "cdf" if args.cdf else "probability", device)) | |
| else: | |
| window = config["data"]["window"] | |
| model.eval() | |
| with torch.no_grad(): | |
| logits = model(field.window(0, 0, window, config["data"]["halo"]).unsqueeze(0).to(device), | |
| torch.tensor([lead], device=device), window)[0] | |
| probabilities = logits.softmax(0).cpu().numpy().astype(np.float32) | |
| if not np.isfinite(probabilities).all(): | |
| raise FloatingPointError("inference probabilities are not finite") | |
| output = Path(config["paths"]["predictions"]) | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| np.savez_compressed(output, probabilities=probabilities, cdf=np.cumsum(probabilities, axis=0), | |
| target=field.target_window(0, 0, window, lead).numpy(), rates=CLASS_RATES, | |
| lead_minutes=np.int32(lead), coverage=np.array(config["inference"]["coverage"]), | |
| is_complete=np.bool_(config["inference"]["is_complete"])) | |
| print(output) | |