File size: 2,808 Bytes
0fa8141 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 | """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()
|