Download scripts/inference.py from OneScience-Group/SmaAtUNet: direct link, hf CLI and curl.
- Browser
- Download file 1.69 kB
-
https://huggingface.co/OneScience-Group/SmaAtUNet/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/SmaAtUNet/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/SmaAtUNet/resolve/main/scripts/inference.py
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() | |