| """Run fold-held-out inference for every valid date, lead, variable, and station.""" |
|
|
| import sys |
| from pathlib import Path |
| import numpy as np |
| import torch |
| import yaml |
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(ROOT)) |
| from model.improver_aifs import ImproverAIFS |
| from fake_data import generate_chunk |
|
|
|
|
| def main(): |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) |
| meta = np.load(ROOT / config["data"]["root"] / "protocol.npz", allow_pickle=True) |
| checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location="cpu", weights_only=False) |
| fold_states = checkpoint.get("model", {}).get("fold_states", []) |
| if checkpoint["format_version"] != config["data"]["format_version"] or len(fold_states) != 2: |
| raise ValueError("checkpoint must contain two valid-time folds") |
| dates, leads, stations = 2, 241, config["data"]["station_count"] |
| raw = np.empty((dates, leads, 3, stations), np.float32) |
| truth = np.empty_like(raw); expected = np.empty_like(raw) |
| probabilities = [np.empty((dates, leads, len(v), stations), np.float32) for v in config["model"]["thresholds"]] |
| thresholds = [torch.tensor(v, dtype=torch.float32) for v in config["model"]["thresholds"]] |
| chunk = config["data"]["station_chunk_size"] |
| for fold in range(2): |
| model = ImproverAIFS(checkpoint["model_config"]); model.load_state_dict(fold_states[fold]); model.eval() |
| for start in range(0, stations, chunk): |
| section = slice(start, min(start + chunk, stations)) |
| patches, analyses, elevation = generate_chunk(meta, fold, section, False, int(config["seed"]) + 100) |
| source_expected, source_probabilities = [], [] |
| with torch.no_grad(): |
| for source in range(3): |
| patch = torch.from_numpy(patches[:, :, :, source]) |
| patch[:, :, 0] += (-0.0098 * torch.from_numpy(elevation)).view(1, 1, -1, 1, 1) |
| patch -= model.bias[source, :, :, section].unsqueeze(0).unsqueeze(-1).unsqueeze(-1) |
| centre = patch[..., 1, 1] |
| source_expected.append(centre) |
| raw_probabilities = [] |
| for variable, values in enumerate(thresholds): |
| width = model.fuzzy_widths[variable] |
| p = ((patch[:, :, variable].unsqueeze(2) - values.view(1, 1, -1, 1, 1, 1) + width) / (2 * width)).clamp(0, 1) |
| raw_probabilities.append(model.neighborhood(model.recursive_filter(p))) |
| source_probabilities.append(model.calibrate(raw_probabilities, source)) |
| blend = model.blend_expected(torch.stack(source_expected, dim=3)) |
| blend_probability = model.blend_probabilities(source_probabilities) |
| raw[fold, ..., section] = patches[0, :, :, 0, :, 1, 1] |
| truth[fold, ..., section] = analyses[0] |
| expected[fold, ..., section] = blend[0].numpy() |
| for variable in range(3): probabilities[variable][fold, ..., section] = blend_probability[variable][0].numpy() |
| arrays = [raw, truth, expected, *probabilities] |
| if not all(np.isfinite(array).all() for array in arrays): |
| raise ValueError("inference produced non-finite values") |
| if not all(((array >= 0) & (array <= 1)).all() for array in probabilities): |
| raise ValueError("probabilities must be within [0, 1]") |
| output = ROOT / config["paths"]["inference"] |
| output.parent.mkdir(parents=True, exist_ok=True) |
| payload = {"raw_aifs": raw, "analyses": truth, "blend_expected": expected, "valid_dates": meta["valid_dates"], "lead_hours": meta["lead_hours"], |
| "station_id": meta["station_id"], "station_latitude": meta["station_latitude"], "station_longitude": meta["station_longitude"], |
| "variables": meta["variables"], "units": meta["units"], "fold_id": np.arange(2)} |
| for variable, name in enumerate(meta["variables"]): |
| payload[f"thresholds_{name}"] = thresholds[variable].numpy(); payload[f"probability_blend_{name}"] = probabilities[variable] |
| np.savez_compressed(output, **payload) |
| print(f"saved={output.relative_to(ROOT)} expected_shape={expected.shape} folds=2 stations=569 thresholds=61/47/49") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|