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