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