multimodalart's picture
multimodalart HF Staff
Right-size ZeroGPU duration from mesh header, fix Gradio 6 theme/css, add second example scan
000083b verified
Raw
History Blame Contribute Delete
9.95 kB
"""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))))
@spaces.GPU(duration=_gpu_duration)
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)