"""Checkpoint-backed one-step and autoregressive RainNet inference.""" import json from pathlib import Path import sys import h5py import numpy as np import torch import torch.nn.functional as F import yaml ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from model.rainnet import build_rainnet def load_config(): with (ROOT / "conf/config.yaml").open(encoding="utf-8") as handle: return yaml.safe_load(handle) def inverse_transform(tensor): return torch.clamp(torch.exp(tensor) - 0.01, min=0.0) def describe(name, array): print( f"{name}: shape={array.shape}, dtype={array.dtype}, " f"min={array.min():.8f}, max={array.max():.8f}, mean={array.mean():.8f}" ) def main(): config = load_config() torch.manual_seed(config["seed"]) requested_device = config["device"] if requested_device == "auto": device = torch.device("cuda" if torch.cuda.is_available() else "cpu") else: device = torch.device(requested_device) checkpoint_path = ROOT / config["inference"]["checkpoint"] if not checkpoint_path.exists(): raise FileNotFoundError(f"Required checkpoint does not exist: {checkpoint_path}") model = build_rainnet(**config["model"]).to(device) checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) model.load_state_dict(checkpoint["model_state_dict"]) model.eval() print("checkpoint_loaded: True") print(f"checkpoint path: {checkpoint_path}") data = config["data"] data_path = ROOT / data["path"] with h5py.File(data_path, "r") as handle: keys = sorted(handle.keys()) test_start = data["train_frames"] + data["val_frames"] test_keys = keys[test_start : test_start + data["test_frames"]] raw = np.stack([handle[key][...] for key in test_keys]).astype(np.float32) needed = data["input_steps"] + config["inference"]["rollout_steps"] if len(raw) < needed: raise ValueError(f"Inference needs {needed} test frames, found {len(raw)}") inputs = torch.from_numpy(raw[: data["input_steps"]]).unsqueeze(0).to(device) targets = raw[data["input_steps"] : needed] log_window = torch.log(inputs + 0.01) pad_h = data["padded_height"] - data["raw_height"] pad_w = data["padded_width"] - data["raw_width"] pad = (pad_w // 2, pad_w - pad_w // 2, pad_h // 2, pad_h - pad_h // 2) crop = (pad_h // 2, pad_w // 2) predictions = [] with torch.inference_mode(): for _ in range(config["inference"]["rollout_steps"]): padded = F.pad(log_window, pad, mode="reflect") padded_prediction = model(padded) prediction = padded_prediction[ :, :, crop[0] : crop[0] + data["raw_height"], crop[1] : crop[1] + data["raw_width"] ] predictions.append(inverse_transform(prediction).squeeze(0).squeeze(0).cpu().numpy()) log_window = torch.cat((log_window[:, 1:], prediction), dim=1) predictions = np.stack(predictions).astype(np.float32) persistence = np.repeat(raw[data["input_steps"] - 1][None], len(predictions), axis=0).astype(np.float32) output_dir = ROOT / config["inference"]["output_dir"] output_dir.mkdir(parents=True, exist_ok=True) np.save(output_dir / "inputs.npy", raw[: data["input_steps"]]) np.save(output_dir / "predictions.npy", predictions) np.save(output_dir / "targets.npy", targets) np.save(output_dir / "persistence.npy", persistence) metadata = { "units": "mm/5min", "interval_minutes": data["interval_minutes"], "rollout_steps": len(predictions), "input_keys": test_keys[: data["input_steps"]], "target_keys": test_keys[data["input_steps"] : needed], "checkpoint": str(checkpoint_path), } with (output_dir / "metadata.json").open("w", encoding="utf-8") as handle: json.dump(metadata, handle, indent=2) describe("inputs", raw[: data["input_steps"]]) describe("predictions", predictions) describe("targets", targets) describe("persistence", persistence) print(f"One-step inference shape: {predictions[:1].shape}") print(f"Autoregressive rollout steps: {len(predictions)}") if __name__ == "__main__": main()