WoFS-StormCal / scripts /inference.py
zhangrenchao's picture
Publish WoFS-StormCal engineering reproduction
fa2b79f verified
Raw
History Blame Contribute Delete
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()