Download scripts/inference.py from OneScience-Group/MassConservingCNN: direct link, hf CLI and curl.
- Browser
- Download file 2.81 kB
-
https://huggingface.co/OneScience-Group/MassConservingCNN/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/MassConservingCNN/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/MassConservingCNN/resolve/main/scripts/inference.py
2.81 kB
| """Run validated inference and save normalized and physical fields.""" | |
| from pathlib import Path | |
| import sys | |
| 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.massconservingcnn import MassConservingCNN | |
| from train import MSWDataset, 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=False) | |
| required = {"model", "optimizer_state_dict", "model_config", "epoch", "eta", | |
| "format_version", "variable_order", "normalization", "climate_mean_uh", "climate_std_uhr", "seed"} | |
| if not required.issubset(checkpoint): | |
| raise ValueError(f"incomplete checkpoint, missing {sorted(required - set(checkpoint))}") | |
| if checkpoint["format_version"] != config["data"]["format_version"] or checkpoint["variable_order"] != ["u", "h", "r"]: | |
| raise ValueError("checkpoint protocol mismatch") | |
| dataset = MSWDataset(ROOT / config["data"]["root"] / "validation.npz", config) | |
| loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), shuffle=False) | |
| model = MassConservingCNN(**checkpoint["model_config"]).to(device) | |
| model.load_state_dict(checkpoint["model"]); model.eval() | |
| outputs = [] | |
| with torch.no_grad(): | |
| for inputs, _ in loader: | |
| outputs.append(model(inputs.to(device)).cpu().numpy()) | |
| predictions = np.concatenate(outputs).astype(np.float32) | |
| if predictions.shape != dataset.data["targets"].shape or predictions.dtype != np.float32 or not np.isfinite(predictions).all(): | |
| raise ValueError("invalid inference output") | |
| means = np.asarray(checkpoint["climate_mean_uh"], dtype=np.float32) | |
| stds = np.asarray(checkpoint["climate_std_uhr"], dtype=np.float32) | |
| physical = predictions.copy() | |
| physical[:, :2] = predictions[:, :2] * stds[None, :2, None] + means[None, :, None] | |
| physical[:, 2] = predictions[:, 2] * stds[2] | |
| output = ROOT / config["paths"]["inference"] | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| np.savez_compressed(output, predictions=predictions, predictions_physical=physical, | |
| inputs=dataset.data["inputs"], xa=dataset.data["xa"], targets=dataset.data["targets"], | |
| targets_physical=dataset.data["targets_physical"], radar=dataset.data["radar"], | |
| format_version=np.asarray(config["data"]["format_version"]), variable_order=np.asarray(["u", "h", "r"])) | |
| print(f"predictions={output.relative_to(ROOT)} shape={predictions.shape} dtype={predictions.dtype}") | |
| if __name__ == "__main__": | |
| main() | |