File size: 2,022 Bytes
9be39c5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Forecast fifteen 6-minute radar frames from five observations."""

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.convlstm import ConvLSTM
from train import RadarDataset, 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=True)
    model = ConvLSTM(checkpoint["model_config"]).to(device)
    model.load_state_dict(checkpoint["model"])
    model.eval()
    loader = DataLoader(RadarDataset(ROOT / config["data"]["root"] / "test.npz", config), batch_size=1)
    inputs_all, targets_all, predictions_all = [], [], []
    with torch.no_grad():
        for inputs, targets in loader:
            prediction, _ = model(inputs.to(device))
            inputs_all.append(inputs.numpy())
            targets_all.append(targets.numpy())
            predictions_all.append(prediction.cpu().numpy())
    output = ROOT / config["paths"]["inference_dir"] / "predictions.npz"
    output.parent.mkdir(parents=True, exist_ok=True)
    np.savez_compressed(output, inputs=np.concatenate(inputs_all), targets=np.concatenate(targets_all),
                        predictions=np.concatenate(predictions_all),
                        input_lead_minutes=np.arange(-24, 1, int(config["data"]["interval_minutes"]), dtype=np.int64),
                        forecast_lead_minutes=np.arange(1, int(config["data"]["output_frames"]) + 1, dtype=np.int64)
                                              * int(config["data"]["interval_minutes"]),
                        normalized_value_range=np.asarray([0.0, 1.0], np.float32),
                        data_type=np.asarray("normalized_radar_echo_grayscale"))
    print(f"predictions={output.relative_to(ROOT)}")


if __name__ == "__main__":
    main()