"""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()