| """Predict the next six precipitation maps from twelve input maps.""" |
|
|
| 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.smaatunet import SmaAtUNet |
| from train import PrecipitationDataset, 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 = SmaAtUNet(checkpoint["model_config"]).to(device) |
| model.load_state_dict(checkpoint["model"]) |
| model.eval() |
| loader = DataLoader(PrecipitationDataset(ROOT / config["data"]["root"] / "test.npz", config), batch_size=1) |
| predictions, inputs_all, targets_all, attention_all = [], [], [], [] |
| with torch.no_grad(): |
| for inputs, targets in loader: |
| prediction, attention = model(inputs.to(device), return_attention=True) |
| inputs_all.append(inputs.numpy()) |
| targets_all.append(targets.numpy()) |
| predictions.append(prediction.cpu().numpy()) |
| attention_all.append(attention[0].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), attention=np.concatenate(attention_all)) |
| print(f"predictions={output.relative_to(ROOT)}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|