OneScience-Group/FireCubeNet
025
1"""Infer next-day center-pixel wildfire probabilities."""2 3import sys4from pathlib import Path5 6import numpy as np7import torch8import yaml9from torch.utils.data import DataLoader10 11 12ROOT = Path(__file__).resolve().parents[1]13sys.path.insert(0, str(ROOT))14from model.firecubenet import FireCubeNet15from train import WildfireDataset, device_from_config16 17 18def main():19 config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())20 device = device_from_config(config)21 checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location=device, weights_only=False)22 if checkpoint["format_version"] != config["data"]["format_version"]:23 raise ValueError("checkpoint and data format versions differ")24 dataset = WildfireDataset(ROOT / config["data"]["root"] / "test.npz", config)25 loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), shuffle=False)26 model = FireCubeNet(**checkpoint["model_config"]).to(device)27 model.load_state_dict(checkpoint["model_state_dict"])28 model.eval()29 mean = torch.from_numpy(checkpoint["channel_mean"]).to(device).view(1, 1, -1, 1, 1)30 std = torch.from_numpy(checkpoint["channel_std"]).to(device).view(1, 1, -1, 1, 1)31 probabilities = []32 with torch.no_grad():33 for inputs, _ in loader:34 probabilities.append(torch.sigmoid(model((inputs.to(device) - mean) / std)).cpu().numpy())35 probabilities = np.concatenate(probabilities).astype(np.float32)36 if probabilities.shape != dataset.data["labels"].shape or not np.isfinite(probabilities).all():37 raise FloatingPointError("invalid inference probabilities")38 output = ROOT / config["paths"]["inference_dir"] / "predictions.npz"39 output.parent.mkdir(parents=True, exist_ok=True)40 np.savez_compressed(41 output, probabilities=probabilities, labels=dataset.data["labels"],42 timestamps=dataset.data["timestamps_unix_s"], coords=dataset.data["coords"],43 format_version=np.asarray(config["data"]["format_version"]),44 )45 print(f"predictions={output.relative_to(ROOT)} shape={probabilities.shape}")46 47 48if __name__ == "__main__":49 main()50 