PXDepth-Demo / app.py
sysu111's picture
Add PXDepth demo
b308d33
Raw
History Blame Contribute Delete
16.6 kB
"""Hugging Face Gradio Space for PXDepth."""
from __future__ import annotations
import shutil
import tempfile
import time
from pathlib import Path
from typing import Optional
# ZeroGPU patches torch during import, so spaces must be imported first.
try:
import spaces
gpu = spaces.GPU(duration=90)
except ImportError:
gpu = lambda fn: fn
import gradio as gr
import numpy as np
import torch
import torch.nn.functional as F
import utils3d
from PIL import Image
from pxdepth.inference import area_size_from_area, resize_image, resize_map
from pxdepth.model import PXDepth
from pxdepth.utils.ply import write_point_cloud_ply
from pxdepth.utils.vis import colorize_depth
PXDEPTH_REPO = "yuanzhy29/PXDepth"
MOGE2_REPO = "Ruicheng/moge-2-vitl-normal"
PXDEPTH_SIZE = (1022, 770)
MOGE2_TOKEN_AREA = 1200
MOGE2_PATCH_SIZE = 14
MAX_INPUT_PIXELS = 12_000_000
OUTPUT_MAX_AGE = 60 * 60
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
CSS = """
html, body {
height: auto !important;
min-height: 100% !important;
overflow-y: auto !important;
overscroll-behavior-y: auto !important;
}
.gradio-container {
height: auto !important;
min-height: 100vh !important;
overflow: visible !important;
}
#pxdepth-demo { max-width: 1280px; margin: 0 auto; }
#img-display-input, #img-display-output { max-height: 72vh; }
#img-display-output img { object-fit: contain !important; }
#model-3d { min-height: 55vh; }
#examples-strip .gallery {
flex-wrap: nowrap !important;
overflow-x: auto;
overflow-y: hidden;
padding-bottom: 0.5rem;
scroll-behavior: smooth;
scroll-snap-type: x proximity;
scrollbar-width: thin;
-webkit-overflow-scrolling: touch;
}
#examples-strip .gallery-item {
flex: 0 0 auto;
scroll-snap-align: start;
}
"""
PAGE_JS = """
() => {
document.documentElement.style.overflowY = "auto";
document.body.style.overflowY = "auto";
const install = () => {
const viewer = document.querySelector("#model-3d");
if (viewer && viewer.dataset.pageWheel !== "true") {
viewer.dataset.pageWheel = "true";
viewer.addEventListener("wheel", (event) => {
if (event.ctrlKey || event.metaKey) return;
event.preventDefault();
event.stopImmediatePropagation();
window.scrollBy({ top: event.deltaY, left: 0, behavior: "auto" });
}, { passive: false, capture: true });
}
const examples = document.querySelector("#examples-strip .gallery");
if (examples && examples.dataset.horizontalWheel !== "true") {
examples.dataset.horizontalWheel = "true";
examples.addEventListener("wheel", (event) => {
if (event.ctrlKey || event.metaKey) return;
if (Math.abs(event.deltaY) <= Math.abs(event.deltaX)) return;
event.preventDefault();
examples.scrollLeft += event.deltaY;
}, { passive: false });
}
};
install();
new MutationObserver(install).observe(document.body, { childList: true, subtree: true });
}
"""
def load_model() -> PXDepth:
"""Load PXDepth and its MoGe-2 metric-scale reference once at startup."""
print("Loading PXDepth...")
model = PXDepth.from_pretrained(PXDEPTH_REPO, strict=True).eval()
try:
from moge.model.v2 import MoGeModel
except ImportError as exc:
raise RuntimeError(
"MoGe-2 is required by this demo. Check the Space requirements."
) from exc
print("Loading MoGe-2...")
model._reference_model = MoGeModel.from_pretrained(MOGE2_REPO).eval()
model = model.to(DEVICE).eval()
print(f"Models loaded on {DEVICE}.")
return model
MODEL = load_model()
def resize_for_tokens(image: torch.Tensor, tokens: int, patch: int) -> torch.Tensor:
"""Preserve aspect ratio and resize an RGB tensor to a patch-token area."""
height, width = area_size_from_area(
image.shape[-2],
image.shape[-1],
tokens * patch * patch,
patch,
)
if (height, width) == tuple(image.shape[-2:]):
return image
return F.interpolate(
image.unsqueeze(0),
(height, width),
mode="bilinear",
align_corners=False,
)[0]
def cleanup_outputs(root: Path) -> None:
"""Remove stale per-session files from the Space's ephemeral storage."""
if not root.exists():
return
cutoff = time.time() - OUTPUT_MAX_AGE
for path in root.iterdir():
try:
if path.is_dir() and path.stat().st_mtime < cutoff:
shutil.rmtree(path, ignore_errors=True)
except OSError:
continue
def session_dir(request: Optional[gr.Request]) -> Path:
"""Create a clean output directory for the current browser session."""
session = getattr(request, "session_hash", None) or "local"
session = "".join(char for char in session if char.isalnum() or char in "-_")
root = Path(tempfile.gettempdir()) / "pxdepth-demo"
root.mkdir(parents=True, exist_ok=True)
cleanup_outputs(root)
output = root / (session or "local")
shutil.rmtree(output, ignore_errors=True)
output.mkdir(parents=True, exist_ok=True)
return output
def sample_points(
points: np.ndarray,
colors: np.ndarray,
max_points: int,
) -> tuple[np.ndarray, np.ndarray]:
"""Deterministically subsample a point cloud for browser rendering."""
if points.shape[0] <= max_points:
return points, colors
indices = np.linspace(0, points.shape[0] - 1, max_points, dtype=np.int64)
return points[indices], colors[indices]
def filter_flying_points(
points: np.ndarray,
colors: np.ndarray,
neighbors: int = 30,
std_ratio: float = 2.0,
) -> tuple[np.ndarray, np.ndarray]:
"""Remove sparse statistical outliers from an already sampled cloud."""
if points.shape[0] <= neighbors + 1:
return points, colors
from scipy.spatial import cKDTree
tree = cKDTree(points.astype(np.float64, copy=False))
mean_distance = np.empty(points.shape[0], dtype=np.float32)
for start in range(0, points.shape[0], 100_000):
stop = min(start + 100_000, points.shape[0])
try:
distances, _ = tree.query(
points[start:stop],
k=neighbors + 1,
workers=-1,
)
except TypeError:
distances, _ = tree.query(points[start:stop], k=neighbors + 1)
mean_distance[start:stop] = np.asarray(
distances[:, 1:],
dtype=np.float32,
).mean(axis=1)
finite = np.isfinite(mean_distance)
if not finite.any():
return points, colors
values = mean_distance[finite]
threshold = float(values.mean() + std_ratio * values.std())
keep = finite & (mean_distance <= threshold)
return (points[keep], colors[keep]) if keep.any() else (points, colors)
def write_viewer_glb(
path: Path,
points: np.ndarray,
colors: np.ndarray,
) -> None:
"""Write the browser point cloud using the stable GLB viewer path."""
import trimesh
display_points = points * np.array([1.0, -1.0, -1.0], dtype=np.float32)
trimesh.PointCloud(display_points, colors=colors).export(path)
def update_viewer(
cache_path: Optional[str],
filter_points: bool,
max_points: int,
) -> Optional[str]:
"""Rebuild the viewer from cached points without running either model."""
if not cache_path or not Path(cache_path).is_file():
return None
with np.load(cache_path) as cache:
points = cache["points"]
colors = cache["colors"]
points, colors = sample_points(points, colors, int(max_points))
if filter_points:
points, colors = filter_flying_points(points, colors)
if points.shape[0] == 0:
raise gr.Error("No points remain after filtering.")
cache_file = Path(cache_path)
tag = f"{int(max_points)}_{int(filter_points)}"
viewer_path = cache_file.with_name(f"pointcloud_viewer_{tag}.glb")
write_viewer_glb(viewer_path, points, colors)
for old_path in cache_file.parent.glob("pointcloud_viewer_*.*"):
if old_path != viewer_path:
old_path.unlink(missing_ok=True)
return str(viewer_path)
@gpu
@torch.inference_mode()
def predict_gpu(image: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Run only model inference while holding the ZeroGPU allocation."""
tensor = (
torch.from_numpy(image.copy())
.to(device=DEVICE, dtype=torch.float32)
.permute(2, 0, 1)
/ 255.0
)
model_image, _ = resize_image(
tensor,
PXDEPTH_SIZE,
True,
MODEL.patch_size,
)
reference_image = resize_for_tokens(
tensor,
MOGE2_TOKEN_AREA,
MOGE2_PATCH_SIZE,
)
result = MODEL.infer(
model_image,
ref_image=reference_image,
apply_mask=False,
use_fp16=DEVICE.type == "cuda",
use_fp32=DEVICE.type != "cuda",
)
return (
result["depth"].float().cpu().numpy(),
result["mask"].cpu().numpy(),
result["intrinsics"].float().cpu().numpy(),
)
def on_submit(
image: Optional[np.ndarray],
apply_mask: bool,
filter_points: bool,
max_points: int,
request: gr.Request,
):
"""Run inference, build visualizations, and export downloadable files."""
if image is None:
raise gr.Error("Please upload an image first.")
if image.ndim != 3 or image.shape[-1] < 3:
raise gr.Error("The input must be an RGB image.")
if image.shape[0] * image.shape[1] > MAX_INPUT_PIXELS:
raise gr.Error(
"The uploaded image is too large. Please use an image below 12 megapixels."
)
image = np.ascontiguousarray(image[..., :3].astype(np.uint8))
original_size = image.shape[:2]
depth_raw, mask_raw, intrinsics_np = predict_gpu(image)
# Restore outputs and reconstruct the point map on CPU so ZeroGPU is held
# only for neural-network inference.
depth = resize_map(torch.from_numpy(depth_raw), original_size).float()
mask = resize_map(torch.from_numpy(mask_raw), original_size, is_mask=True)
intrinsics = torch.from_numpy(intrinsics_np).float()
finite = torch.isfinite(depth) & (depth > 0)
valid = finite & mask if apply_mask else finite
points = utils3d.pt.depth_map_to_point_map(
torch.where(finite, depth, torch.zeros_like(depth)),
intrinsics=intrinsics,
)
depth_np = depth.numpy().astype(np.float32)
mask_np = mask.numpy().astype(bool)
valid_np = valid.numpy().astype(bool)
depth_vis = colorize_depth(np.where(mask_np, depth_np, np.inf), mask=None)
output = session_dir(request)
depth_npy = output / "depth.npy"
depth_png = output / "depth.png"
mask_png = output / "mask.png"
ply_path = output / "pointcloud.ply"
cache_path = output / "viewer_data.npz"
np.save(depth_npy, depth_np)
Image.fromarray(depth_vis).save(depth_png)
Image.fromarray(mask_np.astype(np.uint8) * 255, mode="L").save(mask_png)
points_np = points.numpy().reshape(-1, 3)
colors_np = image.reshape(-1, 3).astype(np.float32) / 255.0
keep = valid_np.reshape(-1) & np.isfinite(points_np).all(axis=1)
points_full, colors_full = points_np[keep], colors_np[keep]
if points_full.shape[0] == 0:
raise gr.Error("No valid 3D points were produced for this image.")
write_point_cloud_ply(ply_path, points_full, colors_full)
colors_uint8 = np.clip(colors_full * 255.0, 0, 255).astype(np.uint8)
np.savez(cache_path, points=points_full.astype(np.float32), colors=colors_uint8)
viewer_path = update_viewer(
str(cache_path),
filter_points,
max_points,
)
files = [str(depth_png), str(depth_npy), str(mask_png), str(ply_path)]
return (image, depth_vis), viewer_path, files, str(cache_path)
def build_demo() -> gr.Blocks:
"""Construct the public Gradio interface."""
description = """
Official demo for **PXDepth: Pixel-Space Modeling for Structure Preserving Monocular Depth Estimation**.
See the [paper](https://arxiv.org/abs/2608.16984),
[project page](https://yuanzhy29.github.io/PXDepth-Page/), and
[GitHub repository](https://github.com/yuanzhy29/PXDepth).
"""
with gr.Blocks(theme=gr.themes.Soft(), css=CSS, js=PAGE_JS) as demo:
viewer_cache = gr.State(value=None)
with gr.Column(elem_id="pxdepth-demo"):
gr.Markdown("# PXDepth")
gr.Markdown(description)
gr.Markdown("### Point Cloud & Depth Prediction Demo")
with gr.Row():
with gr.Column():
input_image = gr.Image(
label="Input Image",
image_mode="RGB",
type="numpy",
placeholder="# Drop an image here\n— or —\nClick to upload",
elem_id="img-display-input",
)
with gr.Accordion(label="Settings", open=False):
apply_mask = gr.Checkbox(
label="Apply valid-depth mask to point cloud",
value=True,
)
filter_points = gr.Checkbox(
label="Filter Flying Points",
info="Statistical outlier filtering; does not rerun the model.",
value=False,
)
max_points = gr.Slider(
50_000,
500_000,
value=200_000,
step=50_000,
label="3D Viewer Max Points",
info="Updates only the viewer; the downloaded PLY retains all valid points.",
)
submit = gr.Button("Predict", variant="primary")
with gr.Column():
with gr.Tabs():
with gr.Tab("3D View"):
model_3d = gr.Model3D(
label="3D Point Map",
clear_color=(1.0, 1.0, 1.0, 1.0),
height="55vh",
elem_id="model-3d",
)
with gr.Tab("Depth"):
depth_map = gr.ImageSlider(
label="RGB / Depth",
image_mode="RGB",
type="numpy",
slider_position=50,
elem_id="img-display-output",
)
with gr.Tab("Download"):
downloads = gr.File(
label="Download Files",
file_count="multiple",
type="filepath",
)
examples = Path("example_images")
example_files = (
sorted(
str(path)
for path in examples.iterdir()
if path.suffix.lower() in {".jpg", ".jpeg", ".png", ".webp"}
)
if examples.exists()
else []
)
if example_files:
gr.Examples(
example_files,
input_image,
cache_examples=False,
examples_per_page=len(example_files),
elem_id="examples-strip",
)
submit.click(
on_submit,
[input_image, apply_mask, filter_points, max_points],
[depth_map, model_3d, downloads, viewer_cache],
show_progress="full",
concurrency_limit=1,
)
viewer_inputs = [viewer_cache, filter_points, max_points]
filter_points.change(
update_viewer,
viewer_inputs,
model_3d,
show_progress="minimal",
)
max_points.release(
update_viewer,
viewer_inputs,
model_3d,
show_progress="minimal",
)
return demo
demo = build_demo()
if __name__ == "__main__":
demo.queue(default_concurrency_limit=1).launch()