SmaAtUNet / scripts /inference.py
zhangrenchao's picture
Add engineering reproduction package
8b64ae3 verified
Raw
History Blame Contribute Delete
1.69 kB
"""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()