"""Sapiens2 pointmap Gradio Space. Image → per-pixel 3D pointmap (camera frame, metric units). Visualized as a .ply point cloud rendered with Gradio's Model3D component for interactive 3D viewing. Foreground mask is mandatory. Everything runs at the model's NATIVE resolution (max 1024×768 grid → at most ~786K points before subsampling to 200K). No huge interpolations. """ import sys import os sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) import tempfile import cv2 import gradio as gr import numpy as np import open3d as o3d import spaces import torch import torch.nn.functional as F from PIL import Image from torchvision import transforms from huggingface_hub import hf_hub_download from sapiens.dense.models import PointmapEstimator, init_model # registers in registry _ = PointmapEstimator # ----------------------------------------------------------------------------- # Config ASSETS_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "assets") CONFIGS_DIR = os.path.join(ASSETS_DIR, "configs") POINTMAP_MODELS = { "0.4B": { "repo": "facebook/sapiens2-pointmap-0.4b", "filename": "sapiens2_0.4b_pointmap.safetensors", "config": os.path.join(CONFIGS_DIR, "sapiens2_0.4b_pointmap_render_people-1024x768.py"), }, "0.8B": { "repo": "facebook/sapiens2-pointmap-0.8b", "filename": "sapiens2_0.8b_pointmap.safetensors", "config": os.path.join(CONFIGS_DIR, "sapiens2_0.8b_pointmap_render_people-1024x768.py"), }, "1B": { "repo": "facebook/sapiens2-pointmap-1b", "filename": "sapiens2_1b_pointmap.safetensors", "config": os.path.join(CONFIGS_DIR, "sapiens2_1b_pointmap_render_people-1024x768.py"), }, "5B": { "repo": "facebook/sapiens2-pointmap-5b", "filename": "sapiens2_5b_pointmap.safetensors", "config": os.path.join(CONFIGS_DIR, "sapiens2_5b_pointmap_render_people-1024x768.py"), }, } DEFAULT_SIZE = "0.4B" # iteration mode — only this is preloaded; others lazy-load on click FG_REPO = "facebook/sapiens-seg-foreground-1b-torchscript" FG_FILENAME = "sapiens_1b_seg_foreground_epoch_8_torchscript.pt2" DEVICE = "cuda" if torch.cuda.is_available() else "cpu" _fg_transform = transforms.Compose([ transforms.Resize((1024, 768)), transforms.ToTensor(), transforms.Normalize(mean=[123.5 / 255, 116.5 / 255, 103.5 / 255], std=[58.5 / 255, 57.0 / 255, 57.5 / 255]), ]) # ----------------------------------------------------------------------------- # Model cache _pointmap_model_cache: dict = {} _fg_model = None def _get_pointmap_model(size: str): if size not in _pointmap_model_cache: spec = POINTMAP_MODELS[size] ckpt = hf_hub_download(repo_id=spec["repo"], filename=spec["filename"]) model = init_model(spec["config"], ckpt, device=DEVICE) _pointmap_model_cache[size] = model return _pointmap_model_cache[size] def _get_fg_model(): global _fg_model if _fg_model is None: ckpt = hf_hub_download(repo_id=FG_REPO, filename=FG_FILENAME) _fg_model = torch.jit.load(ckpt).eval().to(DEVICE) return _fg_model print("[startup] pre-loading 0.4B (iteration mode) + fg/bg ...") _get_pointmap_model(DEFAULT_SIZE) _get_fg_model() print("[startup] ready.") # ----------------------------------------------------------------------------- # Inference (always at native resolution) def _estimate_pointmap(image_bgr: np.ndarray, model) -> np.ndarray: data = model.pipeline(dict(img=image_bgr)) data = model.data_preprocessor(data) inputs, data_samples = data["inputs"], data["data_samples"] if inputs.ndim == 3: inputs = inputs.unsqueeze(0) with torch.no_grad(): pointmap, scale = model(inputs) pointmap = pointmap / scale # → metric pad_left, pad_right, pad_top, pad_bottom = data_samples["meta"]["padding_size"] pointmap = pointmap[ :, :, pad_top : inputs.shape[2] - pad_bottom, pad_left : inputs.shape[3] - pad_right, ] return pointmap.squeeze(0).cpu().float().numpy().transpose(1, 2, 0) # (H_native, W_native, 3) def _foreground_mask(image_pil: Image.Image, target_h: int, target_w: int) -> np.ndarray: fg = _get_fg_model() inputs = _fg_transform(image_pil).unsqueeze(0).to(DEVICE) with torch.no_grad(): out = fg(inputs) out = F.interpolate(out, size=(target_h, target_w), mode="bilinear", align_corners=False) return (out.argmax(dim=1)[0] > 0).cpu().numpy() def _depth_to_rgb(depth: np.ndarray, mask: np.ndarray) -> np.ndarray: """Inverse-depth turbo colormap (matches sapiens2 vis_pointmap.py). Background pixels are left at 0 — caller should overlay them.""" valid = np.isfinite(depth) & (depth > 1e-3) & mask rgb = np.zeros((*depth.shape, 3), dtype=np.uint8) if not valid.any(): return rgb inv = np.zeros_like(depth, dtype=np.float32) inv[valid] = 1.0 / depth[valid] p1, p99 = np.percentile(inv[valid], [1, 99]) lo, hi = float(p1), float(p99) if hi <= lo: hi = lo + 1e-3 norm = ((inv - lo) / (hi - lo)).clip(0, 1) grey = (norm * 255.0).astype(np.uint8) color = cv2.applyColorMap(grey, cv2.COLORMAP_TURBO)[:, :, ::-1] # cv2 is BGR → RGB rgb[valid] = color[valid] return rgb # ----------------------------------------------------------------------------- # Point cloud export (camera marker + cloud, native-res grid) def _camera_marker(radius: float = 0.04, n_points: int = 800, color=(0.20, 0.55, 0.96)) -> o3d.geometry.PointCloud: """Tiny slate-blue Fibonacci sphere at the world origin.""" i = np.arange(n_points) phi = np.arccos(1 - 2 * (i + 0.5) / n_points) theta = np.pi * (1 + 5 ** 0.5) * (i + 0.5) pts = np.stack([ radius * np.sin(phi) * np.cos(theta), radius * np.sin(phi) * np.sin(theta), radius * np.cos(phi), ], axis=1) pc = o3d.geometry.PointCloud() pc.points = o3d.utility.Vector3dVector(pts.astype(np.float64)) pc.colors = o3d.utility.Vector3dVector(np.tile(color, (n_points, 1)).astype(np.float64)) return pc def _make_ply(image_pil_native: Image.Image, pointmap_hwc: np.ndarray, mask_hw: np.ndarray, max_points: int = 200_000) -> str: """`image_pil_native` MUST already be sized to `pointmap_hwc.shape[:2]` so point colors line up. Output .ply: foreground points + camera marker.""" h, w = pointmap_hwc.shape[:2] image_rgb = np.asarray(image_pil_native.resize((w, h), Image.LANCZOS)) pts = pointmap_hwc.reshape(-1, 3) cols = image_rgb.reshape(-1, 3).astype(np.float32) / 255.0 z = pts[:, 2] finite = np.isfinite(pts).all(axis=1) & (z > 0.05) & (z < 25.0) & mask_hw.reshape(-1) pts, cols = pts[finite], cols[finite] if len(pts) > max_points: idx = np.random.default_rng(0).choice(len(pts), size=max_points, replace=False) pts, cols = pts[idx], cols[idx] pc = o3d.geometry.PointCloud() pc.points = o3d.utility.Vector3dVector(pts.astype(np.float64)) pc.colors = o3d.utility.Vector3dVector(cols.astype(np.float64)) pc += _camera_marker() out_path = tempfile.NamedTemporaryFile(delete=False, suffix=".ply").name o3d.io.write_point_cloud(out_path, pc, write_ascii=False) return out_path # ----------------------------------------------------------------------------- # Gradio handler @spaces.GPU(duration=120) def predict(image: Image.Image, size: str): if image is None: return None, None image_pil = image.convert("RGB") image_bgr = cv2.cvtColor(np.array(image_pil), cv2.COLOR_RGB2BGR) model = _get_pointmap_model(size) pointmap = _estimate_pointmap(image_bgr, model) # (H_n, W_n, 3) at most 1024 in either dim h_n, w_n = pointmap.shape[:2] mask = _foreground_mask(image_pil, h_n, w_n) # native-res mask, fast # Depth heatmap (right pane). Solid mid-grey background with the foreground # turbo-coloured by inverse depth. Mirrors sapiens2 vis_pointmap.py colormap. depth = pointmap[:, :, 2] depth_rgb = _depth_to_rgb(depth, mask) BG_GREY = 200 depth_rgb[~mask] = BG_GREY w0, h0 = image_pil.size depth_pil = Image.fromarray(depth_rgb).resize((w0, h0), Image.LANCZOS) # PLY (download in accordion). Native-res, ≤200K points. ply_path = _make_ply(image_pil, pointmap, mask) return depth_pil, ply_path # ----------------------------------------------------------------------------- # UI EXAMPLES = sorted( os.path.join(ASSETS_DIR, "images", n) for n in os.listdir(os.path.join(ASSETS_DIR, "images")) if n.lower().endswith((".jpg", ".jpeg", ".png")) ) CUSTOM_CSS = """ :root, body, .gradio-container, button, input, select, textarea, .gradio-container *:not(code):not(pre) { font-family: "Helvetica Neue", Helvetica, Arial, sans-serif !important; -webkit-font-smoothing: antialiased; -moz-osx-font-smoothing: grayscale; } #title { text-align: center; font-size: 44px; font-weight: 700; letter-spacing: -0.01em; margin: 28px 0 4px; background: linear-gradient(90deg, #1d4ed8 0%, #6d28d9 50%, #be185d 100%); -webkit-background-clip: text; -webkit-text-fill-color: transparent; background-clip: text; } #subtitle { text-align: center; font-size: 12px; color: #64748b; letter-spacing: 0.18em; margin: 0 0 14px; text-transform: uppercase; font-weight: 500; } #badges { display: flex; justify-content: center; flex-wrap: wrap; gap: 8px; margin: 0 0 32px; } .pill { display: inline-flex; align-items: center; gap: 6px; padding: 7px 14px; border-radius: 999px; background: #f1f5f9; color: #0f172a !important; font-size: 13px; font-weight: 500; letter-spacing: 0.01em; text-decoration: none !important; border: 1px solid #e2e8f0; transition: background 150ms ease, transform 150ms ease, border-color 150ms ease; } .pill:hover { background: #0f172a; color: #f8fafc !important; border-color: #0f172a; transform: translateY(-1px); } .pill svg { width: 14px; height: 14px; } """ HEADER_HTML = """