"""Infer next-day center-pixel wildfire probabilities.""" import sys from pathlib import Path 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.firecubenet import FireCubeNet from train import WildfireDataset, 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) if checkpoint["format_version"] != config["data"]["format_version"]: raise ValueError("checkpoint and data format versions differ") dataset = WildfireDataset(ROOT / config["data"]["root"] / "test.npz", config) loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), shuffle=False) model = FireCubeNet(**checkpoint["model_config"]).to(device) model.load_state_dict(checkpoint["model_state_dict"]) model.eval() mean = torch.from_numpy(checkpoint["channel_mean"]).to(device).view(1, 1, -1, 1, 1) std = torch.from_numpy(checkpoint["channel_std"]).to(device).view(1, 1, -1, 1, 1) probabilities = [] with torch.no_grad(): for inputs, _ in loader: probabilities.append(torch.sigmoid(model((inputs.to(device) - mean) / std)).cpu().numpy()) probabilities = np.concatenate(probabilities).astype(np.float32) if probabilities.shape != dataset.data["labels"].shape or not np.isfinite(probabilities).all(): raise FloatingPointError("invalid inference probabilities") output = ROOT / config["paths"]["inference_dir"] / "predictions.npz" output.parent.mkdir(parents=True, exist_ok=True) np.savez_compressed( output, probabilities=probabilities, labels=dataset.data["labels"], timestamps=dataset.data["timestamps_unix_s"], coords=dataset.data["coords"], format_version=np.asarray(config["data"]["format_version"]), ) print(f"predictions={output.relative_to(ROOT)} shape={probabilities.shape}") if __name__ == "__main__": main()