""" app.py — Gradio entry point for the GeoAP Hugging Face Space. Local: python app.py (http://127.0.0.1:7860) Space: HF runs this file automatically (sdk: gradio in README.md). Layout: upload (any raster, incl. GeoTIFF) -> run detection or segmentation -> annotated image + Counts table -> interactive Leaflet map of the GeoJSON. """ from __future__ import annotations import json import logging import os import tempfile from pathlib import Path from typing import Any, Dict, List, Optional import cv2 import gradio as gr import numpy as np from src.inference import PredictionResult, run_detection, run_segmentation, to_geojson from src.io_raster import SUPPORTED_EXTS, Raster, load_raster from src.mapview import PLACEHOLDER, build_map_html, class_color_map from src.settings import ROOT, load_settings from src.viz import hex_for logging.basicConfig(level=logging.INFO, format="[%(asctime)s] %(levelname)s %(name)s — %(message)s") logger = logging.getLogger("geoap.space") # ── ZeroGPU ────────────────────────────────────────────────────────────────── # On ZeroGPU hardware a GPU is attached only for the duration of a function # marked with @spaces.GPU. The `spaces` package exists only on HF, so locally # we fall back to a no-op decorator and nothing changes. try: import spaces # type: ignore def gpu(duration: int = 120): return spaces.GPU(duration=duration) except ImportError: # local run / non-ZeroGPU Space def gpu(duration: int = 120): def _identity(fn): return fn return _identity # Gradio 6 moved `theme` from Blocks(...) to launch(...). Support both. GRADIO_MAJOR = int(gr.__version__.split(".")[0]) _THEME_ON_LAUNCH = GRADIO_MAJOR >= 6 S = load_settings() DET_COLORS = class_color_map(S.detection_classes, hex_for) SEG_COLORS = class_color_map(S.segmentation_classes, hex_for) EXAMPLES = sorted( p for p in (ROOT / "examples").glob("*") if p.suffix.lower() in set(SUPPORTED_EXTS) ) def _to_rgb(bgr: np.ndarray) -> np.ndarray: return cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB) def _write(text: str, name: str) -> str: out = Path(tempfile.gettempdir()) / name out.write_text(text, encoding="utf-8") return str(out) def _count_rows(result: PredictionResult, class_names: List[str]) -> List[List[Any]]: """Counts table: class, count, colour swatch hex (the only legend now).""" rows: List[List[Any]] = [ [name, int(result.counts.get(name, 0)), hex_for(i)] for i, name in enumerate(class_names) ] rows.append(["TOTAL", int(sum(result.counts.values())), ""]) return rows # ── Upload ─────────────────────────────────────────────────────────────────── def load_input(file_path: Optional[str]): """Read the uploaded raster once; keep it (and its geotransform) in state.""" if not file_path: return None, None, "", PLACEHOLDER try: raster = load_raster(file_path) except Exception as exc: # noqa: BLE001 logger.exception("load failed") return None, None, f"**Could not read this file** — {exc}", PLACEHOLDER info = f"`{Path(file_path).name}`\n\n" + "\n\n".join(raster.info_lines()) return raster.rgb, raster, info, PLACEHOLDER # ── Prediction ─────────────────────────────────────────────────────────────── def _finish(result: PredictionResult, raster: Raster, class_names: List[str], colors: Dict[str, str], stem: str): gj = to_geojson(result, px_to_lonlat=raster.px_to_lonlat) iframe, page = build_map_html( gj, colors=colors, image_bgr=result.annotated_bgr, geographic=raster.georeferenced, ) downloads = [ _write(json.dumps(gj, indent=2), f"geoap_{stem}.geojson"), _write(page, f"geoap_{stem}_map.html"), ] stats = result.to_dict() stats.pop("objects", None) stats["georeferenced"] = raster.georeferenced return _to_rgb(result.annotated_bgr), _count_rows(result, class_names), stats, downloads, iframe def _no_input(): return None, [], {"error": "Upload an image first."}, None, PLACEHOLDER @gpu(duration=120) def predict_detection(raster, conf, iou, tile_size, overlap, progress=gr.Progress()): if raster is None: return _no_input() progress(0.1, desc="Loading detection model…") res = run_detection( raster.rgb, S, conf=conf, iou=iou, tile_size=int(tile_size), overlap=overlap ) progress(0.9, desc="Building map…") return _finish(res, raster, S.detection_classes, DET_COLORS, "detection") @gpu(duration=180) def predict_segmentation(raster, conf, iou, tile_size, overlap, progress=gr.Progress()): if raster is None: return _no_input() progress(0.05, desc="Loading segmentation model…") def cb(frac: float, desc: str) -> None: progress(0.05 + 0.85 * frac, desc=desc) res = run_segmentation( raster.rgb, S, conf=conf, iou=iou, tile_size=int(tile_size), overlap=overlap, progress_cb=cb, ) return _finish(res, raster, S.segmentation_classes, SEG_COLORS, "segmentation") # ── UI ─────────────────────────────────────────────────────────────────────── def build_demo() -> gr.Blocks: blocks_kwargs: Dict[str, Any] = {"title": S.title} if not _THEME_ON_LAUNCH: blocks_kwargs["theme"] = gr.themes.Soft() with gr.Blocks(**blocks_kwargs) as demo: gr.Markdown(f"# {S.title}\n{S.description}") raster_state = gr.State() with gr.Row(): with gr.Column(scale=1): file_in = gr.File( label="Image (GeoTIFF, TIFF, JP2, PNG, JPG, WEBP, BMP)", file_count="single", file_types=SUPPORTED_EXTS, type="filepath", ) preview = gr.Image(label="Input preview", height=260, interactive=False) file_info = gr.Markdown() conf = gr.Slider(0.05, 0.95, value=S.confidence_threshold, step=0.05, label="Confidence") iou = gr.Slider(0.1, 0.9, value=S.iou_threshold, step=0.05, label="IoU / merge threshold") tile = gr.Dropdown([512, 640, 768, 1024], value=S.tile_size, label="Tile size") ov = gr.Slider(0.0, 0.5, value=S.overlap, step=0.05, label="Tile overlap") with gr.Row(): det_btn = gr.Button("Run detection", variant="primary") seg_btn = gr.Button("Run segmentation", variant="primary") if EXAMPLES: gr.Examples([[str(p)] for p in EXAMPLES], inputs=[file_in]) with gr.Column(scale=2): out_img = gr.Image(label="Result", height=560) counts = gr.Dataframe( headers=["class", "count", "colour"], datatype=["str", "number", "str"], label="Counts", interactive=False, wrap=True, ) with gr.Accordion("Run details", open=False): stats_json = gr.JSON() geojson_file = gr.File( label="Downloads — GeoJSON + standalone map (.html)", file_count="multiple", ) gr.Markdown("### Map") map_html = gr.HTML(PLACEHOLDER) with gr.Accordion("About the models", open=False): gr.Markdown( f""" | task | base | classes | |---|---|---| | detection | `{S.model.detection_base}` | {", ".join(S.detection_classes)} | | segmentation | `{S.model.segmentation_base}` | {", ".join(S.segmentation_classes)} | **Hub repos** — detection `{S.hub.detection_repo or "not set"}`, segmentation `{S.hub.segmentation_repo or "not set"}` **Inference** — images are cut into `{S.tile_size}px` tiles with `{int(S.overlap * 100)}%` overlap. Detection uses SAHI sliced prediction and merges duplicate boxes across seams with `{S.nms_type}`. Segmentation runs per tile and composes masks back onto a full-resolution canvas. Anything longer than `{S.max_side}px` on a side is downscaled first. **Map** — a georeferenced input (GeoTIFF with a CRS) puts features on an OpenStreetMap basemap in EPSG:4326. A plain image falls back to pixel coordinates drawn over the result itself. Reprojection needs `rasterio`. Weights resolve local `weights/*.pt` → Hub model repo → base COCO checkpoint. COCO class names in the output mean the fine-tuned weights were not found. """ ) file_in.change( load_input, inputs=[file_in], outputs=[preview, raster_state, file_info, map_html], ) det_btn.click( predict_detection, inputs=[raster_state, conf, iou, tile, ov], outputs=[out_img, counts, stats_json, geojson_file, map_html], ) seg_btn.click( predict_segmentation, inputs=[raster_state, conf, iou, tile, ov], outputs=[out_img, counts, stats_json, geojson_file, map_html], ) return demo if __name__ == "__main__": launch_kwargs: Dict[str, Any] = {"server_name": "0.0.0.0", "show_error": True} if _THEME_ON_LAUNCH: launch_kwargs["theme"] = gr.themes.Soft() # GEOAP_SHARE=1 opens a temporary public gradio.live tunnel so others can # test the app without a Space. Never set on HF (SPACE_ID is defined there). if os.getenv("GEOAP_SHARE") == "1" and not os.getenv("SPACE_ID"): launch_kwargs["share"] = True logger.info("share mode on — a public *.gradio.live URL will be printed below") build_demo().queue(max_size=12).launch(**launch_kwargs)