Spaces:
Running on Zero
Running on Zero
multimodalart HF Staff
Right-size ZeroGPU duration from mesh header, fix Gradio 6 theme/css, add second example scan
000083b verified | """AQ3D — Adaptive Query Transformer for 3D Instance Segmentation. | |
| Upload an indoor surface mesh; the Space runs the official `aq3d_scannet200_volt` | |
| checkpoint and paints every detected object instance with its ScanNet200 class | |
| color. | |
| """ | |
| import os | |
| os.environ.setdefault("NUMBA_DISABLE_CUDA", "1") # keep numba off the GPU | |
| os.environ.setdefault("NUMBA_CACHE_DIR", "/tmp/numba-cache") | |
| import spaces # noqa: E402 (before torch) | |
| import json # noqa: E402 | |
| import tempfile # noqa: E402 | |
| import time # noqa: E402 | |
| from typing import List, Optional, Tuple # noqa: E402 | |
| import gradio as gr # noqa: E402 | |
| import numpy as np # noqa: E402 | |
| import torch # noqa: E402 | |
| from huggingface_hub import hf_hub_download # noqa: E402 | |
| import pipeline as P # noqa: E402 | |
| import superpoints # noqa: E402 | |
| from model import AQ3D # noqa: E402 | |
| REPO_ID = "kenomo/aq3d" | |
| CKPT = "weights/aq3d_scannet200_volt.pth" | |
| superpoints.warmup() # JIT the Numba kernels once | |
| _ckpt_path = hf_hub_download(REPO_ID, CKPT) | |
| _state = torch.load(_ckpt_path, map_location="cpu", weights_only=False)["state_dict"] | |
| _state = {k: v for k, v in _state.items() if not k.startswith("criterion.")} | |
| model = AQ3D(num_classes=len(P.CLASS_NAMES)) | |
| model.load_state_dict(_state, strict=True) | |
| model.eval().to("cuda") | |
| del _state | |
| HEADERS = ["#", "class", "score", "points", "color"] | |
| def _empty(msg: str): | |
| return None, [], msg | |
| # --------------------------------------------------------------------------- # | |
| # GPU-duration estimate | |
| # | |
| # Runtime is dominated by the vertex count (superpoint graph segmentation, the | |
| # 0.6 x |superpoints| adaptive queries and their NMS). Measured on this Space: | |
| # 211k vertices -> 8.2 s, 526k vertices -> 19.2 s, i.e. ~0.035 s per 1k vertices. | |
| # The vertex count is read straight out of the file header (cheap, no parsing of | |
| # the geometry) so every visitor only reserves the quota their own scan needs. | |
| # --------------------------------------------------------------------------- # | |
| def _vertex_count(path: str) -> Optional[int]: | |
| """Vertex count from a glTF-binary / PLY header, without loading geometry.""" | |
| try: | |
| with open(path, "rb") as fh: | |
| head = fh.read(20) | |
| if head[:4] == b"glTF": | |
| chunk_len = int.from_bytes(head[12:16], "little") | |
| if head[16:20] != b"JSON": | |
| return None | |
| gltf = json.loads(fh.read(chunk_len).decode("utf-8", "replace")) | |
| accessors = gltf.get("accessors", []) | |
| total = 0 | |
| for mesh in gltf.get("meshes", []): | |
| for prim in mesh.get("primitives", []): | |
| i = prim.get("attributes", {}).get("POSITION") | |
| if isinstance(i, int) and 0 <= i < len(accessors): | |
| total += int(accessors[i].get("count", 0)) | |
| return total or None | |
| if head[:3] == b"ply": | |
| fh.seek(0) | |
| for line in fh.read(8192).split(b"\n"): | |
| if line.startswith(b"element vertex"): | |
| return int(line.split()[2]) | |
| except Exception: | |
| pass | |
| return None | |
| def _gpu_duration(mesh_file: Optional[str], *args, **kwargs) -> int: | |
| if not mesh_file or not os.path.exists(mesh_file): | |
| return 25 | |
| verts = _vertex_count(mesh_file) | |
| if verts is None: # OBJ / unknown container: size proxy | |
| verts = os.path.getsize(mesh_file) / 60.0 | |
| seconds = (1.0 + 0.035 * verts / 1000.0) * 1.5 + 5.0 # fit + 50% + fork cost | |
| return int(min(75, max(20, round(seconds)))) | |
| def segment( | |
| mesh_file: Optional[str], | |
| up_axis: str = "Auto", | |
| scale: float = 1.0, | |
| auto_fit: bool = True, | |
| threshold: float = 0.35, | |
| max_instances: int = 40, | |
| progress=gr.Progress(track_tqdm=True), | |
| ) -> Tuple[Optional[str], List[List], str]: | |
| """Segment an indoor 3D scan into labelled object instances with AQ3D. | |
| Args: | |
| mesh_file: Path to a triangle-mesh scan (.ply, .obj or .glb) of an indoor scene. | |
| up_axis: Which axis points up in the uploaded mesh ("Auto", "Z", "Y" or "X"). | |
| scale: Multiplier applied to the mesh coordinates to bring them into metres. | |
| auto_fit: Rescale the scene automatically when its footprint is not room sized. | |
| threshold: Minimum instance confidence to keep, between 0 and 1. | |
| max_instances: Maximum number of instances to display. | |
| Returns: | |
| A GLB mesh colored by predicted instance, a table of the detected | |
| instances, and a short status message. | |
| """ | |
| if not mesh_file: | |
| return _empty("Please upload a mesh first.") | |
| t_start = time.time() | |
| try: | |
| mesh = P.load_mesh(mesh_file) | |
| except ValueError as exc: | |
| return _empty(f"❌ {exc}") | |
| rgb = P.mesh_vertex_colors(mesh) | |
| faces = np.ascontiguousarray(mesh.faces, dtype=np.int64) | |
| verts, used_axis, used_scale = P.orient_and_scale( | |
| np.ascontiguousarray(mesh.vertices, dtype=np.float32), up_axis, scale, auto_fit | |
| ) | |
| batch, spts = P.build_batch(verts, faces, rgb, torch.device("cuda")) | |
| with torch.no_grad(): | |
| out = model(batch) | |
| labels, scores, masks_binary, npoints = P.decode_predictions( | |
| out, spts, num_classes=len(P.CLASS_NAMES) | |
| ) | |
| del out, batch | |
| torch.cuda.empty_cache() | |
| colored, rows = P.colorize( | |
| verts, faces, spts, labels, scores, masks_binary, npoints, | |
| float(threshold), int(max_instances), | |
| ) | |
| path = tempfile.mktemp(suffix=".glb") | |
| colored.export(path) | |
| extent = verts.max(0) - verts.min(0) | |
| status = ( | |
| f"✅ **{len(rows)} instances** above {threshold:.2f} · " | |
| f"{len(verts):,} vertices → {int(spts.max()) + 1:,} superpoints · " | |
| f"scene {extent[0]:.1f} × {extent[1]:.1f} × {extent[2]:.1f} m " | |
| f"(up axis `{used_axis}`, scale ×{used_scale:.3g}) · {time.time() - t_start:.1f}s" | |
| ) | |
| return path, rows, status | |
| CSS = """ | |
| .dark .gradio-container { color: var(--body-text-color); } | |
| #viewer { height: 520px; } | |
| """ | |
| with gr.Blocks(title="AQ3D 3D Instance Segmentation") as demo: | |
| gr.Markdown( | |
| """ | |
| # 🪑 AQ3D — 3D Instance Segmentation | |
| Upload an indoor **surface mesh** (`.ply` / `.obj` / `.glb`) and AQ3D will find | |
| every object in it and label it with one of the **198 ScanNet200 classes**. | |
| [Paper](https://huggingface.co/papers/2608.30618) · | |
| [Code](https://github.com/kenomo/aq3d) · | |
| [Weights](https://huggingface.co/kenomo/aq3d) — running `aq3d_scannet200_volt` | |
| (Volt-B backbone + adaptive-query decoder). | |
| """ | |
| ) | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| mesh_in = gr.Model3D(label="Input scan", elem_id="viewer") | |
| run_btn = gr.Button("Segment scene", variant="primary") | |
| with gr.Column(scale=1): | |
| mesh_out = gr.Model3D(label="Instance segmentation", elem_id="viewer") | |
| status = gr.Markdown() | |
| table = gr.Dataframe( | |
| headers=HEADERS, label="Detected instances", wrap=True, | |
| datatype=["number", "str", "number", "number", "str"], | |
| ) | |
| with gr.Accordion("Options", open=False): | |
| with gr.Row(): | |
| threshold = gr.Slider(0.0, 1.0, value=0.35, step=0.01, | |
| label="Confidence threshold") | |
| max_instances = gr.Slider(1, 150, value=40, step=1, | |
| label="Max instances shown") | |
| gr.Markdown( | |
| "AQ3D expects **metric, Z-up** room scans. Override the automatic guess " | |
| "here if the scene comes out mislabelled." | |
| ) | |
| with gr.Row(): | |
| up_axis = gr.Radio(["Auto", "Z", "Y", "X"], value="Auto", label="Up axis") | |
| scale = gr.Number(value=1.0, label="Scale factor", minimum=1e-4) | |
| auto_fit = gr.Checkbox(value=True, label="Auto-fit to room size") | |
| gr.Examples( | |
| examples=[ | |
| ["examples/attic.glb"], | |
| ["examples/historic-interior.glb"], | |
| ], | |
| inputs=[mesh_in], | |
| outputs=[mesh_out, table, status], | |
| fn=segment, | |
| cache_examples=True, | |
| cache_mode="lazy", | |
| label="Example room scans (CC BY 4.0, via Objaverse / Zenodo)", | |
| ) | |
| gr.Markdown( | |
| """ | |
| ### Notes | |
| * Superpoints come from a faithful Numba port of the ScanNet | |
| `segmentator` (Felzenszwalb–Huttenlocher) mesh segmentation used by AQ3D, | |
| so a **triangle mesh is required** — raw point clouds are rejected. | |
| * Preprocessing mirrors the official ScanNet200 validation config: | |
| mean-centred coordinates, colors normalised to [-1, 1], 2 cm voxel grid, | |
| superpoint attention pooling, superpoint NMS (0.8), adaptive top-k. | |
| * Example scans: *"my room and the mess therein"* | |
| ([Zenodo 10380976](https://zenodo.org/records/10380976)) and a LiDAR capture of a | |
| historic building interior ([Zenodo 10325220](https://zenodo.org/records/10325220)), | |
| both CC BY 4.0. | |
| """ | |
| ) | |
| run_btn.click( | |
| segment, | |
| inputs=[mesh_in, up_axis, scale, auto_fit, threshold, max_instances], | |
| outputs=[mesh_out, table, status], | |
| ) | |
| if __name__ == "__main__": | |
| # Gradio 6 moved `theme` / `css` from the Blocks constructor to launch(). | |
| demo.queue().launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True) | |