File size: 2,657 Bytes
a00d152
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
51
52
53
54
"""Run all six storm-surge regressors on station-day test records."""

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.globalsurgeml import GlobalSurgeML
from train import CONFIGURATIONS, SurgeDataset, 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")
    dataset = SurgeDataset(ROOT / config["data"]["root"] / "test.npz", config)
    loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), shuffle=False)
    predictions = {name: [] for name in CONFIGURATIONS}
    models = {}
    for index, (name, (feature_key, method)) in enumerate(CONFIGURATIONS.items()):
        model = GlobalSurgeML(dataset.data[feature_key].shape[1], method, checkpoint["model_config"], int(config["seed"]) + index)
        model.load_state_dict(checkpoint["model"][name])
        models[name] = model.to(device).eval()
    with torch.no_grad():
        for batch in loader:
            for name, (feature_key, _) in CONFIGURATIONS.items():
                mean = torch.from_numpy(checkpoint["feature_means"][name]).to(device)
                scale = torch.from_numpy(checkpoint["feature_scales"][name]).to(device)
                predictions[name].append(models[name]((batch[feature_key].to(device) - mean) / scale).cpu().numpy())
    payload = {f"predictions_{name}": np.concatenate(values) for name, values in predictions.items()}
    for value in payload.values():
        if not np.isfinite(value).all():
            raise FloatingPointError("inference produced NaN or Inf")
    source = dataset.data
    output = ROOT / config["paths"]["inference_dir"] / "predictions.npz"
    output.parent.mkdir(parents=True, exist_ok=True)
    np.savez_compressed(output, **payload, targets_m=source["targets_m"], gtsr_m=source["gtsr_m"],
                        latitude_degrees=source["latitude_degrees"], longitude_degrees=source["longitude_degrees"],
                        timestamps_unix_s=source["timestamps_unix_s"], target_unit=np.asarray("m"),
                        format_version=np.asarray(config["data"]["format_version"]))
    print(f"predictions={output.relative_to(ROOT)} shape={source['targets_m'].shape}")


if __name__ == "__main__":
    main()