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