Spaces:
Runtime error
Runtime error
| """Janus — Feature Extraction via SAM 3 with Tiling + Live Dashboard. | |
| Upload large GeoTIFFs, select feature types, and watch as SAM 3 processes | |
| tile by tile with live progress and confidence tracking. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| import os | |
| import tempfile | |
| import time | |
| from collections import defaultdict | |
| from pathlib import Path | |
| import gradio as gr | |
| import geopandas as gpd | |
| import spaces | |
| import matplotlib | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| import rasterio | |
| import torch | |
| from affine import Affine | |
| from PIL import Image | |
| from pyproj import CRS | |
| from rasterio.features import shapes | |
| from rasterio.plot import show | |
| from rasterio.windows import Window | |
| from shapely.geometry import shape | |
| # Experimental backends from janus_core (additive — legacy modes still live here) | |
| try: | |
| from janus_core.config import ( | |
| ExperimentalConfig, YoloConfig, Sam2Config, | |
| PrefilterConfig, TilingConfig, RouterConfig, BenchmarkConfig, | |
| ) | |
| from janus_core.backends.base import SegmentResult | |
| from janus_core.backends.yolo_seg import YoloSegBackend | |
| from janus_core.backends.sam2_refine import Sam2RefineBackend | |
| JANUS_CORE_AVAILABLE = True | |
| except ImportError as _e: | |
| print(f"[CORE] janus_core unavailable: {_e}") | |
| JANUS_CORE_AVAILABLE = False | |
| # Lazy backend singletons (created on first use) | |
| _YOLO_BACKEND = None | |
| _SAM2_BACKEND = None | |
| def _get_yolo_backend(yolo_config): | |
| """Get-or-create the YOLO backend, returns None if unavailable.""" | |
| global _YOLO_BACKEND | |
| if not JANUS_CORE_AVAILABLE: | |
| return None | |
| if _YOLO_BACKEND is None or _YOLO_BACKEND.cfg is not yolo_config: | |
| _YOLO_BACKEND = YoloSegBackend(yolo_config) | |
| if not _YOLO_BACKEND.is_available(): | |
| return None | |
| return _YOLO_BACKEND | |
| def _get_sam2_backend(sam2_config): | |
| """Get-or-create the SAM 2 refinement backend, returns None if unavailable.""" | |
| global _SAM2_BACKEND | |
| if not JANUS_CORE_AVAILABLE: | |
| return None | |
| if _SAM2_BACKEND is None or _SAM2_BACKEND.cfg is not sam2_config: | |
| _SAM2_BACKEND = Sam2RefineBackend(sam2_config) | |
| if not _SAM2_BACKEND.is_available(): | |
| return None | |
| return _SAM2_BACKEND | |
| # --------------------------------------------------------------------------- | |
| # Globals | |
| # --------------------------------------------------------------------------- | |
| MODEL = None | |
| PROCESSOR = None | |
| def get_device(): | |
| """Detect GPU at runtime, not import time.""" | |
| import subprocess | |
| print(f"[GPU DEBUG] torch.cuda.is_available() = {torch.cuda.is_available()}") | |
| print(f"[GPU DEBUG] torch.version.cuda = {torch.version.cuda}") | |
| print(f"[GPU DEBUG] torch.backends.cudnn.enabled = {torch.backends.cudnn.enabled}") | |
| try: | |
| print(f"[GPU DEBUG] torch.cuda.device_count() = {torch.cuda.device_count()}") | |
| except Exception as e: | |
| print(f"[GPU DEBUG] torch.cuda.device_count() ERROR: {e}") | |
| try: | |
| result = subprocess.run(["nvidia-smi"], capture_output=True, text=True, timeout=5) | |
| print(f"[GPU DEBUG] nvidia-smi output:\n{result.stdout[:500]}") | |
| if result.stderr: | |
| print(f"[GPU DEBUG] nvidia-smi stderr: {result.stderr[:200]}") | |
| except Exception as e: | |
| print(f"[GPU DEBUG] nvidia-smi ERROR: {e}") | |
| try: | |
| import os | |
| print(f"[GPU DEBUG] CUDA_VISIBLE_DEVICES = {os.environ.get('CUDA_VISIBLE_DEVICES', 'NOT SET')}") | |
| print(f"[GPU DEBUG] NVIDIA_VISIBLE_DEVICES = {os.environ.get('NVIDIA_VISIBLE_DEVICES', 'NOT SET')}") | |
| except Exception as e: | |
| print(f"[GPU DEBUG] env check ERROR: {e}") | |
| if torch.cuda.is_available(): | |
| name = torch.cuda.get_device_name(0) | |
| print(f"[GPU DEBUG] Using GPU: {name}") | |
| return "cuda" | |
| print("[GPU DEBUG] Falling back to CPU") | |
| return "cpu" | |
| matplotlib.use("Agg") | |
| # --------------------------------------------------------------------------- | |
| # Feature presets | |
| # --------------------------------------------------------------------------- | |
| FEATURE_PRESETS = { | |
| "Building": { | |
| "prompt": "building", | |
| "min_area": 20.0, "max_area": 50000.0, | |
| "min_compactness": 0.25, "min_rectangularity": 0.5, | |
| "color": "#10b981", | |
| }, | |
| "Rooftop": { | |
| "prompt": "rooftop", | |
| "min_area": 20.0, "max_area": 50000.0, | |
| "min_compactness": 0.25, "min_rectangularity": 0.5, | |
| "color": "#34d399", | |
| }, | |
| "Road": { | |
| "prompt": "road", | |
| "min_area": 10.0, "max_area": 500000.0, | |
| "min_compactness": 0.0, "min_rectangularity": 0.0, | |
| "color": "#f59e0b", | |
| }, | |
| "Waterbody": { | |
| "prompt": "water", | |
| "min_area": 50.0, "max_area": 5000000.0, | |
| "min_compactness": 0.0, "min_rectangularity": 0.0, | |
| "color": "#3b82f6", | |
| }, | |
| "Vegetation": { | |
| "prompt": "tree", | |
| "min_area": 30.0, "max_area": 5000000.0, | |
| "min_compactness": 0.0, "min_rectangularity": 0.0, | |
| "color": "#22c55e", | |
| }, | |
| "Parking Lot": { | |
| "prompt": "parking lot", | |
| "min_area": 100.0, "max_area": 200000.0, | |
| "min_compactness": 0.2, "min_rectangularity": 0.4, | |
| "color": "#8b5cf6", | |
| }, | |
| } | |
| WORLD_EXT_MAP = { | |
| ".png": ".pgw", ".jpg": ".jgw", ".jpeg": ".jgw", | |
| ".tif": ".tfw", ".tiff": ".tfw", | |
| } | |
| # --------------------------------------------------------------------------- | |
| # Model | |
| # --------------------------------------------------------------------------- | |
| def load_model(): | |
| global MODEL, PROCESSOR | |
| if MODEL is None: | |
| from transformers import Sam3Model, Sam3Processor | |
| device = get_device() | |
| print(f"[MODEL] Loading SAM 3 on {device}...") | |
| MODEL = Sam3Model.from_pretrained("facebook/sam3").to(device) | |
| PROCESSOR = Sam3Processor.from_pretrained("facebook/sam3") | |
| print(f"[MODEL] SAM 3 loaded on {device}") | |
| return MODEL, PROCESSOR | |
| # Grounding DINO (used for two-stage detection mode) | |
| GDINO_MODEL = None | |
| GDINO_PROCESSOR = None | |
| def load_grounding_dino(): | |
| global GDINO_MODEL, GDINO_PROCESSOR | |
| if GDINO_MODEL is None: | |
| from transformers import AutoModelForZeroShotObjectDetection, AutoProcessor | |
| device = get_device() | |
| print(f"[GDINO] Loading Grounding DINO on {device}...") | |
| GDINO_MODEL = AutoModelForZeroShotObjectDetection.from_pretrained( | |
| "IDEA-Research/grounding-dino-tiny" | |
| ).to(device) | |
| GDINO_PROCESSOR = AutoProcessor.from_pretrained( | |
| "IDEA-Research/grounding-dino-tiny" | |
| ) | |
| print(f"[GDINO] Grounding DINO loaded on {device}") | |
| return GDINO_MODEL, GDINO_PROCESSOR | |
| # RAMP — aerial building specialist (TFLite, runs on CPU) | |
| # Trained on aerial imagery from 8 LMIC countries (Ghana, India, Malawi, Myanmar, | |
| # Oman, Sierra Leone, South Sudan, St Vincent). Eff-UNet architecture, 256x256 input. | |
| RAMP_MODEL = None | |
| RAMP_MODEL_URL = ( | |
| "https://api-prod.fair.hotosm.org/api/v1/workspace/download/ramp/baseline.tflite" | |
| ) | |
| RAMP_MODEL_PATH = "/tmp/ramp_baseline.tflite" | |
| def load_ramp_model(): | |
| """Download (if needed) and load the RAMP TFLite model. | |
| Runs on CPU via tflite-runtime — no GPU needed, ~50MB dependency. | |
| """ | |
| global RAMP_MODEL | |
| if RAMP_MODEL is not None: | |
| return RAMP_MODEL | |
| # Download model if not cached | |
| if not os.path.exists(RAMP_MODEL_PATH): | |
| print(f"[RAMP] Downloading model from {RAMP_MODEL_URL}...") | |
| import urllib.request | |
| urllib.request.urlretrieve(RAMP_MODEL_URL, RAMP_MODEL_PATH) | |
| size_mb = os.path.getsize(RAMP_MODEL_PATH) / (1024 * 1024) | |
| print(f"[RAMP] Downloaded {size_mb:.1f} MB to {RAMP_MODEL_PATH}") | |
| # Prefer ai-edge-litert (Google's TFLite successor, has Python 3.13 wheels) | |
| # Fall back to tflite_runtime, then full TensorFlow | |
| try: | |
| from ai_edge_litert.interpreter import Interpreter | |
| RAMP_MODEL = Interpreter(model_path=RAMP_MODEL_PATH) | |
| except ImportError: | |
| try: | |
| import tflite_runtime.interpreter as tflite | |
| RAMP_MODEL = tflite.Interpreter(model_path=RAMP_MODEL_PATH) | |
| except ImportError: | |
| try: | |
| import tensorflow as tf | |
| RAMP_MODEL = tf.lite.Interpreter(model_path=RAMP_MODEL_PATH) | |
| except ImportError: | |
| raise RuntimeError( | |
| "Need ai-edge-litert, tflite-runtime, or tensorflow installed. " | |
| "Add ai-edge-litert to requirements.txt" | |
| ) | |
| print(f"[RAMP] Model loaded from {RAMP_MODEL_PATH}") | |
| return RAMP_MODEL | |
| # --------------------------------------------------------------------------- | |
| # Ingest | |
| # --------------------------------------------------------------------------- | |
| def _parse_world_file(world_path: str) -> Affine: | |
| lines = Path(world_path).read_text().strip().splitlines() | |
| if len(lines) < 6: | |
| raise gr.Error(f"World file must have 6 lines, got {len(lines)}") | |
| return Affine( | |
| float(lines[0]), float(lines[2]), float(lines[4]), | |
| float(lines[1]), float(lines[3]), float(lines[5]), | |
| ) | |
| def ingest(image_path: str, world_path: str | None, crs: str) -> str: | |
| tmp = tempfile.mkdtemp(prefix="janus_") | |
| geotiff = os.path.join(tmp, "input.tif") | |
| if image_path.lower().endswith((".tif", ".tiff")): | |
| try: | |
| with rasterio.open(image_path) as src: | |
| if src.crs is not None and src.transform != Affine.identity(): | |
| import shutil | |
| shutil.copy2(image_path, geotiff) | |
| return geotiff | |
| except Exception: | |
| pass | |
| if world_path is None: | |
| img_p = Path(image_path) | |
| expected_ext = WORLD_EXT_MAP.get(img_p.suffix.lower(), ".pgw") | |
| auto_wf = img_p.with_suffix(expected_ext) | |
| if auto_wf.exists(): | |
| world_path = str(auto_wf) | |
| else: | |
| raise gr.Error(f"No world file found. Expected: {auto_wf.name}") | |
| transform = _parse_world_file(world_path) | |
| target_crs = CRS.from_user_input(crs) | |
| img = Image.open(image_path).convert("RGB") | |
| img_array = np.array(img) | |
| h, w, bands = img_array.shape | |
| profile = { | |
| "driver": "GTiff", "dtype": "uint8", | |
| "width": w, "height": h, "count": bands, | |
| "crs": target_crs, "transform": transform, | |
| } | |
| with rasterio.open(geotiff, "w", **profile) as dst: | |
| for b in range(bands): | |
| dst.write(img_array[:, :, b], b + 1) | |
| return geotiff | |
| # --------------------------------------------------------------------------- | |
| # Tiling | |
| # --------------------------------------------------------------------------- | |
| def compute_tile_windows(img_w: int, img_h: int, tile_size: int, overlap: int) -> list[Window]: | |
| step = tile_size - overlap | |
| windows = [] | |
| for y in range(0, img_h, step): | |
| for x in range(0, img_w, step): | |
| tw = min(tile_size, img_w - x) | |
| th = min(tile_size, img_h - y) | |
| if tw < tile_size // 4 or th < tile_size // 4: | |
| continue | |
| windows.append(Window(x, y, tw, th)) | |
| return windows | |
| # --------------------------------------------------------------------------- | |
| # Per-tile segmentation | |
| # --------------------------------------------------------------------------- | |
| def segment_tile(tile_rgb: np.ndarray, prompt: str, confidence: float): | |
| """Single-stage: SAM 3 with text prompt only.""" | |
| model, processor = load_model() | |
| device = get_device() | |
| image = Image.fromarray(tile_rgb) | |
| inputs = processor(images=image, text=prompt, return_tensors="pt").to(device) | |
| with torch.no_grad(): | |
| outputs = model(**inputs) | |
| results = processor.post_process_instance_segmentation( | |
| outputs, threshold=confidence, mask_threshold=0.5, | |
| target_sizes=inputs.get("original_sizes").tolist(), | |
| )[0] | |
| masks = results["masks"].cpu().numpy() | |
| scores = results["scores"].cpu().numpy() | |
| print( | |
| f"[SAM3-1stage] prompt='{prompt}' threshold={confidence} -> " | |
| f"{len(masks)} masks (scores: {scores.tolist() if len(scores) < 20 else f'min={float(scores.min()):.2f} max={float(scores.max()):.2f}'})" | |
| ) | |
| return masks, scores | |
| def segment_tile_two_stage( | |
| tile_rgb: np.ndarray, | |
| prompt: str, | |
| confidence: float, | |
| box_threshold: float = 0.15, | |
| text_threshold: float = 0.15, | |
| ): | |
| """Two-stage: Grounding DINO -> bounding boxes -> SAM 3 box-prompt segmentation.""" | |
| device = get_device() | |
| image = Image.fromarray(tile_rgb) | |
| # Stage 1: Grounding DINO -> boxes | |
| gdino_model, gdino_processor = load_grounding_dino() | |
| # Aerial imagery benefits from multi-prompt phrasing | |
| if prompt.lower().strip() == "building": | |
| gdino_prompt = "building. rooftop. house. roof." | |
| else: | |
| gdino_prompt = prompt.lower().strip().rstrip(".") + "." | |
| gdino_inputs = gdino_processor( | |
| images=image, text=gdino_prompt, return_tensors="pt" | |
| ).to(device) | |
| with torch.no_grad(): | |
| gdino_outputs = gdino_model(**gdino_inputs) | |
| try: | |
| gdino_results = gdino_processor.post_process_grounded_object_detection( | |
| gdino_outputs, | |
| gdino_inputs.input_ids, | |
| threshold=box_threshold, | |
| text_threshold=text_threshold, | |
| target_sizes=[image.size[::-1]], | |
| )[0] | |
| except TypeError: | |
| # Older transformers used box_threshold instead of threshold | |
| gdino_results = gdino_processor.post_process_grounded_object_detection( | |
| gdino_outputs, | |
| gdino_inputs.input_ids, | |
| box_threshold=box_threshold, | |
| text_threshold=text_threshold, | |
| target_sizes=[image.size[::-1]], | |
| )[0] | |
| boxes = gdino_results["boxes"].cpu().numpy() | |
| box_scores = gdino_results["scores"].cpu().numpy() | |
| print( | |
| f"[GDINO] prompt='{gdino_prompt}' threshold={box_threshold} -> " | |
| f"{len(boxes)} boxes (scores: {box_scores.tolist() if len(box_scores) < 20 else f'min={box_scores.min():.2f} max={box_scores.max():.2f}'})" | |
| ) | |
| if len(boxes) == 0: | |
| return np.array([]), np.array([]) | |
| # Stage 2: SAM 3 in box-prompt mode -> masks per box | |
| sam_model, sam_processor = load_model() | |
| input_boxes = [boxes.tolist()] | |
| input_labels = [[1] * len(boxes)] | |
| sam_inputs = sam_processor( | |
| images=image, | |
| input_boxes=input_boxes, | |
| input_boxes_labels=input_labels, | |
| return_tensors="pt", | |
| ).to(device) | |
| with torch.no_grad(): | |
| sam_outputs = sam_model(**sam_inputs) | |
| sam_results = sam_processor.post_process_instance_segmentation( | |
| sam_outputs, | |
| threshold=confidence, | |
| mask_threshold=0.5, | |
| target_sizes=sam_inputs.get("original_sizes").tolist(), | |
| )[0] | |
| masks = sam_results["masks"].cpu().numpy() | |
| sam_scores = sam_results["scores"].cpu().numpy() | |
| # Joint confidence: detection score (DINO) * mask score (SAM 3) | |
| if len(sam_scores) == len(box_scores): | |
| scores = sam_scores * box_scores | |
| else: | |
| scores = sam_scores | |
| return masks, scores | |
| def segment_tile_ramp(tile_rgb: np.ndarray, confidence: float = 0.5): | |
| """RAMP aerial-buildings specialist with watershed post-processing. | |
| RAMP outputs a 4-channel mask: [background, building, boundary, close_contact]. | |
| - Channel 0: background probability | |
| - Channel 1: building interior probability | |
| - Channel 2: building boundary (edges around each building) | |
| - Channel 3: close_contact (where two buildings touch) | |
| We use all four channels: | |
| 1. Build a "core" mask of high-confidence interior pixels (subtract boundary | |
| and close_contact from building) -- these become watershed seeds. | |
| 2. Build a "full extent" mask of the building footprint at a lower threshold. | |
| 3. Run watershed: each seed grows into the full extent, separated by ridges | |
| at the boundary/contact channels. | |
| 4. Each watershed region becomes one building. | |
| This properly separates touching buildings instead of merging them into | |
| a single connected component. | |
| Args: | |
| tile_rgb: HxWx3 uint8 numpy array. | |
| confidence: Probability threshold for building interior (0-1). | |
| Returns: | |
| (masks, scores) tuple matching segment_tile() interface. | |
| """ | |
| interpreter = load_ramp_model() | |
| h, w = tile_rgb.shape[:2] | |
| # Resize to 256x256 for the model | |
| if (h, w) != (256, 256): | |
| from PIL import Image as PILImage | |
| pil = PILImage.fromarray(tile_rgb).resize((256, 256), PILImage.BILINEAR) | |
| tile_resized = np.array(pil) | |
| else: | |
| tile_resized = tile_rgb | |
| # Prepare input tensor — RAMP expects pixels normalized to [0, 1] | |
| interpreter.resize_tensor_input( | |
| interpreter.get_input_details()[0]["index"], (1, 256, 256, 3) | |
| ) | |
| interpreter.allocate_tensors() | |
| input_idx = interpreter.get_input_details()[0]["index"] | |
| output_idx = interpreter.get_output_details()[0]["index"] | |
| inp = (tile_resized.astype(np.float32) / 255.0).reshape(1, 256, 256, 3) | |
| interpreter.set_tensor(input_idx, inp) | |
| interpreter.invoke() | |
| output = interpreter.get_tensor(output_idx) # (1, 256, 256, num_channels) | |
| print( | |
| f"[RAMP] Output shape={output.shape}, " | |
| f"min={float(output.min()):.3f}, max={float(output.max()):.3f}" | |
| ) | |
| # Decompose channels — handle both 4-channel (full RAMP) and 2-channel fallback | |
| if output.shape[-1] >= 4: | |
| building_prob = output[0, :, :, 1] | |
| boundary_prob = output[0, :, :, 2] | |
| contact_prob = output[0, :, :, 3] | |
| elif output.shape[-1] >= 2: | |
| building_prob = output[0, :, :, 1] | |
| boundary_prob = np.zeros_like(building_prob) | |
| contact_prob = np.zeros_like(building_prob) | |
| else: | |
| building_prob = output[0, :, :, 0] | |
| boundary_prob = np.zeros_like(building_prob) | |
| contact_prob = np.zeros_like(building_prob) | |
| print( | |
| f"[RAMP] Channel maxes — building={float(building_prob.max()):.3f}, " | |
| f"boundary={float(boundary_prob.max()):.3f}, " | |
| f"contact={float(contact_prob.max()):.3f}" | |
| ) | |
| # Build full extent mask (lower threshold so we don't lose building edges) | |
| full_extent_thresh = max(0.3, confidence * 0.6) | |
| full_extent = building_prob > full_extent_thresh | |
| if full_extent.sum() == 0: | |
| print( | |
| f"[RAMP] No buildings above full_extent threshold={full_extent_thresh:.2f} " | |
| f"(max prob={float(building_prob.max()):.3f})" | |
| ) | |
| return np.array([]), np.array([]) | |
| # Build "interior cores" by subtracting boundary + contact from building probability | |
| # These act as seeds for watershed — high-confidence centers of each building | |
| cleaned_building = building_prob * (1.0 - boundary_prob) * (1.0 - contact_prob) | |
| seeds = cleaned_building > confidence | |
| # Morphological cleanup on seeds — remove tiny noise, ensure separation | |
| from scipy import ndimage | |
| seeds = ndimage.binary_opening(seeds, iterations=1) | |
| # Label seeds — each becomes a watershed marker | |
| markers, num_seeds = ndimage.label(seeds) | |
| print(f"[RAMP] Watershed seeds={num_seeds} (after morphological opening)") | |
| if num_seeds == 0: | |
| # Fall back to plain connected components on full_extent | |
| labeled, num = ndimage.label(full_extent) | |
| instance_labels = labeled | |
| else: | |
| # Watershed: each seed expands into the full extent, ridges at boundary/contact | |
| try: | |
| from skimage.segmentation import watershed | |
| distance = ndimage.distance_transform_edt(full_extent) | |
| # Use negative distance so watershed flows from cores outward | |
| instance_labels = watershed( | |
| -distance, | |
| markers=markers, | |
| mask=full_extent, | |
| ) | |
| num = int(instance_labels.max()) | |
| except ImportError: | |
| # Fallback if scikit-image not installed | |
| print("[RAMP] scikit-image not available, falling back to connected components") | |
| labeled, num = ndimage.label(full_extent) | |
| instance_labels = labeled | |
| print(f"[RAMP] After watershed: {num} buildings") | |
| if num == 0: | |
| return np.array([]), np.array([]) | |
| # Resize instance labels and probability map back to tile size | |
| if (h, w) != (256, 256): | |
| from PIL import Image as PILImage | |
| # Use NEAREST for labels to preserve discrete IDs | |
| labels_pil = PILImage.fromarray(instance_labels.astype(np.int32)).resize( | |
| (w, h), PILImage.NEAREST | |
| ) | |
| instance_labels = np.array(labels_pil) | |
| prob_pil = PILImage.fromarray((building_prob * 255).astype(np.uint8)).resize( | |
| (w, h), PILImage.BILINEAR | |
| ) | |
| building_prob = np.array(prob_pil).astype(np.float32) / 255.0 | |
| # Build per-building masks and scores | |
| masks = np.zeros((num, h, w), dtype=bool) | |
| scores = np.zeros(num, dtype=np.float32) | |
| keep_idx = [] | |
| for i in range(num): | |
| component = instance_labels == (i + 1) | |
| if component.sum() < 4: # Drop tiny noise (< 4 pixels) | |
| continue | |
| masks[i] = component | |
| scores[i] = float(building_prob[component].mean()) | |
| keep_idx.append(i) | |
| if not keep_idx: | |
| return np.array([]), np.array([]) | |
| masks = masks[keep_idx] | |
| scores = scores[keep_idx] | |
| num = len(keep_idx) | |
| print( | |
| f"[RAMP] Found {num} buildings (max prob={float(building_prob.max()):.3f}, " | |
| f"mean score={float(scores.mean()):.3f})" | |
| ) | |
| return masks, scores | |
| # --------------------------------------------------------------------------- | |
| # Shape metrics | |
| # --------------------------------------------------------------------------- | |
| def _compactness(geom): | |
| if geom.is_empty or geom.length == 0: | |
| return 0.0 | |
| return (4.0 * math.pi * geom.area) / (geom.length ** 2) | |
| def _rectangularity(geom): | |
| if geom.is_empty: | |
| return 0.0 | |
| mrr = geom.minimum_rotated_rectangle | |
| return geom.area / mrr.area if mrr.area > 0 else 0.0 | |
| # --------------------------------------------------------------------------- | |
| # Dedup + filter | |
| # --------------------------------------------------------------------------- | |
| def merge_and_deduplicate(features: list[dict], crs_obj, iou_thresh: float = 0.5) -> gpd.GeoDataFrame: | |
| if not features: | |
| return gpd.GeoDataFrame(columns=["feature_id", "feature_type", "confidence", "geometry"]) | |
| gdf = gpd.GeoDataFrame(features, crs=crs_obj) | |
| if len(gdf) == 0: | |
| return gdf | |
| sindex = gdf.sindex | |
| drop = set() | |
| for idx, row in gdf.iterrows(): | |
| if idx in drop: | |
| continue | |
| for cand_idx in sindex.intersection(row.geometry.bounds): | |
| if cand_idx <= idx or cand_idx in drop: | |
| continue | |
| cand = gdf.loc[cand_idx] | |
| if row["feature_type"] != cand["feature_type"]: | |
| continue | |
| if not row.geometry.intersects(cand.geometry): | |
| continue | |
| try: | |
| inter = row.geometry.intersection(cand.geometry).area | |
| union = row.geometry.union(cand.geometry).area | |
| if union > 0 and inter / union >= iou_thresh: | |
| if row["confidence"] >= cand["confidence"]: | |
| drop.add(cand_idx) | |
| else: | |
| drop.add(idx) | |
| break | |
| except Exception: | |
| continue | |
| gdf = gdf.drop(index=drop).reset_index(drop=True) | |
| gdf["feature_id"] = range(1, len(gdf) + 1) | |
| return gdf | |
| def _orthogonalize_polygon(geom, dominant_angle_deg: float, snap_threshold_deg: float = 15.0): | |
| """Snap polygon edges to dominant orientation (90-degree corners). | |
| Algorithm: | |
| 1. Rotate polygon by -dominant_angle so the dominant orientation is axis-aligned. | |
| 2. For each vertex, compute the angle of the incoming and outgoing edges. | |
| 3. If the edge is nearly axis-aligned (within snap_threshold_deg), snap | |
| the vertex's coordinates to align with the previous vertex along the | |
| relevant axis. | |
| 4. Rotate back. | |
| Returns the orthogonalized polygon, or None if it can't be cleanly snapped. | |
| """ | |
| from shapely.affinity import rotate | |
| from shapely.geometry import Polygon | |
| centroid = geom.centroid | |
| rotated = rotate(geom, -dominant_angle_deg, origin=centroid, use_radians=False) | |
| coords = list(rotated.exterior.coords) | |
| if len(coords) < 4: | |
| return None | |
| # Snap each consecutive vertex pair to be axis-aligned if the edge angle | |
| # is within snap_threshold of horizontal/vertical. | |
| new_coords = [coords[0]] | |
| for i in range(1, len(coords)): | |
| x_prev, y_prev = new_coords[-1] | |
| x_cur, y_cur = coords[i] | |
| dx = x_cur - x_prev | |
| dy = y_cur - y_prev | |
| if dx == 0 and dy == 0: | |
| continue | |
| angle = math.degrees(math.atan2(dy, dx)) | |
| # Compute angle delta from horizontal (0/180) or vertical (90/-90) | |
| from_horiz = min(abs(angle), abs(angle - 180), abs(angle + 180)) | |
| from_vert = min(abs(angle - 90), abs(angle + 90)) | |
| if from_horiz < snap_threshold_deg: | |
| # Snap to horizontal — keep y_prev | |
| new_coords.append((x_cur, y_prev)) | |
| elif from_vert < snap_threshold_deg: | |
| # Snap to vertical — keep x_prev | |
| new_coords.append((x_prev, y_cur)) | |
| else: | |
| # Diagonal edge — leave as-is | |
| new_coords.append((x_cur, y_cur)) | |
| # Ensure closure | |
| if new_coords[0] != new_coords[-1]: | |
| new_coords.append(new_coords[0]) | |
| if len(new_coords) < 4: | |
| return None | |
| try: | |
| snapped = Polygon(new_coords) | |
| if not snapped.is_valid: | |
| snapped = snapped.buffer(0) | |
| if snapped.is_empty or snapped.area < 1.0: | |
| return None | |
| # Rotate back | |
| return rotate(snapped, dominant_angle_deg, origin=centroid, use_radians=False) | |
| except Exception: | |
| return None | |
| def _dominant_angle(geom) -> float: | |
| """Get the dominant orientation angle (in degrees) of a polygon. | |
| Uses the longest edge of the minimum rotated rectangle. | |
| """ | |
| mrr = geom.minimum_rotated_rectangle | |
| coords = list(mrr.exterior.coords) | |
| if len(coords) < 5: | |
| return 0.0 | |
| # MRR has 5 coords (4 corners + closing), edges between consecutive corners | |
| edges = [ | |
| (coords[i], coords[i + 1], math.dist(coords[i], coords[i + 1])) | |
| for i in range(4) | |
| ] | |
| longest = max(edges, key=lambda e: e[2]) | |
| (x0, y0), (x1, y1), _ = longest | |
| return math.degrees(math.atan2(y1 - y0, x1 - x0)) | |
| def simplify_and_orthogonalize( | |
| gdf: gpd.GeoDataFrame, | |
| simplify_tolerance_m: float = 0.5, | |
| orthogonalize: bool = True, | |
| rect_iou_threshold: float = 0.85, | |
| ) -> gpd.GeoDataFrame: | |
| """Clean polygon geometries: simplify, snap to right-angle footprints. | |
| Three-step cleanup: | |
| 1. Douglas-Peucker simplify to remove redundant pixel-edge vertices. | |
| 2. If polygon's IoU with its minimum rotated rectangle is high (>0.85), | |
| replace it with the rectangle — clean rectangular footprint. | |
| 3. Otherwise, snap edges to the dominant orientation (90-degree corners) | |
| while preserving L/T/U shapes. | |
| Args: | |
| gdf: GeoDataFrame of building polygons. | |
| simplify_tolerance_m: Simplification tolerance in meters. | |
| orthogonalize: If False, only simplify. | |
| rect_iou_threshold: IoU above which a polygon is replaced by its MRR. | |
| """ | |
| if len(gdf) == 0: | |
| return gdf | |
| # Project to metric CRS so tolerance and angles are stable | |
| if gdf.crs and gdf.crs.is_geographic: | |
| utm_crs = gdf.estimate_utm_crs() | |
| gdf_proj = gdf.to_crs(utm_crs) | |
| else: | |
| gdf_proj = gdf.copy() | |
| n = len(gdf_proj) | |
| print( | |
| f"[CLEAN] Simplifying {n} polygons " | |
| f"(tolerance={simplify_tolerance_m}m, orthogonalize={orthogonalize})" | |
| ) | |
| new_geoms = [] | |
| rect_replaced = 0 | |
| ortho_applied = 0 | |
| simplified_only = 0 | |
| for geom in gdf_proj.geometry: | |
| if geom.is_empty: | |
| new_geoms.append(geom) | |
| continue | |
| # Step 1: Douglas-Peucker simplification | |
| simplified = geom.simplify(simplify_tolerance_m, preserve_topology=True) | |
| if simplified.is_empty: | |
| new_geoms.append(geom) | |
| continue | |
| # Step 2: If shape is basically a rectangle, replace with MRR | |
| mrr = simplified.minimum_rotated_rectangle | |
| if mrr.area > 0: | |
| iou_with_rect = simplified.intersection(mrr).area / simplified.union(mrr).area | |
| else: | |
| iou_with_rect = 0 | |
| if iou_with_rect > rect_iou_threshold: | |
| new_geoms.append(mrr) | |
| rect_replaced += 1 | |
| continue | |
| if not orthogonalize: | |
| new_geoms.append(simplified) | |
| simplified_only += 1 | |
| continue | |
| # Step 3: Orthogonalize — snap edges to right angles around dominant orientation | |
| try: | |
| angle = _dominant_angle(simplified) | |
| ortho_geom = _orthogonalize_polygon(simplified, angle, snap_threshold_deg=20.0) | |
| if ortho_geom is not None and ortho_geom.is_valid and not ortho_geom.is_empty: | |
| # Sanity check — the orthogonalized version shouldn't deviate too far from original | |
| if ortho_geom.intersection(simplified).area / max(ortho_geom.union(simplified).area, 1e-6) > 0.5: | |
| new_geoms.append(ortho_geom) | |
| ortho_applied += 1 | |
| else: | |
| new_geoms.append(simplified) | |
| simplified_only += 1 | |
| else: | |
| new_geoms.append(simplified) | |
| simplified_only += 1 | |
| except Exception: | |
| new_geoms.append(simplified) | |
| simplified_only += 1 | |
| gdf_proj["geometry"] = new_geoms | |
| print( | |
| f"[CLEAN] Done: {rect_replaced} replaced with rectangles, " | |
| f"{ortho_applied} orthogonalized, {simplified_only} simplified only" | |
| ) | |
| # Reproject back to original CRS if needed | |
| if gdf.crs and gdf.crs.is_geographic: | |
| gdf_proj = gdf_proj.to_crs(gdf.crs) | |
| return gdf_proj | |
| def filter_features(gdf: gpd.GeoDataFrame, preset: dict) -> gpd.GeoDataFrame: | |
| if len(gdf) == 0: | |
| print(f"[FILTER] preset={preset.get('prompt')}: input is empty") | |
| return gdf | |
| n0 = len(gdf) | |
| if gdf.crs and gdf.crs.is_geographic: | |
| proj = gdf.to_crs(gdf.estimate_utm_crs()) | |
| gdf["area_m2"] = proj.geometry.area | |
| pg = proj.geometry | |
| else: | |
| gdf["area_m2"] = gdf.geometry.area | |
| pg = gdf.geometry | |
| print( | |
| f"[FILTER] preset='{preset.get('prompt')}' input={n0} | " | |
| f"area_m2 range: {gdf['area_m2'].min():.1f} - {gdf['area_m2'].max():.1f} | " | |
| f"thresholds: area>={preset['min_area']}, compact>={preset['min_compactness']}, " | |
| f"rect>={preset['min_rectangularity']}" | |
| ) | |
| gdf = gdf[(gdf["area_m2"] >= preset["min_area"]) & (gdf["area_m2"] <= preset["max_area"])].copy() | |
| print(f"[FILTER] after AREA filter: {len(gdf)} (dropped {n0 - len(gdf)})") | |
| if len(gdf) == 0: | |
| return gdf | |
| n1 = len(gdf) | |
| if gdf.crs and gdf.crs.is_geographic: | |
| pg = gdf.to_crs(gdf.estimate_utm_crs()).geometry | |
| else: | |
| pg = gdf.geometry | |
| gdf["compactness"] = pg.apply(_compactness) | |
| print( | |
| f"[FILTER] compactness range: {gdf['compactness'].min():.3f} - {gdf['compactness'].max():.3f}" | |
| ) | |
| if preset["min_compactness"] > 0: | |
| gdf = gdf[gdf["compactness"] >= preset["min_compactness"]].copy() | |
| print(f"[FILTER] after COMPACTNESS filter: {len(gdf)} (dropped {n1 - len(gdf)})") | |
| if len(gdf) == 0: | |
| return gdf | |
| n2 = len(gdf) | |
| if preset["min_rectangularity"] > 0: | |
| if gdf.crs and gdf.crs.is_geographic: | |
| pg = gdf.to_crs(gdf.estimate_utm_crs()).geometry | |
| else: | |
| pg = gdf.geometry | |
| gdf["rectangularity"] = pg.apply(_rectangularity) | |
| print( | |
| f"[FILTER] rectangularity range: {gdf['rectangularity'].min():.3f} - {gdf['rectangularity'].max():.3f}" | |
| ) | |
| gdf = gdf[gdf["rectangularity"] >= preset["min_rectangularity"]].copy() | |
| print(f"[FILTER] after RECTANGULARITY filter: {len(gdf)} (dropped {n2 - len(gdf)})") | |
| print(f"[FILTER] FINAL: {len(gdf)} features kept out of {n0} input") | |
| return gdf | |
| # --------------------------------------------------------------------------- | |
| # Export | |
| # --------------------------------------------------------------------------- | |
| def export_all(features_by_type: dict[str, gpd.GeoDataFrame], tmp_dir: str) -> list[str]: | |
| paths = [] | |
| for feat_type, gdf in features_by_type.items(): | |
| if len(gdf) == 0: | |
| continue | |
| name = feat_type.lower().replace(" ", "_") | |
| gpkg = os.path.join(tmp_dir, f"{name}.gpkg") | |
| gdf.to_file(gpkg, driver="GPKG") | |
| paths.append(gpkg) | |
| gj = os.path.join(tmp_dir, f"{name}.geojson") | |
| gdf.to_file(gj, driver="GeoJSON") | |
| paths.append(gj) | |
| wkt = os.path.join(tmp_dir, f"{name}.wkt") | |
| with open(wkt, "w") as f: | |
| for idx, row in enumerate(gdf.itertuples(), start=1): | |
| c = getattr(row, "confidence", 0.0) | |
| f.write(f"{idx}|{c:.3f}|{row.geometry.wkt}\n") | |
| paths.append(wkt) | |
| return paths | |
| # --------------------------------------------------------------------------- | |
| # Overlay | |
| # --------------------------------------------------------------------------- | |
| def make_overlay( | |
| geotiff: str, | |
| features_by_type: dict[str, gpd.GeoDataFrame], | |
| scanned_bounds: tuple | None = None, | |
| ) -> str: | |
| """Render before/after overlay with optional scan-progress box. | |
| Args: | |
| geotiff: Path to the source GeoTIFF. | |
| features_by_type: Dict of feature_type -> GeoDataFrame. | |
| scanned_bounds: (left, bottom, right, top) in CRS coordinates showing | |
| the area processed so far. Drawn as a dashed box. | |
| """ | |
| fig, axes = plt.subplots(1, 2, figsize=(16, 7), dpi=100) | |
| fig.patch.set_facecolor("#fafafa") | |
| with rasterio.open(geotiff) as src: | |
| max_dim = 2000 | |
| scale = min(1.0, max_dim / max(src.width, src.height)) | |
| out_w = int(src.width * scale) | |
| out_h = int(src.height * scale) | |
| data = src.read( | |
| out_shape=(src.count, out_h, out_w), | |
| resampling=rasterio.enums.Resampling.bilinear, | |
| ) | |
| extent = [src.bounds.left, src.bounds.right, src.bounds.bottom, src.bounds.top] | |
| rgb = np.transpose(data[:3], (1, 2, 0)) | |
| # Left panel: input image | |
| axes[0].imshow(rgb, extent=extent) | |
| axes[0].set_title("Input", fontsize=14, fontweight=600, color="#18181b", pad=12) | |
| axes[0].tick_params(labelsize=7, colors="#71717a") | |
| # Right panel: image + detected features | |
| axes[1].imshow(rgb, extent=extent) | |
| # Draw scan progress box | |
| if scanned_bounds is not None: | |
| from matplotlib.patches import Rectangle | |
| left, bottom, right, top = scanned_bounds | |
| rect = Rectangle( | |
| (left, bottom), right - left, top - bottom, | |
| linewidth=1.5, edgecolor="#71717a", facecolor="none", | |
| linestyle="--", alpha=0.6, | |
| ) | |
| axes[1].add_patch(rect) | |
| # Plot features color-coded by type | |
| total = 0 | |
| legend_items = [] | |
| for feat_type, gdf in features_by_type.items(): | |
| if len(gdf) == 0: | |
| continue | |
| color = FEATURE_PRESETS.get(feat_type, {}).get("color", "#ef4444") | |
| gdf.plot(ax=axes[1], edgecolor=color, facecolor=color, alpha=0.35, linewidth=1) | |
| total += len(gdf) | |
| legend_items.append(f"{feat_type}: {len(gdf)}") | |
| title = f"{total} Features" | |
| if legend_items: | |
| title += f" ({', '.join(legend_items)})" | |
| axes[1].set_title(title, fontsize=12, fontweight=600, color="#18181b", pad=12) | |
| axes[1].tick_params(labelsize=7, colors="#71717a") | |
| for ax in axes: | |
| for spine in ax.spines.values(): | |
| spine.set_color("#e4e4e7") | |
| plt.tight_layout(pad=2) | |
| out = os.path.join(os.path.dirname(geotiff), "overlay.png") | |
| plt.savefig(out, bbox_inches="tight", facecolor="#fafafa") | |
| plt.close() | |
| return out | |
| # --------------------------------------------------------------------------- | |
| # Dashboard formatting | |
| # --------------------------------------------------------------------------- | |
| class ConfidenceTracker: | |
| """Track running confidence stats per feature type.""" | |
| def __init__(self): | |
| self.scores: dict[str, list[float]] = defaultdict(list) | |
| def add(self, feat_type: str, score: float): | |
| self.scores[feat_type].append(score) | |
| def count(self, feat_type: str) -> int: | |
| return len(self.scores[feat_type]) | |
| def total(self) -> int: | |
| return sum(len(v) for v in self.scores.values()) | |
| def stats(self, feat_type: str) -> dict: | |
| s = self.scores[feat_type] | |
| if not s: | |
| return {"count": 0, "avg": 0, "min": 0, "max": 0, "low": 0} | |
| return { | |
| "count": len(s), | |
| "avg": sum(s) / len(s), | |
| "min": min(s), | |
| "max": max(s), | |
| "low": sum(1 for x in s if x < 0.5), | |
| } | |
| def format_dashboard( | |
| image_name: str, | |
| img_w: int, | |
| img_h: int, | |
| tile_idx: int, | |
| total_tiles: int, | |
| elapsed: float, | |
| tracker: ConfidenceTracker, | |
| feature_types: list[str], | |
| status: str = "Processing", | |
| ) -> str: | |
| pct = (tile_idx / total_tiles * 100) if total_tiles > 0 else 0 | |
| elapsed_str = _fmt_time(elapsed) | |
| if tile_idx > 0 and tile_idx < total_tiles: | |
| rate = elapsed / tile_idx | |
| remaining = rate * (total_tiles - tile_idx) | |
| remaining_str = _fmt_time(remaining) | |
| else: | |
| remaining_str = "--" | |
| md = f"**{status}: {image_name}** ({img_w:,} x {img_h:,} px)\n\n" | |
| md += f"Tile {tile_idx} / {total_tiles} | {pct:.0f}% | " | |
| md += f"Elapsed: {elapsed_str} | Remaining: ~{remaining_str}\n\n" | |
| md += "| Feature | Count | Avg Conf | Min Conf | Max Conf | Low (<0.5) |\n" | |
| md += "|---------|-------|----------|----------|----------|------------|\n" | |
| for ft in feature_types: | |
| st = tracker.stats(ft) | |
| if st["count"] > 0: | |
| md += ( | |
| f"| {ft} | {st['count']} | {st['avg']:.2f} | " | |
| f"{st['min']:.2f} | {st['max']:.2f} | {st['low']} |\n" | |
| ) | |
| else: | |
| md += f"| {ft} | 0 | -- | -- | -- | -- |\n" | |
| return md | |
| def format_final_summary( | |
| tracker: ConfidenceTracker, | |
| feature_types: list[str], | |
| final_counts: dict[str, int], | |
| elapsed: float, | |
| ) -> str: | |
| md = f"**Extraction Complete** | Total time: {_fmt_time(elapsed)}\n\n" | |
| md += "| Feature | Raw | Final | Avg Conf | Min | Max | Low (<0.5) |\n" | |
| md += "|---------|-----|-------|----------|-----|-----|------------|\n" | |
| for ft in feature_types: | |
| st = tracker.stats(ft) | |
| final = final_counts.get(ft, 0) | |
| if st["count"] > 0: | |
| md += ( | |
| f"| {ft} | {st['count']} | {final} | " | |
| f"{st['avg']:.2f} | {st['min']:.2f} | {st['max']:.2f} | {st['low']} |\n" | |
| ) | |
| else: | |
| md += f"| {ft} | 0 | 0 | -- | -- | -- | -- |\n" | |
| md += "\n*Adjust the confidence threshold and re-run to improve results.*" | |
| return md | |
| def _fmt_time(secs: float) -> str: | |
| if secs < 60: | |
| return f"{secs:.0f}s" | |
| m, s = divmod(int(secs), 60) | |
| return f"{m}m {s:02d}s" | |
| # --------------------------------------------------------------------------- | |
| # Main pipeline (generator for live updates) | |
| # --------------------------------------------------------------------------- | |
| PREVIEW_INTERVAL = 10 # update overlay every N tiles | |
| CHECKPOINT_INTERVAL = 100 # save downloadable files every N tiles | |
| def run_pipeline( | |
| image_file, | |
| world_file, | |
| crs: str, | |
| feature_types: list[str], | |
| confidence: float, | |
| tile_size: int, | |
| detection_mode: str = "SAM 3 (text prompt only)", | |
| gdino_threshold: float = 0.15, | |
| # Experimental flags (default off — pipeline behaves identically to before) | |
| use_yolo_seg: bool = False, | |
| yolo_model_path: str = "yolov8n-seg.pt", | |
| yolo_conf: float = 0.25, | |
| yolo_iou: float = 0.45, | |
| use_sam2_refinement: bool = False, | |
| sam2_model_size: str = "large", | |
| sam2_prompt_mode: str = "box", | |
| ): | |
| """Generator: yields (overlay_image, dashboard_md, files) after each tile.""" | |
| if image_file is None: | |
| raise gr.Error("Please upload an aerial image.") | |
| if not feature_types: | |
| raise gr.Error("Select at least one feature type.") | |
| start_time = time.time() | |
| tracker = ConfidenceTracker() | |
| all_features: list[dict] = [] | |
| tmp_dir = tempfile.mkdtemp(prefix="janus_") | |
| # -- Ingest -- | |
| yield None, "**Ingesting image...**", None | |
| geotiff = ingest(image_file, world_file, crs) | |
| # Model loads on first segment_tile call (inside @spaces.GPU context) | |
| yield None, "**Starting extraction...** (model loads on first tile)", None | |
| # -- Compute tiles -- | |
| with rasterio.open(geotiff) as src: | |
| img_w, img_h = src.width, src.height | |
| crs_obj = src.crs | |
| transform = src.transform | |
| image_name = Path(image_file).name | |
| overlap = tile_size // 8 | |
| windows = compute_tile_windows(img_w, img_h, tile_size, overlap) | |
| total_tiles = len(windows) | |
| dash = format_dashboard(image_name, img_w, img_h, 0, total_tiles, 0, tracker, feature_types, "Starting") | |
| yield None, dash, None | |
| # -- Process tiles -- | |
| # Track the scanned region for progress visualization | |
| scan_left = scan_bottom = float("inf") | |
| scan_right = scan_top = float("-inf") | |
| with rasterio.open(geotiff) as src: | |
| for tile_idx, window in enumerate(windows): | |
| tile_data = src.read([1, 2, 3], window=window) | |
| tile_rgb = np.transpose(tile_data, (1, 2, 0)) | |
| tile_transform = rasterio.windows.transform(window, transform) | |
| # Update scanned bounds | |
| tile_bounds = rasterio.windows.bounds(window, transform) | |
| scan_left = min(scan_left, tile_bounds[0]) | |
| scan_bottom = min(scan_bottom, tile_bounds[1]) | |
| scan_right = max(scan_right, tile_bounds[2]) | |
| scan_top = max(scan_top, tile_bounds[3]) | |
| for ft in feature_types: | |
| preset = FEATURE_PRESETS[ft] | |
| print( | |
| f"[TILE {tile_idx + 1}/{total_tiles}] feat='{ft}' " | |
| f"prompt='{preset['prompt']}' mode='{detection_mode}' " | |
| f"conf_threshold={confidence}" | |
| ) | |
| # ---- Backend dispatch ---- | |
| weak_boxes = None # populated if a backend produces them (for SAM2) | |
| if detection_mode == "YOLO segmentation (experimental)": | |
| yolo_cfg = YoloConfig( | |
| model_path=yolo_model_path, | |
| conf=yolo_conf, | |
| iou=yolo_iou, | |
| use_as_detector_for_sam2=use_sam2_refinement, | |
| ) | |
| yolo = _get_yolo_backend(yolo_cfg) | |
| if yolo is None: | |
| print( | |
| f"[TILE {tile_idx + 1}/{total_tiles}] YOLO unavailable, " | |
| f"falling back to SAM 3" | |
| ) | |
| masks, scores = segment_tile(tile_rgb, preset["prompt"], confidence) | |
| else: | |
| result = yolo.segment_tile(tile_rgb, ft, confidence) | |
| masks = result.masks if len(result) > 0 else np.array([]) | |
| scores = result.scores | |
| weak_boxes = result.boxes | |
| elif detection_mode == "Grounding DINO + SAM 3 (two-stage)": | |
| masks, scores = segment_tile_two_stage( | |
| tile_rgb, preset["prompt"], confidence, | |
| box_threshold=gdino_threshold, | |
| text_threshold=gdino_threshold, | |
| ) | |
| elif detection_mode == "RAMP (aerial buildings specialist)": | |
| # RAMP only does buildings — skip non-building feature types | |
| if ft not in ("Building", "Rooftop"): | |
| print( | |
| f"[TILE {tile_idx + 1}/{total_tiles}] {ft}: " | |
| f"RAMP is buildings-only, skipping this feature type" | |
| ) | |
| continue | |
| masks, scores = segment_tile_ramp(tile_rgb, confidence) | |
| else: | |
| masks, scores = segment_tile( | |
| tile_rgb, preset["prompt"], confidence | |
| ) | |
| # ---- Optional SAM 2 refinement (experimental) ---- | |
| if ( | |
| use_sam2_refinement | |
| and len(masks) > 0 | |
| and detection_mode != "YOLO segmentation (experimental)" | |
| or (use_sam2_refinement and weak_boxes is not None and len(masks) > 0) | |
| ): | |
| sam2_cfg = Sam2Config( | |
| enabled=True, | |
| model_size=sam2_model_size, | |
| prompt_mode=sam2_prompt_mode, | |
| ) | |
| sam2 = _get_sam2_backend(sam2_cfg) | |
| if sam2 is not None: | |
| weak = SegmentResult( | |
| masks=np.asarray(masks, dtype=bool) if len(masks) else np.zeros((0, 1, 1), dtype=bool), | |
| scores=np.asarray(scores, dtype=np.float32), | |
| boxes=weak_boxes, | |
| ) | |
| try: | |
| refined = sam2.refine(tile_rgb, weak, sam2_prompt_mode) | |
| masks = refined.masks | |
| scores = refined.scores | |
| print( | |
| f"[TILE {tile_idx + 1}/{total_tiles}] {ft}: " | |
| f"SAM2 refined {len(refined)} masks (mode='{sam2_prompt_mode}')" | |
| ) | |
| except Exception as e: | |
| print(f"[SAM2] refinement failed: {e} — keeping unrefined") | |
| else: | |
| print("[SAM2] backend unavailable — keeping unrefined") | |
| print( | |
| f"[TILE {tile_idx + 1}/{total_tiles}] {ft}: model returned " | |
| f"{len(masks)} masks, scores={scores.tolist() if len(scores) < 20 else f'min={float(scores.min()):.2f} max={float(scores.max()):.2f}'}" | |
| ) | |
| if len(masks) == 0: | |
| print(f"[TILE {tile_idx + 1}/{total_tiles}] {ft}: NO MASKS — skipping vectorize") | |
| continue | |
| # Vectorize each mask separately and keep only the LARGEST polygon | |
| # per mask to avoid fragmentation (e.g. masks with holes/islands). | |
| polys_added = 0 | |
| for i, mask in enumerate(masks): | |
| if mask.sum() == 0: | |
| continue | |
| single = (mask > 0).astype(np.int32) | |
| candidate_polys = [] | |
| for geom, _val in shapes( | |
| single, mask=(single > 0), transform=tile_transform | |
| ): | |
| candidate_polys.append(shape(geom)) | |
| if not candidate_polys: | |
| continue | |
| # Take the largest polygon (drop holes/islands) | |
| largest = max(candidate_polys, key=lambda g: g.area) | |
| score = float(scores[i]) if i < len(scores) else 0.0 | |
| tracker.add(ft, score) | |
| all_features.append({ | |
| "geometry": largest, | |
| "feature_type": ft, | |
| "confidence": round(score, 3), | |
| }) | |
| polys_added += 1 | |
| print( | |
| f"[TILE {tile_idx + 1}/{total_tiles}] {ft}: " | |
| f"vectorized {polys_added} polygons from {len(masks)} masks" | |
| ) | |
| elapsed = time.time() - start_time | |
| dash = format_dashboard( | |
| image_name, img_w, img_h, | |
| tile_idx + 1, total_tiles, | |
| elapsed, tracker, feature_types, | |
| ) | |
| # Checkpoint: save downloadable files every N tiles | |
| # Update overlay + save files every PREVIEW_INTERVAL tiles | |
| if (tile_idx + 1) % PREVIEW_INTERVAL == 0 or tile_idx == total_tiles - 1 or tile_idx == 0: | |
| temp_by_type = {} | |
| for ft in feature_types: | |
| ft_feats = [f for f in all_features if f["feature_type"] == ft] | |
| if ft_feats: | |
| temp_by_type[ft] = gpd.GeoDataFrame(ft_feats, crs=crs_obj) | |
| # Save files every update so they're always downloadable | |
| checkpoint_files = None | |
| if temp_by_type: | |
| checkpoint_files = export_all(temp_by_type, tmp_dir) | |
| scanned = (scan_left, scan_bottom, scan_right, scan_top) | |
| overlay = make_overlay(geotiff, temp_by_type, scanned_bounds=scanned) | |
| yield overlay, dash, checkpoint_files | |
| else: | |
| yield gr.update(), dash, gr.update() | |
| # -- Deduplicate + filter -- | |
| elapsed = time.time() - start_time | |
| yield gr.update(), format_dashboard( | |
| image_name, img_w, img_h, total_tiles, total_tiles, | |
| elapsed, tracker, feature_types, "Deduplicating & filtering", | |
| ), None | |
| final_by_type: dict[str, gpd.GeoDataFrame] = {} | |
| final_counts: dict[str, int] = {} | |
| for ft in feature_types: | |
| ft_feats = [f for f in all_features if f["feature_type"] == ft] | |
| if not ft_feats: | |
| final_counts[ft] = 0 | |
| continue | |
| gdf = merge_and_deduplicate(ft_feats, crs_obj) | |
| gdf = filter_features(gdf, FEATURE_PRESETS[ft]) | |
| # Apply simplification + orthogonalization to building-like features only | |
| if ft in ("Building", "Rooftop", "Parking Lot") and len(gdf) > 0: | |
| gdf = simplify_and_orthogonalize( | |
| gdf, | |
| simplify_tolerance_m=0.5, | |
| orthogonalize=True, | |
| rect_iou_threshold=0.85, | |
| ) | |
| final_by_type[ft] = gdf | |
| final_counts[ft] = len(gdf) | |
| # -- Export -- | |
| export_paths = export_all(final_by_type, tmp_dir) | |
| # -- Final overlay -- | |
| final_overlay = make_overlay(geotiff, final_by_type) | |
| elapsed = time.time() - start_time | |
| final_dash = format_final_summary(tracker, feature_types, final_counts, elapsed) | |
| yield final_overlay, final_dash, export_paths | |
| # --------------------------------------------------------------------------- | |
| # Gradio UI | |
| # --------------------------------------------------------------------------- | |
| CUSTOM_CSS = """ | |
| @import url('https://fonts.googleapis.com/css2?family=Outfit:wght@300;400;500;600;700&family=JetBrains+Mono:wght@400;500&display=swap'); | |
| :root { | |
| --body-background-fill: #fafafa !important; | |
| --block-background-fill: #ffffff !important; | |
| --block-border-color: #e4e4e7 !important; | |
| --block-label-text-color: #3f3f46 !important; | |
| --block-title-text-color: #18181b !important; | |
| --button-primary-background-fill: #18181b !important; | |
| --button-primary-text-color: #fafafa !important; | |
| --button-primary-background-fill-hover: #27272a !important; | |
| --input-background-fill: #ffffff !important; | |
| --border-color-primary: #e4e4e7 !important; | |
| } | |
| * { font-family: 'Outfit', system-ui, -apple-system, sans-serif !important; } | |
| code, pre, .code, [class*="mono"] { font-family: 'JetBrains Mono', monospace !important; } | |
| .gradio-container { max-width: 1400px !important; margin: 0 auto !important; background: #fafafa !important; } | |
| .gr-button-primary { | |
| border-radius: 12px !important; font-weight: 600 !important; | |
| letter-spacing: -0.01em !important; padding: 12px 32px !important; | |
| transition: all 0.2s cubic-bezier(0.16, 1, 0.3, 1) !important; | |
| box-shadow: 0 1px 3px rgba(0,0,0,0.08) !important; | |
| } | |
| .gr-button-primary:hover { transform: translateY(-1px) !important; box-shadow: 0 4px 12px rgba(0,0,0,0.1) !important; } | |
| .gr-button-primary:active { transform: translateY(0) scale(0.98) !important; } | |
| .gr-panel, .gr-box, .gr-form { | |
| border-radius: 20px !important; border: 1px solid #e4e4e7 !important; | |
| box-shadow: 0 20px 40px -15px rgba(0,0,0,0.04) !important; | |
| } | |
| .gr-input, .gr-textbox textarea, select { | |
| border-radius: 12px !important; border: 1px solid #e4e4e7 !important; | |
| font-size: 14px !important; transition: border-color 0.2s ease !important; | |
| } | |
| .gr-input:focus, .gr-textbox textarea:focus { | |
| border-color: #18181b !important; box-shadow: 0 0 0 3px rgba(24,24,27,0.06) !important; | |
| } | |
| h1 { font-size: 2.25rem !important; font-weight: 700 !important; letter-spacing: -0.03em !important; line-height: 1.1 !important; color: #18181b !important; } | |
| h2, h3, .gr-block-label { font-weight: 600 !important; letter-spacing: -0.02em !important; color: #3f3f46 !important; } | |
| .markdown-text p { color: #52525b !important; line-height: 1.6 !important; max-width: 65ch !important; } | |
| footer { display: none !important; } | |
| """ | |
| HEADER_MD = """ | |
| # Janus [Auto-Synced] | |
| ### GIS-ready feature extraction from aerial imagery | |
| Upload a GeoTIFF (any size), select feature types to extract, and watch SAM 3 process | |
| tile by tile with live confidence tracking. | |
| """ | |
| def _gpu_status(): | |
| try: | |
| if torch.cuda.is_available(): | |
| return ( | |
| f"Running on **{torch.cuda.get_device_name(0)}** " | |
| f"({torch.cuda.get_device_properties(0).total_memory / (1024**3):.0f} GB)" | |
| ) | |
| except Exception: | |
| pass | |
| return "No GPU detected — inference will be slow" | |
| gpu_status = _gpu_status() | |
| with gr.Blocks(css=CUSTOM_CSS, title="Janus — Feature Extraction") as demo: | |
| gr.Markdown(HEADER_MD) | |
| gr.Markdown(f"*{gpu_status}*") | |
| with gr.Row(equal_height=False): | |
| with gr.Column(scale=3): | |
| with gr.Group(): | |
| image_input = gr.File( | |
| label="Aerial Image", | |
| file_types=[".png", ".jpg", ".jpeg", ".tif", ".tiff"], | |
| type="filepath", | |
| ) | |
| world_input = gr.File( | |
| label="World File (optional for GeoTIFF)", | |
| file_types=[".pgw", ".jgw", ".tfw"], | |
| type="filepath", | |
| ) | |
| crs_input = gr.Textbox(value="EPSG:4326", label="CRS") | |
| feature_checks = gr.CheckboxGroup( | |
| choices=list(FEATURE_PRESETS.keys()), | |
| value=["Building"], | |
| label="Feature Types", | |
| info="Select one or more feature types to extract", | |
| ) | |
| detection_mode = gr.Radio( | |
| choices=[ | |
| "SAM 3 (text prompt only)", | |
| "Grounding DINO + SAM 3 (two-stage)", | |
| "RAMP (aerial buildings specialist)", | |
| "YOLO segmentation (experimental)", | |
| ], | |
| value="SAM 3 (text prompt only)", | |
| label="Detection Mode", | |
| info=( | |
| "SAM 3 is general-purpose. Two-stage adds Grounding DINO for " | |
| "better recall. RAMP is a specialist trained on aerial buildings. " | |
| "YOLO segmentation is multi-class and works well as a weak detector " | |
| "for SAM 2 refinement (toggle below)." | |
| ), | |
| ) | |
| gdino_threshold_slider = gr.Slider( | |
| minimum=0.05, maximum=0.50, value=0.15, step=0.01, | |
| label="Grounding DINO Threshold (two-stage only)", | |
| info=( | |
| "Lower = more candidate boxes (catches more buildings, " | |
| "more false positives). Only used in two-stage mode." | |
| ), | |
| ) | |
| with gr.Row(): | |
| confidence_slider = gr.Slider( | |
| minimum=0.1, maximum=0.95, value=0.5, step=0.05, | |
| label="Confidence Threshold", | |
| ) | |
| tile_size_dropdown = gr.Dropdown( | |
| choices=[256, 512, 1024, 2048], | |
| value=256, | |
| label="Tile Size (px)", | |
| info="Larger = fewer tiles but more GPU memory", | |
| ) | |
| with gr.Accordion("⚠ Experimental Models", open=False) as experimental_section: | |
| gr.Markdown( | |
| "*These backends are under active development. RAMP is the current " | |
| "production recommendation for buildings.*" | |
| ) | |
| use_yolo_seg_check = gr.Checkbox( | |
| value=False, | |
| label="YOLO segmentation toggle", | |
| info=( | |
| "Used when 'YOLO segmentation' is selected as the Detection Mode. " | |
| "Out of the box, YOLO uses a generic COCO-pretrained model — best " | |
| "as a weak detector feeding SAM 2 refinement until a fine-tuned " | |
| "aerial checkpoint is provided." | |
| ), | |
| ) | |
| with gr.Row(): | |
| yolo_model_path_input = gr.Textbox( | |
| value="yolov8n-seg.pt", | |
| label="YOLO model path / hub id", | |
| ) | |
| yolo_conf_slider = gr.Slider( | |
| minimum=0.05, maximum=0.9, value=0.25, step=0.01, | |
| label="YOLO confidence", | |
| ) | |
| yolo_iou_slider = gr.Slider( | |
| minimum=0.1, maximum=0.95, value=0.45, step=0.05, | |
| label="YOLO IoU", | |
| ) | |
| use_sam2_refinement_check = gr.Checkbox( | |
| value=False, | |
| label="SAM 2 refinement", | |
| info=( | |
| "Take boxes/masks from the current Detection Mode and pass them " | |
| "to SAM 2 for cleaner per-instance masks. Best paired with " | |
| "Grounding DINO + SAM 3 or YOLO." | |
| ), | |
| ) | |
| with gr.Row(): | |
| sam2_size_dropdown = gr.Dropdown( | |
| choices=["tiny", "small", "base", "large"], | |
| value="large", | |
| label="SAM 2 model size", | |
| ) | |
| sam2_prompt_dropdown = gr.Dropdown( | |
| choices=["box", "point", "mask", "hybrid"], | |
| value="box", | |
| label="SAM 2 prompt mode", | |
| ) | |
| run_btn = gr.Button("Extract Features", variant="primary", size="lg") | |
| with gr.Column(scale=5): | |
| dashboard = gr.Markdown(value="*Upload an image and click Extract Features to begin.*") | |
| output_image = gr.Image(label="Live Preview", type="filepath", show_download_button=True) | |
| file_output = gr.Files(label="Download GIS Files") | |
| run_btn.click( | |
| fn=run_pipeline, | |
| inputs=[ | |
| image_input, world_input, crs_input, feature_checks, | |
| confidence_slider, tile_size_dropdown, detection_mode, gdino_threshold_slider, | |
| # Experimental | |
| use_yolo_seg_check, yolo_model_path_input, yolo_conf_slider, yolo_iou_slider, | |
| use_sam2_refinement_check, sam2_size_dropdown, sam2_prompt_dropdown, | |
| ], | |
| outputs=[output_image, dashboard, file_output], | |
| ) | |
| if __name__ == "__main__": | |
| demo.launch() | |