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()
|