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