| """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() |
|
|