MassConservingCNN / scripts /inference.py
zhangrenchao's picture
Upload folder using huggingface_hub
0fa8141 verified
Raw
History Blame Contribute Delete
2.81 kB
"""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()