File size: 2,344 Bytes
fa2b79f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
"""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()