CoolFace
Modelpublic

OneScience-Group/MassConservingCNN

sourceHugging Faceapache-2.0updated 16d agoView on Hugging Face
0likes28downloads
inference.py55 linesDownload Raw Back to scripts
1"""Run validated inference and save normalized and physical fields."""2 3from pathlib import Path4import sys5 6import numpy as np7import torch8import yaml9from torch.utils.data import DataLoader10 11 12ROOT = Path(__file__).resolve().parents[1]13sys.path.insert(0, str(ROOT))14from model.massconservingcnn import MassConservingCNN15from train import MSWDataset, device_from_config16 17 18def main():19    config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())20    device = device_from_config(config)21    checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location=device, weights_only=False)22    required = {"model", "optimizer_state_dict", "model_config", "epoch", "eta",23                "format_version", "variable_order", "normalization", "climate_mean_uh", "climate_std_uhr", "seed"}24    if not required.issubset(checkpoint):25        raise ValueError(f"incomplete checkpoint, missing {sorted(required - set(checkpoint))}")26    if checkpoint["format_version"] != config["data"]["format_version"] or checkpoint["variable_order"] != ["u", "h", "r"]:27        raise ValueError("checkpoint protocol mismatch")28    dataset = MSWDataset(ROOT / config["data"]["root"] / "validation.npz", config)29    loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), shuffle=False)30    model = MassConservingCNN(**checkpoint["model_config"]).to(device)31    model.load_state_dict(checkpoint["model"]); model.eval()32    outputs = []33    with torch.no_grad():34        for inputs, _ in loader:35            outputs.append(model(inputs.to(device)).cpu().numpy())36    predictions = np.concatenate(outputs).astype(np.float32)37    if predictions.shape != dataset.data["targets"].shape or predictions.dtype != np.float32 or not np.isfinite(predictions).all():38        raise ValueError("invalid inference output")39    means = np.asarray(checkpoint["climate_mean_uh"], dtype=np.float32)40    stds = np.asarray(checkpoint["climate_std_uhr"], dtype=np.float32)41    physical = predictions.copy()42    physical[:, :2] = predictions[:, :2] * stds[None, :2, None] + means[None, :, None]43    physical[:, 2] = predictions[:, 2] * stds[2]44    output = ROOT / config["paths"]["inference"]45    output.parent.mkdir(parents=True, exist_ok=True)46    np.savez_compressed(output, predictions=predictions, predictions_physical=physical,47                        inputs=dataset.data["inputs"], xa=dataset.data["xa"], targets=dataset.data["targets"],48                        targets_physical=dataset.data["targets_physical"], radar=dataset.data["radar"],49                        format_version=np.asarray(config["data"]["format_version"]), variable_order=np.asarray(["u", "h", "r"]))50    print(f"predictions={output.relative_to(ROOT)} shape={predictions.shape} dtype={predictions.dtype}")51 52 53if __name__ == "__main__":54    main()55