| """Apply all four UnetDif heads and save 3-hour and 24-hour products.""" |
|
|
| 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.unetdif import UnetDif, apply_correction |
|
|
|
|
| def main(): |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location=device, weights_only=True) |
| required_checkpoint = {"epoch", "model", "optimizer", "model_config", "loss_normalizers", "loss_components", "format_version"} |
| if not required_checkpoint.issubset(checkpoint): |
| raise ValueError(f"checkpoint is missing {sorted(required_checkpoint - checkpoint.keys())}") |
| if checkpoint["format_version"] != config["data"]["format_version"]: |
| raise ValueError("checkpoint and configured data format versions differ") |
| model = UnetDif(**checkpoint["model_config"]).to(device) |
| model.load_state_dict(checkpoint["model"]); model.eval() |
| archive = np.load(ROOT / config["data"]["root"] / "test.npz", allow_pickle=False) |
| inputs = archive["inputs"].astype(np.float32, copy=False) |
| expected = (int(config["data"]["time_steps"]), int(config["data"]["channels"]), |
| int(config["data"]["height"]), int(config["data"]["width"])) |
| if inputs.shape[1:] != expected or archive["targets"].shape[1:] != (expected[0], expected[2], expected[3]): |
| raise ValueError(f"test NPZ does not satisfy the fixed day contract {expected}") |
| if archive["lead_hours"].tolist() != config["data"]["lead_hours"]: |
| raise ValueError("test NPZ lead_hours do not match configuration") |
| flat = torch.from_numpy(inputs.reshape(-1, *inputs.shape[2:])).to(device) |
| heads_all = {key: [] for key in ("dry_logits", "false_alarm_logits", "positive_bias", "negative_bias")} |
| corrected = [] |
| batch_size = int(config["inference"]["batch_size"]) |
| with torch.no_grad(): |
| for start in range(0, len(flat), batch_size): |
| batch = flat[start:start + batch_size] |
| heads = model(batch) |
| corrected.append(apply_correction(batch[:, int(config["data"]["precipitation_channel"])], heads, |
| float(config["inference"]["probability_threshold"])).cpu()) |
| for key in heads_all: |
| heads_all[key].append(heads[key].cpu()) |
| day_shape = inputs.shape[:2] + inputs.shape[-2:] |
| corrected_3h = torch.cat(corrected).numpy().reshape(day_shape) |
| raw_3h = inputs[:, :, int(config["data"]["precipitation_channel"])] |
| output = ROOT / config["paths"]["inference"] |
| output.parent.mkdir(parents=True, exist_ok=True) |
| payload = {key: torch.cat(value).numpy().reshape(day_shape) for key, value in heads_all.items()} |
| np.savez_compressed( |
| output, inputs=inputs, group_id=archive["group_id"], lead_hours=archive["lead_hours"], |
| raw_3h=raw_3h, target_3h=archive["targets"], corrected_3h=corrected_3h, |
| raw_24h=raw_3h.sum(1), target_24h=archive["targets"].sum(1), corrected_24h=corrected_3h.sum(1), **payload, |
| ) |
| print(f"saved={output.relative_to(ROOT)} corrected_3h={corrected_3h.shape}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|