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