gaussian_studio / app.py
dgarch424's picture
Upload 21 files
728fc83 verified
Raw
History Blame Contribute Delete
23.2 kB
"""
app.py
──────
Gradio UI for the Image → Background Removal → Depth → 3D Gaussian
Splatting pipeline.
Layout
──────
┌───────────────────────────────────────────────────────┐
│ HEADER / TITLE │
├─── INPUT IMAGE ────────────────────────────────────────┤
│ [Upload Image] │
├─── SETTINGS ACCORDION ────────────────────────────────┤
│ [Stage 1 model ▼] [Stage 2 model ▼] [Stage 3 model ▼]│
│ Custom model ID text inputs + Validate buttons │
├─── STAGE TOGGLES / RUN ───────────────────────────────┤
│ [▶ BG removal] [▶ Depth] [▶ Reconstruction] │
│ [🚀 Run Pipeline] │
├─── OUTPUTS ───────────────────────────────────────────┤
│ BG Removed │ Depth Colourmap │ Depth 16-bit │
│ Log / status │
│ [⬇ Download PLY] │
└───────────────────────────────────────────────────────┘
"""
from __future__ import annotations
# ZeroGPU: must be imported before anything touches torch/CUDA. `import spaces`
# activates a monkey-patch so torch.cuda.is_available() reports True and
# .to("cuda") succeeds at module scope even though no physical GPU is attached
# to this process yet — real GPU access is granted only inside functions
# decorated with @spaces.GPU (see run_pipeline() below). Off-ZeroGPU hardware
# (CPU Basic, local dev, dedicated GPU Spaces) this import is a harmless no-op.
import spaces
import sys
import types
# Patch 1: missing audioop for Python 3.13 / gradio 4.x
if "audioop" not in sys.modules:
sys.modules["audioop"] = types.ModuleType("audioop")
# Patch 2: fix gradio_client bug where schema can be bool instead of dict.
# This causes both TypeError and APIInfoParseError in get_api_info().
# Patch both get_type and _json_schema_to_python_type to guard against non-dict schemas.
import gradio_client.utils as _gcu
_original_get_type = _gcu.get_type # type: ignore[attr-defined]
_original_json_schema_to_python_type = _gcu._json_schema_to_python_type # type: ignore[attr-defined]
def _patched_get_type(schema):
if not isinstance(schema, dict):
return "Any"
return _original_get_type(schema)
def _patched_json_schema_to_python_type(schema, defs=None):
if not isinstance(schema, dict):
return "Any"
return _original_json_schema_to_python_type(schema, defs)
_gcu.get_type = _patched_get_type # type: ignore[attr-defined]
_gcu._json_schema_to_python_type = _patched_json_schema_to_python_type # type: ignore[attr-defined]
import logging
import os
# Patch 3: HF Hub token + faster transfers
# ── On HF Spaces, set HF_TOKEN in Settings → Repository secrets.
# ── hf_transfer is ~3-5x faster for large model downloads; opt-in via env var.
_hf_token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN")
if _hf_token:
os.environ.setdefault("HUGGING_FACE_HUB_TOKEN", _hf_token)
try:
from huggingface_hub import login as _hf_login
_hf_login(token=_hf_token, add_to_git_credential=False)
except Exception:
pass # non-fatal; individual loaders pass token directly
if os.environ.get("HF_HUB_ENABLE_HF_TRANSFER", "").lower() not in ("0", "false", ""):
try:
import hf_transfer # noqa: F401 # speeds up downloads when available
os.environ.setdefault("HF_HUB_ENABLE_HF_TRANSFER", "1")
except ImportError:
pass
import tempfile
from pathlib import Path
import gradio as gr
import numpy as np
from PIL import Image
from configs.model_registry import (
get_display_names,
get_config_by_display_name,
ModelConfig,
BACKGROUND_REMOVAL_MODELS,
DEPTH_ESTIMATION_MODELS,
RECONSTRUCTION_MODELS,
)
from pipeline import SpatialPipeline, PipelineResult
from utils.hf_utils import validate_custom_model
# ── Logging ───────────────────────────────────────────────────────────────────
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)-8s %(name)s %(message)s",
datefmt="%H:%M:%S",
)
logger = logging.getLogger("app")
# ── Pipeline singleton (shared across Gradio requests) ────────────────────────
OUTPUT_DIR = Path("outputs")
OUTPUT_DIR.mkdir(exist_ok=True)
_pipeline = SpatialPipeline(output_dir=OUTPUT_DIR)
# ── Model dropdown options ─────────────────────────────────────────────────────
BG_NAMES = get_display_names("background_removal")
DEPTH_NAMES = get_display_names("depth_estimation")
RECON_NAMES = get_display_names("reconstruction")
CUSTOM_SENTINEL = "✏️ Custom HF model ID / URL"
# ── CSS ───────────────────────────────────────────────────────────────────────
CSS = """
/* ── Global ── */
:root {
--brand-bg: #0f1117;
--brand-surface: #181d27;
--brand-border: #2a3045;
--brand-accent: #5b6ef5;
--brand-accent2: #a78bfa;
--brand-text: #e2e8f0;
--brand-muted: #64748b;
--brand-success: #34d399;
--brand-warn: #fbbf24;
--brand-err: #f87171;
--radius: 10px;
font-family: 'Inter', 'Segoe UI', system-ui, sans-serif;
}
body, .gradio-container {
background: var(--brand-bg) !important;
color: var(--brand-text) !important;
}
/* Header */
#header-md h1 { font-size: 2rem; font-weight: 700; letter-spacing: -0.03em; margin-bottom: 0.2rem; }
#header-md p { color: var(--brand-muted); font-size: 0.95rem; margin: 0; }
#header-md span.accent { color: var(--brand-accent2); }
/* Stage badges */
.stage-badge {
display: inline-block;
padding: 2px 10px;
border-radius: 20px;
font-size: 0.78rem;
font-weight: 600;
letter-spacing: 0.04em;
text-transform: uppercase;
margin-right: 6px;
}
.s1 { background: #1e3a5f; color: #93c5fd; }
.s2 { background: #1a3a2a; color: #6ee7b7; }
.s3 { background: #3a1e5f; color: #c4b5fd; }
/* Run button */
#run-btn { background: var(--brand-accent) !important; color: #fff !important; font-weight: 700 !important; }
#run-btn:hover { background: #4255e0 !important; }
/* Log box */
#log-box textarea {
font-family: 'JetBrains Mono', 'Fira Code', monospace !important;
font-size: 0.78rem !important;
background: #0a0d14 !important;
color: #a3e635 !important;
}
/* Output image labels */
.output-label { font-size: 0.8rem; color: var(--brand-muted); text-transform: uppercase; letter-spacing: 0.06em; }
"""
# ── Helpers ────────────────────────────────────────────────────────────────────
def _resolve_config(stage: str, dropdown_val: str, custom_id: str) -> ModelConfig | None:
"""Return a ModelConfig, substituting custom_id when the sentinel is selected."""
if dropdown_val == CUSTOM_SENTINEL:
if not custom_id or not custom_id.strip():
return None
# Find the custom slot in the registry and patch its model_id
cfg = get_config_by_display_name(stage, CUSTOM_SENTINEL) # type: ignore[arg-type]
if cfg is None:
return None
from dataclasses import replace
return replace(cfg, model_id=custom_id.strip())
return get_config_by_display_name(stage, dropdown_val) # type: ignore[arg-type]
def _log_lines(*args) -> str:
return "\n".join(str(a) for a in args if a)
# ── Core run function ──────────────────────────────────────────────────────────
# ZeroGPU: this is the single Gradio-bound entry point (see run_btn.click(fn=run_pipeline, ...)
# below), so it's the right place to decorate — one GPU slot is requested for the whole
# bg-removal → depth → reconstruction run, not one per stage (each entry into a
# @spaces.GPU function pays a process-fork + CUDA-reattach cost, so decorating each stage
# separately would be both slower and more likely to lose the GPU mid-pipeline).
#
# duration=120 covers a full 3-stage run, including a cold model load/download the first
# time a stage's model changes (that load happens inside this call, since models are
# swapped on demand from the dropdowns/custom IDs — see SpatialPipeline._get_loader).
# Tune this to what you observe in your Space's logs: raise it if large/custom models
# time out, lower it if your runs are consistently short (shorter duration = higher queue
# priority). See https://huggingface.co/docs/hub/spaces-zerogpu#duration-management
@spaces.GPU(duration=120)
def run_pipeline(
input_image,
bg_dropdown: str,
bg_custom: str,
depth_dropdown: str,
depth_custom: str,
recon_dropdown: str,
recon_custom: str,
run_bg: bool,
run_depth: bool,
run_recon: bool,
progress=gr.Progress(track_tqdm=True),
) -> tuple:
"""
Main Gradio handler. Returns a tuple matching the .outputs list in the UI.
Order: (bg_removed_image, depth_colour, depth_16, ply_file, log_text)
"""
log = []
def emit(msg: str):
log.append(msg)
if input_image is None:
return None, None, None, None, "⚠️ Please upload an image to begin."
if isinstance(input_image, np.ndarray):
pil_input = Image.fromarray(input_image)
elif isinstance(input_image, Image.Image):
pil_input = input_image
else:
return None, None, None, None, "⚠️ Unsupported image input."
# --- Validate stage selection ---
stages = []
if run_bg: stages.append("bgremove")
if run_depth: stages.append("depth")
if run_recon: stages.append("recon")
if not stages:
return None, None, None, None, "⚠️ Please enable at least one stage."
# --- Resolve model configs ---
bg_cfg = _resolve_config("background_removal", bg_dropdown, bg_custom)
depth_cfg = _resolve_config("depth_estimation", depth_dropdown, depth_custom)
recon_cfg = _resolve_config("reconstruction", recon_dropdown, recon_custom)
if "bgremove" in stages and bg_cfg is None:
return None, None, None, None, "❌ Background removal: no valid model selected."
if "depth" in stages and depth_cfg is None:
return None, None, None, None, "❌ Depth estimation: no valid model selected."
if "recon" in stages and recon_cfg is None:
return None, None, None, None, "❌ Reconstruction: no valid model selected."
# Use registry defaults if a stage is skipped (needed for type-safety)
bg_cfg = bg_cfg or get_config_by_display_name("background_removal", BG_NAMES[0])
depth_cfg = depth_cfg or get_config_by_display_name("depth_estimation", DEPTH_NAMES[0])
recon_cfg = recon_cfg or get_config_by_display_name("reconstruction", RECON_NAMES[0])
emit(f"🚀 Starting pipeline | stages: {', '.join(stages)}")
emit(f" BG removal: {bg_cfg.display_name}")
emit(f" Depth: {depth_cfg.display_name}")
emit(f" Recon: {recon_cfg.display_name}")
# Progress relay
def on_progress(stage: str, message: str):
emit(f"[{stage.upper()}] {message}")
progress(0, desc=message)
_pipeline.progress_callback = on_progress
try:
result: PipelineResult = _pipeline.run(
input_image=pil_input,
bgremove_config=bg_cfg,
depth_config=depth_cfg,
recon_config=recon_cfg,
run_stages=tuple(stages),
)
except Exception as exc:
logger.exception("Pipeline crashed")
emit(f"💥 Pipeline crashed: {exc}")
return None, None, None, None, "\n".join(log)
if result.errors:
for e in result.errors:
emit(f"❌ {e}")
return None, None, None, None, "\n".join(log)
# --- Build timing summary ---
emit("")
emit("─── Results ─────────────────────────────────────")
if result.bgremove_elapsed: emit(f" Stage 1 {result.bgremove_elapsed:.1f}s model={result.bgremove_model}")
if result.depth_elapsed: emit(f" Stage 2 {result.depth_elapsed:.1f}s model={result.depth_model}")
if result.recon_elapsed: emit(f" Stage 3 {result.recon_elapsed:.1f}s points={result.point_count:,}")
emit(f" Total {result.total_elapsed:.1f}s")
if result.ply_path: emit(f" PLY {result.ply_path}")
# --- Package outputs ---
ply_file = result.ply_path if result.ply_path and Path(result.ply_path).exists() else None
bg_preview = result.bg_rgba if result.bg_rgba is not None else result.input_image
return (
bg_preview, # background-removed preview (RGBA), or raw input if stage skipped
result.depth_colourmap, # depth false-colour (PIL)
result.depth_uint16, # depth 16-bit (PIL)
ply_file, # path string or None
"\n".join(log),
)
def validate_model_id(model_id: str) -> str:
ok, msg = validate_custom_model(model_id)
return msg
# ── Gradio UI ──────────────────────────────────────────────────────────────────
def build_ui() -> gr.Blocks:
with gr.Blocks(css=CSS, title="Image → 3DGS Pipeline") as demo:
# ── Header ────────────────────────────────────────────────────────────
gr.HTML("""
<div id="header-md" style="padding:1.5rem 0 0.5rem;">
<h1>🌐 Image → <span class="accent">3D Gaussian Splatting</span></h1>
<p>Upload a photo → remove the background → estimate dense depth → export a 3D point cloud or Gaussian splat scaffold.</p>
</div>
""")
# ── Input image ───────────────────────────────────────────────────────
input_image_upload = gr.Image(label="📷 Input Image", type="pil")
# ── Model selection ────────────────────────────────────────────────────
with gr.Accordion("⚙️ Model Selection", open=True):
gr.HTML("""
<p style="color:#64748b;font-size:0.85rem;margin:0 0 1rem;">
Select a preset model for each stage, or choose <em>Custom</em> and paste any
HuggingFace model ID (e.g. <code>ZhengPeng7/BiRefNet</code>) or direct HTTPS URL.
</p>""")
with gr.Row():
# Stage 1
with gr.Column():
gr.HTML('<span class="stage-badge s1">Stage 1</span><strong>Background Removal</strong>')
bg_dropdown = gr.Dropdown(
choices=BG_NAMES,
value=BG_NAMES[0],
label="Background-removal model",
interactive=True,
)
bg_custom = gr.Textbox(
label="Custom model ID or URL",
placeholder="org/model-name or https://…",
visible=False,
)
bg_validate_btn = gr.Button("🔍 Validate", size="sm", visible=False)
bg_validate_out = gr.Textbox(label="", lines=1, interactive=False, visible=False)
# Stage 2
with gr.Column():
gr.HTML('<span class="stage-badge s2">Stage 2</span><strong>Image → Depth</strong>')
depth_dropdown = gr.Dropdown(
choices=DEPTH_NAMES,
value=DEPTH_NAMES[0],
label="Depth estimation model",
interactive=True,
)
depth_custom = gr.Textbox(
label="Custom model ID or URL",
placeholder="org/model-name or https://…",
visible=False,
)
depth_validate_btn = gr.Button("🔍 Validate", size="sm", visible=False)
depth_validate_out = gr.Textbox(label="", lines=1, interactive=False, visible=False)
# Stage 3
with gr.Column():
gr.HTML('<span class="stage-badge s3">Stage 3</span><strong>RGBD → 3D</strong>')
recon_dropdown = gr.Dropdown(
choices=RECON_NAMES,
value=RECON_NAMES[0],
label="Reconstruction method",
interactive=True,
)
recon_custom = gr.Textbox(
label="Custom model ID or URL",
placeholder="org/model-name or https://…",
visible=False,
)
recon_validate_btn = gr.Button("🔍 Validate", size="sm", visible=False)
recon_validate_out = gr.Textbox(label="", lines=1, interactive=False, visible=False)
# Show/hide custom input on sentinel selection
bg_dropdown.change(
lambda v: (gr.update(visible=v == CUSTOM_SENTINEL),
gr.update(visible=v == CUSTOM_SENTINEL),
gr.update(visible=v == CUSTOM_SENTINEL)),
inputs=[bg_dropdown],
outputs=[bg_custom, bg_validate_btn, bg_validate_out],
)
depth_dropdown.change(
lambda v: (gr.update(visible=v == CUSTOM_SENTINEL),
gr.update(visible=v == CUSTOM_SENTINEL),
gr.update(visible=v == CUSTOM_SENTINEL)),
inputs=[depth_dropdown],
outputs=[depth_custom, depth_validate_btn, depth_validate_out],
)
recon_dropdown.change(
lambda v: (gr.update(visible=v == CUSTOM_SENTINEL),
gr.update(visible=v == CUSTOM_SENTINEL),
gr.update(visible=v == CUSTOM_SENTINEL)),
inputs=[recon_dropdown],
outputs=[recon_custom, recon_validate_btn, recon_validate_out],
)
bg_validate_btn.click(validate_model_id, inputs=[bg_custom], outputs=[bg_validate_out])
depth_validate_btn.click(validate_model_id, inputs=[depth_custom], outputs=[depth_validate_out])
recon_validate_btn.click(validate_model_id, inputs=[recon_custom], outputs=[recon_validate_out])
# ── Stage enable toggles ─────────────────────────────────────────────
with gr.Row():
run_bg_chk = gr.Checkbox(value=True, label="▶ Stage 1: Background removal")
run_depth_chk = gr.Checkbox(value=True, label="▶ Stage 2: Depth estimation")
run_recon_chk = gr.Checkbox(value=True, label="▶ Stage 3: 3D Reconstruction")
run_btn = gr.Button("🚀 Run Pipeline", variant="primary", elem_id="run-btn")
# ── Outputs ────────────────────────────────────────────────────────────
with gr.Row():
out_bg = gr.Image(label="Background Removed", type="pil", interactive=False)
out_depth_col = gr.Image(label="Depth Map (colour)", type="pil", interactive=False)
out_depth_16 = gr.Image(label="Depth Map (16-bit)", type="pil", interactive=False)
with gr.Row():
out_ply = gr.File(label="⬇ Download Point Cloud (.ply)")
log_box = gr.Textbox(
label="Pipeline log",
lines=10,
interactive=False,
elem_id="log-box",
)
# ── Wire up ────────────────────────────────────────────────────────────
run_btn.click(
fn=run_pipeline,
inputs=[
input_image_upload,
bg_dropdown, bg_custom,
depth_dropdown, depth_custom,
recon_dropdown, recon_custom,
run_bg_chk, run_depth_chk, run_recon_chk,
],
outputs=[out_bg, out_depth_col, out_depth_16, out_ply, log_box],
)
# ── Footer ─────────────────────────────────────────────────────────────
gr.HTML("""
<div style="text-align:center;color:#64748b;font-size:0.8rem;padding:1rem 0;">
Background removal uses
<a href="https://huggingface.co/ZhengPeng7/BiRefNet" style="color:#a78bfa;">BiRefNet</a>
(MIT) — the same model family
<a href="https://huggingface.co/spaces/VAST-AI/TripoSplat" style="color:#a78bfa;">TripoSplat</a>
uses for its own foreground matting stage.
Models run locally on this Space's hardware.
PLY files are compatible with
<a href="https://github.com/graphdeco-inria/gaussian-splatting" style="color:#a78bfa;">
graphdeco-inria/gaussian-splatting</a> and
<a href="https://github.com/antimatter15/splat" style="color:#a78bfa;">antimatter15/splat</a>.
</div>
""")
return demo
# ── Entry point ────────────────────────────────────────────────────────────────
if __name__ == "__main__":
demo = build_ui()
demo.queue(max_size=3)
demo.launch(
show_error=True,
share=False,
)