Download scripts/inference.py from OneScience-Group/WoFS-StormCal: direct link, hf CLI and curl.
- Browser
- Download file 2.34 kB
-
https://huggingface.co/OneScience-Group/WoFS-StormCal/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/WoFS-StormCal/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/WoFS-StormCal/resolve/main/scripts/inference.py
2.34 kB
| """Restore a checkpoint and infer calibrated tornado, hail, and wind probabilities.""" | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import yaml | |
| from torch.utils.data import DataLoader | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| from model.wofsstormcal import WoFSStormCal | |
| from train import HazardDataset, device_from_config | |
| def main(): | |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) | |
| device = device_from_config(config) | |
| checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location=device, weights_only=False) | |
| if checkpoint["format_version"] != config["data"]["format_version"]: | |
| raise ValueError("checkpoint and data format versions differ") | |
| model = WoFSStormCal(int(config["model"]["calibration_points"])).to(device) | |
| model.load_state_dict(checkpoint["model"]); model.eval() | |
| dataset = HazardDataset(ROOT / config["data"]["root"] / "test.npz", config) | |
| loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), shuffle=False) | |
| predictions = [] | |
| with torch.no_grad(): | |
| for features, _, lead_group in loader: | |
| prediction = model(features.to(device), lead_group.to(device)) | |
| if prediction.shape != (len(features), 3): | |
| raise RuntimeError("model output must have shape [N,3]") | |
| predictions.append(prediction.cpu().numpy()) | |
| predictions = np.concatenate(predictions) | |
| if predictions.shape != (len(dataset), 3) or not np.isfinite(predictions).all(): | |
| raise FloatingPointError("inference output is invalid") | |
| source = dataset.data | |
| output = ROOT / config["paths"]["inference_dir"] / "predictions.npz" | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| np.savez_compressed(output, probabilities=predictions, targets=source["targets"], | |
| lead_group=source["lead_group"], lead_start_minutes=source["lead_start_minutes"], | |
| lead_end_minutes=source["lead_end_minutes"], hazards=source["hazards"], | |
| lead_group_names=source["lead_group_names"], format_version=source["format_version"]) | |
| print(f"predictions={output.relative_to(ROOT)} shape={predictions.shape} range=({predictions.min():.3f},{predictions.max():.3f})") | |
| if __name__ == "__main__": | |
| main() | |