RainNet / scripts /inference.py
Niulh's picture
Publish RainNet model package
9f29df6 verified
Raw
History Blame Contribute Delete
4.28 kB
"""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()