CoolFace
Modelpublic

OneScience-Group/FireCubeNet

sourceHugging Facemitupdated 18d agoView on Hugging Face
0likes25downloads
inference.py50 linesDownload Raw Back to scripts
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