OneScience-Group/MassConservingCNN
028
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 