File size: 2,146 Bytes
ef2ae28
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Infer next-day center-pixel wildfire 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.firecubenet import FireCubeNet
from train import WildfireDataset, 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 = WildfireDataset(ROOT / config["data"]["root"] / "test.npz", config)
    loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), shuffle=False)
    model = FireCubeNet(**checkpoint["model_config"]).to(device)
    model.load_state_dict(checkpoint["model_state_dict"])
    model.eval()
    mean = torch.from_numpy(checkpoint["channel_mean"]).to(device).view(1, 1, -1, 1, 1)
    std = torch.from_numpy(checkpoint["channel_std"]).to(device).view(1, 1, -1, 1, 1)
    probabilities = []
    with torch.no_grad():
        for inputs, _ in loader:
            probabilities.append(torch.sigmoid(model((inputs.to(device) - mean) / std)).cpu().numpy())
    probabilities = np.concatenate(probabilities).astype(np.float32)
    if probabilities.shape != dataset.data["labels"].shape or not np.isfinite(probabilities).all():
        raise FloatingPointError("invalid inference probabilities")
    output = ROOT / config["paths"]["inference_dir"] / "predictions.npz"
    output.parent.mkdir(parents=True, exist_ok=True)
    np.savez_compressed(
        output, probabilities=probabilities, labels=dataset.data["labels"],
        timestamps=dataset.data["timestamps_unix_s"], coords=dataset.data["coords"],
        format_version=np.asarray(config["data"]["format_version"]),
    )
    print(f"predictions={output.relative_to(ROOT)} shape={probabilities.shape}")


if __name__ == "__main__":
    main()