"""Evaluate binary wildfire danger predictions and create task-specific plots.""" import json from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np import yaml ROOT = Path(__file__).resolve().parents[1] def roc_curve_and_auc(labels, probabilities): order = np.argsort(-probabilities, kind="stable") sorted_labels = labels[order] positives = max(int(labels.sum()), 1) negatives = max(int((1 - labels).sum()), 1) true_positive_rate = np.r_[0.0, np.cumsum(sorted_labels) / positives, 1.0] false_positive_rate = np.r_[0.0, np.cumsum(1 - sorted_labels) / negatives, 1.0] return false_positive_rate, true_positive_rate, float(np.trapz(true_positive_rate, false_positive_rate)) def main(): config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) data = np.load(ROOT / config["paths"]["inference_dir"] / "predictions.npz") if str(data["format_version"]) != config["data"]["format_version"]: raise ValueError("incompatible prediction format") probabilities = data["probabilities"].reshape(-1) labels = data["labels"].reshape(-1).astype(np.int64) if probabilities.shape != labels.shape or not np.isfinite(probabilities).all() or not np.isin(labels, (0, 1)).all(): raise ValueError("probabilities/labels are invalid") threshold = float(config["evaluation"]["threshold"]) predictions = (probabilities >= threshold).astype(np.int64) tp = int(((predictions == 1) & (labels == 1)).sum()) fp = int(((predictions == 1) & (labels == 0)).sum()) tn = int(((predictions == 0) & (labels == 0)).sum()) fn = int(((predictions == 0) & (labels == 1)).sum()) precision = tp / (tp + fp) if tp + fp else 0.0 recall = tp / (tp + fn) if tp + fn else 0.0 f1 = 2 * precision * recall / (precision + recall) if precision + recall else 0.0 fpr, tpr, auroc = roc_curve_and_auc(labels, probabilities) report = { "samples": int(len(labels)), "threshold": threshold, "precision": precision, "recall": recall, "f1": f1, "auroc": auroc, "confusion_matrix": {"true_negative": tn, "false_positive": fp, "false_negative": fn, "true_positive": tp}, "note": "Synthetic engineering validation; not paper test-set performance." } if not np.isfinite([precision, recall, f1, auroc]).all(): raise FloatingPointError("evaluation contains non-finite metrics") output = ROOT / config["paths"]["evaluation_dir"] output.mkdir(parents=True, exist_ok=True) (output / "metrics.json").write_text(json.dumps(report, indent=2) + "\n") order = np.argsort(data["timestamps"]) figure, axes = plt.subplots(1, 2, figsize=(11, 4.2)) axes[0].plot(fpr, tpr, color="firebrick", linewidth=2, label=f"ConvLSTM (AUROC={auroc:.3f})") axes[0].plot([0, 1], [0, 1], "k--", linewidth=1) axes[0].set(xlabel="False positive rate", ylabel="True positive rate", title="Next-day wildfire ROC") axes[0].legend() colors = np.where(labels[order] == 1, "firebrick", "steelblue") axes[1].scatter(np.arange(len(labels)), probabilities[order], c=colors, s=45) axes[1].axhline(threshold, color="black", linestyle="--", linewidth=1, label="threshold=0.5") axes[1].set(xlabel="Chronological sample", ylabel="Wildfire danger probability", title="Center-pixel next-day danger", ylim=(0, 1)) axes[1].legend() figure.tight_layout() figure.savefig(output / "wildfire_danger.png", dpi=150) plt.close(figure) print(f"evaluation={output.relative_to(ROOT)} f1={f1:.3f} auroc={auroc:.3f}") if __name__ == "__main__": main()