| """Compute Satlas semantic, instance, geometric, and global metrics.""" |
|
|
| import argparse |
| import json |
| from collections import deque |
| from pathlib import Path |
|
|
| import matplotlib.pyplot as plt |
| import numpy as np |
| import torch |
| from torch.nn import functional as F |
| import yaml |
|
|
| ROOT = Path(__file__).resolve().parents[1] |
|
|
|
|
| def precision_recall_f1(tp, fp, fn): |
| precision = tp / max(tp + fp, 1); recall = tp / max(tp + fn, 1) |
| return float(precision), float(recall), float(2 * precision * recall / max(precision + recall, 1e-12)) |
|
|
|
|
| def connected_components(mask): |
| """Return 4-connected component coordinate arrays without SciPy.""" |
| mask = np.asarray(mask, dtype=bool); visited = np.zeros_like(mask); components = [] |
| height, width = mask.shape |
| for row, col in np.argwhere(mask): |
| if visited[row, col]: continue |
| queue = deque([(int(row), int(col))]); visited[row, col] = True; points = [] |
| while queue: |
| y, x = queue.popleft(); points.append((y, x)) |
| for ny, nx in ((y - 1, x), (y + 1, x), (y, x - 1), (y, x + 1)): |
| if 0 <= ny < height and 0 <= nx < width and mask[ny, nx] and not visited[ny, nx]: |
| visited[ny, nx] = True; queue.append((ny, nx)) |
| components.append(np.asarray(points, dtype=np.int32)) |
| return components |
|
|
|
|
| def local_peaks(scores, threshold, radius): |
| tensor = torch.from_numpy(scores)[None, None] |
| pooled = F.max_pool2d(tensor, kernel_size=2 * radius + 1, stride=1, padding=radius) |
| candidates = np.argwhere(np.logical_and(scores >= threshold, scores == pooled[0, 0].numpy())) |
| if not len(candidates): return np.empty((0, 2), dtype=np.float32) |
| order = sorted(candidates, key=lambda point: scores[tuple(point)], reverse=True); selected = [] |
| for point in order: |
| if all(np.sum((point - previous) ** 2) > radius ** 2 for previous in selected): selected.append(point) |
| return np.asarray(selected, dtype=np.float32).reshape(-1, 2) |
|
|
|
|
| def point_counts(predictions, targets, threshold, nms_radius, tolerance): |
| tp = fp = fn = 0 |
| for scores, target in zip(predictions[:, 0], targets[:, 0]): |
| predicted = local_peaks(scores, threshold, nms_radius) |
| truth_components = connected_components(target >= 0.5) |
| truth = np.asarray([component.mean(0) for component in truth_components], dtype=np.float32).reshape(-1, 2) |
| candidates = sorted((float(np.linalg.norm(p - t)), pi, ti) for pi, p in enumerate(predicted) |
| for ti, t in enumerate(truth) if np.linalg.norm(p - t) <= tolerance) |
| matched_pred, matched_truth = set(), set() |
| for _, pi, ti in candidates: |
| if pi not in matched_pred and ti not in matched_truth: matched_pred.add(pi); matched_truth.add(ti) |
| tp += len(matched_pred); fp += len(predicted) - len(matched_pred); fn += len(truth) - len(matched_truth) |
| return tp, fp, fn |
|
|
|
|
| def component_mask(component, shape): |
| mask = np.zeros(shape, dtype=bool); mask[component[:, 0], component[:, 1]] = True; return mask |
|
|
|
|
| def polygon_counts(predictions, targets, threshold, iou_threshold): |
| tp = fp = fn = 0; matched_ious = [] |
| for scores, target in zip(predictions[:, 0], targets[:, 0]): |
| predicted = connected_components(scores >= threshold); truth = connected_components(target >= 0.5) |
| pred_masks = [component_mask(component, scores.shape) for component in predicted] |
| truth_masks = [component_mask(component, scores.shape) for component in truth] |
| candidates = [] |
| for pi, pred in enumerate(pred_masks): |
| for ti, actual in enumerate(truth_masks): |
| union = np.logical_or(pred, actual).sum(); iou = np.logical_and(pred, actual).sum() / max(union, 1) |
| if iou >= iou_threshold: candidates.append((float(iou), pi, ti)) |
| matched_pred, matched_truth = set(), set() |
| for iou, pi, ti in sorted(candidates, reverse=True): |
| if pi not in matched_pred and ti not in matched_truth: |
| matched_pred.add(pi); matched_truth.add(ti); matched_ious.append(iou) |
| tp += len(matched_pred); fp += len(predicted) - len(matched_pred); fn += len(truth) - len(matched_truth) |
| return tp, fp, fn, float(np.mean(matched_ious)) if matched_ious else 0.0 |
|
|
|
|
| def dilate(mask, radius): |
| height, width = mask.shape[-2:]; padded = np.pad(mask, ((0, 0), (radius, radius), (radius, radius))) |
| neighborhoods = [padded[:, dy:dy + height, dx:dx + width] |
| for dy in range(2 * radius + 1) for dx in range(2 * radius + 1) |
| if (dy - radius) ** 2 + (dx - radius) ** 2 <= radius ** 2] |
| return np.logical_or.reduce(neighborhoods) |
|
|
|
|
| def tolerant_line_metrics(prediction, target, threshold, tolerance): |
| predicted, truth = prediction[:, 0] >= threshold, target[:, 0] >= 0.5 |
| matched_pred = np.logical_and(predicted, dilate(truth, tolerance)).sum() |
| matched_truth = np.logical_and(truth, dilate(predicted, tolerance)).sum() |
| precision = matched_pred / max(predicted.sum(), 1); recall = matched_truth / max(truth.sum(), 1) |
| return float(precision), float(recall), float(2 * precision * recall / max(precision + recall, 1e-12)) |
|
|
|
|
| def mean_iou(prediction, target, classes): |
| values = [] |
| for label in range(classes): |
| pred, truth = prediction == label, target == label; union = np.logical_or(pred, truth).sum() |
| if union: values.append(np.logical_and(pred, truth).sum() / union) |
| return float(np.mean(values)) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml"); parser.add_argument("--prediction", type=Path) |
| parser.add_argument("--target", type=Path); parser.add_argument("--output-dir", type=Path); args = parser.parse_args() |
| config = yaml.safe_load(args.config.read_text()); evaluation = config["evaluation"] |
| pred_path = args.prediction or ROOT / config["paths"]["inference_dir"] / "predictions.npz" |
| target_path = args.target or ROOT / config["data"]["root"] / "test.npz" |
| if not pred_path.is_file(): raise FileNotFoundError("Run inference before evaluation") |
| prediction, target = np.load(pred_path), np.load(target_path) |
| if str(prediction["protocol"]) != str(target["protocol"]) or str(target["protocol"]) != config["data"]["protocol"]: |
| raise ValueError("prediction, target, and configuration protocols do not match") |
| if str(prediction["source"]) != str(target["source"]): raise ValueError("prediction and target sources do not match") |
| if prediction["sample_ids"].shape != target["sample_ids"].shape or not np.array_equal(prediction["sample_ids"], target["sample_ids"]): |
| raise ValueError("prediction sample_ids do not exactly match target identity/order") |
| segmentation = prediction["segmentation"].argmax(1) |
| point = precision_recall_f1(*point_counts(prediction["point"], target["point"], evaluation["point_peak_threshold"], |
| evaluation["point_nms_radius"], evaluation["point_distance_tolerance"])) |
| polygon_counts_result = polygon_counts(prediction["polygon"], target["polygon"], evaluation["mask_threshold"], |
| evaluation["polygon_iou_threshold"]) |
| polygon = precision_recall_f1(*polygon_counts_result[:3]); line = tolerant_line_metrics( |
| prediction["polyline"], target["polyline"], evaluation["mask_threshold"], evaluation["polyline_distance_tolerance"]) |
| metrics = { |
| "segmentation_mIoU": mean_iou(segmentation, target["segmentation"], config["model"]["segmentation_classes"]), |
| "regression_MAE": float(np.abs(prediction["regression"] - target["regression"]).mean()), |
| "point_precision": point[0], "point_recall": point[1], "point_F1": point[2], |
| "polygon_precision": polygon[0], "polygon_recall": polygon[1], "polygon_F1": polygon[2], |
| "polygon_matched_IoU": polygon_counts_result[3], |
| "polyline_precision": line[0], "polyline_recall": line[1], "polyline_F1": line[2], |
| "property_accuracy": float((prediction["property"].argmax(1) == target["property"]).mean()), |
| "classification_accuracy": float((prediction["classification"].argmax(1) == target["classification"]).mean()), |
| "samples": int(len(segmentation)), "protocol": str(prediction["protocol"]), "checkpoint": str(prediction["checkpoint"]), |
| "source": str(prediction["source"]), "identity_verified": True, |
| } |
| output = args.output_dir or ROOT / config["paths"]["evaluation_dir"]; output.mkdir(parents=True, exist_ok=True) |
| (output / "metrics.json").write_text(json.dumps(metrics, indent=2) + "\n") |
| last_highres = np.flatnonzero(target["valid_highres_times"][0])[-1] |
| last_lowres = np.flatnonzero(target["valid_lowres_times"][0])[-1] |
| image = np.clip(target["highres_images"][0, last_highres].transpose(1, 2, 0), 0, 1) |
| figure, axes = plt.subplots(2, 4, figsize=(12, 6)); panels = [ |
| ("NAIP RGB", image, None), ("Segmentation GT", target["segmentation"][0], "tab20"), |
| ("Segmentation pred", segmentation[0], "tab20"), ("Regression", prediction["regression"][0, 0], "viridis"), |
| ("Point", prediction["point"][0, 0], "magma"), ("Polygon", prediction["polygon"][0, 0], "magma"), |
| ("Polyline", prediction["polyline"][0, 0], "magma"), |
| ("Sentinel band", target["lowres_images"][0, last_lowres, 0], "viridis")] |
| for axis, (title, panel, cmap) in zip(axes.flat, panels): axis.imshow(panel, cmap=cmap); axis.set_title(title); axis.axis("off") |
| figure.tight_layout(); figure.savefig(output / "multitask_predictions.png", dpi=150); plt.close(figure) |
| print(json.dumps(metrics, indent=2)); print(f"evaluation={output}") |
|
|
|
|
| if __name__ == "__main__": main() |
|
|