From e5a642a7f7cbbccc046b2b67e12b5766771f2d77 Mon Sep 17 00:00:00 2001 From: Jens Date: Mon, 27 Jul 2026 06:46:42 +0200 Subject: [PATCH] Add auditable SAM roof label refinement --- backend/tests/test_regional_yolo_dataset.py | 33 +++ .../tests/test_sam_roof_label_refinement.py | 23 ++ scripts/build_regional_yolo_dataset.py | 77 +++++++ scripts/refine_yolo_labels_with_sam.py | 198 ++++++++++++++++++ 4 files changed, 331 insertions(+) create mode 100644 backend/tests/test_regional_yolo_dataset.py create mode 100644 backend/tests/test_sam_roof_label_refinement.py create mode 100644 scripts/build_regional_yolo_dataset.py create mode 100644 scripts/refine_yolo_labels_with_sam.py diff --git a/backend/tests/test_regional_yolo_dataset.py b/backend/tests/test_regional_yolo_dataset.py new file mode 100644 index 00000000..f55f5e74 --- /dev/null +++ b/backend/tests/test_regional_yolo_dataset.py @@ -0,0 +1,33 @@ +from __future__ import annotations + +import importlib.util +from pathlib import Path + + +SCRIPT = Path(__file__).parents[2] / "scripts" / "build_regional_yolo_dataset.py" +SPEC = importlib.util.spec_from_file_location("regional_yolo_dataset", SCRIPT) +assert SPEC and SPEC.loader +MODULE = importlib.util.module_from_spec(SPEC) +SPEC.loader.exec_module(MODULE) + + +def test_select_paths_is_region_and_split_safe() -> None: + manifest = { + "samples": [ + {"sample_slug": "f-train", "region": "flanders", "split": "train"}, + {"sample_slug": "f-val", "region": "flanders", "split": "val"}, + {"sample_slug": "f-test", "region": "flanders", "split": "test"}, + {"sample_slug": "w-train", "region": "wallonia", "split": "train"}, + ] + } + summary = { + "tiles": [ + {"sample_slug": "f-train", "split": "train", "image_path": "/f-train.png"}, + {"sample_slug": "f-val", "split": "val", "image_path": "/f-val.png"}, + {"sample_slug": "f-test", "split": "val", "image_path": "/f-test.png"}, + {"sample_slug": "w-train", "split": "train", "image_path": "/w-train.png"}, + ] + } + train, val = MODULE.select_paths(summary, manifest, "flanders") + assert train == ["/f-train.png"] + assert val == ["/f-val.png"] diff --git a/backend/tests/test_sam_roof_label_refinement.py b/backend/tests/test_sam_roof_label_refinement.py new file mode 100644 index 00000000..7f7c9abd --- /dev/null +++ b/backend/tests/test_sam_roof_label_refinement.py @@ -0,0 +1,23 @@ +from __future__ import annotations + +import importlib.util +from pathlib import Path + + +SCRIPT = Path(__file__).parents[2] / "scripts" / "refine_yolo_labels_with_sam.py" +SPEC = importlib.util.spec_from_file_location("sam_roof_refinement", SCRIPT) +assert SPEC and SPEC.loader +MODULE = importlib.util.module_from_spec(SPEC) +SPEC.loader.exec_module(MODULE) + + +def test_plausible_refinement_is_fail_closed() -> None: + source = (10.0, 10.0, 30.0, 30.0) + assert MODULE.plausible_refinement(source, (8.0, 9.0, 31.0, 32.0), min_iou=0.15, min_area_ratio=0.25, max_area_ratio=4.0) + assert not MODULE.plausible_refinement(source, (100.0, 100.0, 120.0, 120.0), min_iou=0.15, min_area_ratio=0.25, max_area_ratio=4.0) + assert not MODULE.plausible_refinement(source, (0.0, 0.0, 100.0, 100.0), min_iou=0.15, min_area_ratio=0.25, max_area_ratio=4.0) + + +def test_yolo_round_trip_shape() -> None: + line = MODULE.yolo_line((10.0, 20.0, 30.0, 40.0), 100, 100) + assert line == "0 0.20000000 0.30000000 0.20000000 0.20000000" diff --git a/scripts/build_regional_yolo_dataset.py b/scripts/build_regional_yolo_dataset.py new file mode 100644 index 00000000..aa0b2494 --- /dev/null +++ b/scripts/build_regional_yolo_dataset.py @@ -0,0 +1,77 @@ +#!/usr/bin/env python3 +"""Create a checksummed regional YOLO view without copying protected data.""" + +from __future__ import annotations + +import argparse +import hashlib +import json +from pathlib import Path +from typing import Any + + +def sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for chunk in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def select_paths(summary: dict[str, Any], manifest: dict[str, Any], region: str) -> tuple[list[str], list[str]]: + samples = {item["sample_slug"]: item for item in manifest["samples"]} + train: list[str] = [] + val: list[str] = [] + for tile in summary["tiles"]: + sample = samples[tile["sample_slug"]] + if sample["region"] != region or not tile.get("kept", True): + continue + if sample["split"] == "train" and tile["split"] == "train": + train.append(tile["image_path"]) + elif sample["split"] == "val" and tile["split"] == "val": + val.append(tile["image_path"]) + if not train or not val: + raise ValueError(f"Region {region!r} must contain train and validation images") + return sorted(train), sorted(val) + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("--summary", type=Path, required=True) + parser.add_argument("--corpus-manifest", type=Path, required=True) + parser.add_argument("--region", required=True) + parser.add_argument("--output-dir", type=Path, required=True) + args = parser.parse_args() + summary = json.loads(args.summary.read_text(encoding="utf-8")) + manifest = json.loads(args.corpus_manifest.read_text(encoding="utf-8")) + train, val = select_paths(summary, manifest, args.region) + args.output_dir.mkdir(parents=True, exist_ok=True) + train_list = args.output_dir / "train.txt" + val_list = args.output_dir / "val.txt" + train_list.write_text("\n".join(train) + "\n", encoding="utf-8") + val_list.write_text("\n".join(val) + "\n", encoding="utf-8") + dataset_yaml = args.output_dir / "dataset.yaml" + dataset_yaml.write_text( + f"path: {args.output_dir}\ntrain: {train_list}\nval: {val_list}\nnames:\n 0: building\n", + encoding="utf-8", + ) + evidence = { + "schema_version": 1, + "status": "ok", + "region": args.region, + "train_image_count": len(train), + "validation_image_count": len(val), + "summary": str(args.summary), + "summary_sha256": sha256(args.summary), + "corpus_manifest": str(args.corpus_manifest), + "corpus_manifest_sha256": sha256(args.corpus_manifest), + "dataset_yaml": str(dataset_yaml), + "protected_splits_in_training": [], + } + (args.output_dir / "regional-dataset.json").write_text(json.dumps(evidence, indent=2), encoding="utf-8") + print(json.dumps(evidence, indent=2)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/refine_yolo_labels_with_sam.py b/scripts/refine_yolo_labels_with_sam.py new file mode 100644 index 00000000..2d155192 --- /dev/null +++ b/scripts/refine_yolo_labels_with_sam.py @@ -0,0 +1,198 @@ +#!/usr/bin/env python3 +"""Refine official footprint boxes into auditable image-visible SAM roof boxes.""" + +from __future__ import annotations + +import argparse +import gc +import hashlib +import json +import os +import shutil +from pathlib import Path +from typing import Any + + +def sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for chunk in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def iou(left: tuple[float, float, float, float], right: tuple[float, float, float, float]) -> float: + ix = max(0.0, min(left[2], right[2]) - max(left[0], right[0])) + iy = max(0.0, min(left[3], right[3]) - max(left[1], right[1])) + intersection = ix * iy + union = (left[2] - left[0]) * (left[3] - left[1]) + (right[2] - right[0]) * ( + right[3] - right[1] + ) - intersection + return intersection / union if union > 0 else 0.0 + + +def plausible_refinement( + source: tuple[float, float, float, float], + refined: tuple[float, float, float, float], + *, + min_iou: float, + min_area_ratio: float, + max_area_ratio: float, +) -> bool: + source_area = (source[2] - source[0]) * (source[3] - source[1]) + refined_area = (refined[2] - refined[0]) * (refined[3] - refined[1]) + ratio = refined_area / source_area if source_area > 0 else 0.0 + return min_area_ratio <= ratio <= max_area_ratio and iou(source, refined) >= min_iou + + +def read_boxes(path: Path, width: int, height: int) -> list[tuple[float, float, float, float]]: + boxes = [] + for line in path.read_text(encoding="utf-8").splitlines() if path.is_file() else []: + class_id, cx, cy, bw, bh = map(float, line.split()) + if class_id != 0: + raise ValueError(f"Unexpected class in {path}: {class_id}") + boxes.append(((cx - bw / 2) * width, (cy - bh / 2) * height, (cx + bw / 2) * width, (cy + bh / 2) * height)) + return boxes + + +def yolo_line(box: tuple[float, float, float, float], width: int, height: int) -> str: + x1, y1, x2, y2 = box + return f"0 {(x1+x2)/(2*width):.8f} {(y1+y2)/(2*height):.8f} {(x2-x1)/width:.8f} {(y2-y1)/height:.8f}" + + +def expanded(box: tuple[float, float, float, float], factor: float, width: int, height: int) -> list[float]: + x1, y1, x2, y2 = box + cx, cy = (x1 + x2) / 2, (y1 + y2) / 2 + hw, hh = (x2 - x1) * factor / 2, (y2 - y1) * factor / 2 + return [max(0, cx - hw), max(0, cy - hh), min(width - 1, cx + hw), min(height - 1, cy + hh)] + + +def link_or_copy(source: Path, target: Path) -> None: + target.parent.mkdir(parents=True, exist_ok=True) + try: + os.link(source, target) + except OSError: + shutil.copy2(source, target) + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("--summary", type=Path, required=True) + parser.add_argument("--model", type=Path, required=True) + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument("--device", default="cuda:0") + parser.add_argument("--prompt-scale", type=float, default=1.5) + parser.add_argument("--min-source-iou", type=float, default=0.15) + parser.add_argument("--min-area-ratio", type=float, default=0.25) + parser.add_argument("--max-area-ratio", type=float, default=4.0) + parser.add_argument("--max-prompts-per-pass", type=int, default=96) + parser.add_argument("--force", action="store_true") + args = parser.parse_args() + if args.output_dir.exists(): + if not args.force: + raise SystemExit(f"Output exists: {args.output_dir}") + shutil.rmtree(args.output_dir) + from PIL import Image + import torch + from ultralytics import SAM + + summary = json.loads(args.summary.read_text(encoding="utf-8")) + model = SAM(str(args.model)) + refined_count = fallback_count = 0 + reason_counts: dict[str, int] = {} + output_tiles = [] + for index, tile in enumerate(summary["tiles"], start=1): + source_image = Path(tile["image_path"]) + source_label = Path(tile["label_path"]) + split = tile["split"] + target_image = args.output_dir / "images" / split / source_image.name + target_label = args.output_dir / "labels" / split / source_label.name + link_or_copy(source_image, target_image) + target_label.parent.mkdir(parents=True, exist_ok=True) + with Image.open(source_image) as image: + width, height = image.size + source_boxes = read_boxes(source_label, width, height) + output_boxes = list(source_boxes) + if source_boxes: + for start in range(0, len(source_boxes), args.max_prompts_per_pass): + source_chunk = source_boxes[start : start + args.max_prompts_per_pass] + prompts = [expanded(box, args.prompt_scale, width, height) for box in source_chunk] + result = model.predict(str(source_image), bboxes=prompts, device=args.device, verbose=False)[0] + masks = result.masks.data.cpu().numpy() if result.masks is not None else [] + candidates = [] + for mask in masks: + ys, xs = mask.nonzero() + if len(xs): + candidates.append((float(xs.min()), float(ys.min()), float(xs.max() + 1), float(ys.max() + 1))) + unmatched = set(range(len(candidates))) + for local_index, source_box in enumerate(source_chunk): + ranked = sorted(((iou(source_box, candidates[item]), item) for item in unmatched), reverse=True) + overlap, candidate_index = ranked[0] if ranked else (0.0, -1) + if candidate_index < 0 or overlap < args.min_source_iou: + reason_counts["unmatched_mask"] = reason_counts.get("unmatched_mask", 0) + 1 + fallback_count += 1 + continue + candidate = candidates[candidate_index] + if plausible_refinement( + source_box, + candidate, + min_iou=args.min_source_iou, + min_area_ratio=args.min_area_ratio, + max_area_ratio=args.max_area_ratio, + ): + output_boxes[start + local_index] = candidate + unmatched.remove(candidate_index) + refined_count += 1 + else: + reason_counts["geometry_gate"] = reason_counts.get("geometry_gate", 0) + 1 + fallback_count += 1 + del result, masks + if model.predictor is not None: + model.predictor.reset_image() + gc.collect() + torch.cuda.empty_cache() + target_label.write_text("\n".join(yolo_line(box, width, height) for box in output_boxes) + ("\n" if output_boxes else ""), encoding="utf-8") + output_tile = dict(tile) + output_tile.update({"image_path": str(target_image), "label_path": str(target_label)}) + output_tiles.append(output_tile) + print(f"{index}/{len(summary['tiles'])} {tile['sample_slug']}: {len(source_boxes)}", flush=True) + + output_summary = dict(summary) + output_summary.update( + { + "output_dir": str(args.output_dir), + "dataset_yaml": str(args.output_dir / "dataset.yaml"), + "tiles": output_tiles, + "label_semantics": "sam_visible_roof_with_official_footprint_fallback", + } + ) + summary_path = args.output_dir / "yolo_tile_dataset_summary.json" + summary_path.write_text(json.dumps(output_summary, indent=2), encoding="utf-8") + (args.output_dir / "dataset.yaml").write_text( + f"path: {args.output_dir}\ntrain: images/train\nval: images/val\nnames:\n 0: building\n", + encoding="utf-8", + ) + evidence = { + "schema_version": 1, + "status": "ok", + "source_summary": str(args.summary), + "source_summary_sha256": sha256(args.summary), + "sam_model": str(args.model), + "sam_model_sha256": sha256(args.model), + "device": args.device, + "prompt_scale": args.prompt_scale, + "min_source_iou": args.min_source_iou, + "min_area_ratio": args.min_area_ratio, + "max_area_ratio": args.max_area_ratio, + "max_prompts_per_pass": args.max_prompts_per_pass, + "refined_label_count": refined_count, + "fallback_label_count": fallback_count, + "fallback_reason_counts": reason_counts, + } + (args.output_dir / "sam-refinement.json").write_text(json.dumps(evidence, indent=2), encoding="utf-8") + print(json.dumps(evidence, indent=2)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main())