Spaces:
Running on Zero
Running on Zero
Download app.py from ricklon/InfiniSplat: direct link, hf CLI and curl.
- Browser
- Download file 12.6 kB
-
https://huggingface.co/spaces/ricklon/InfiniSplat/resolve/main/app.py
- Command line
-
hf download hf://spaces/ricklon/InfiniSplat/app.py
-
curl -L -o app.py https://huggingface.co/spaces/ricklon/InfiniSplat/resolve/main/app.py
12.6 kB
| from __future__ import annotations | |
| import asyncio | |
| import contextlib | |
| import json | |
| import math | |
| import os | |
| import shutil | |
| import time | |
| import uuid | |
| from pathlib import Path | |
| from urllib.parse import quote | |
| os.environ.setdefault("GRADIO_SSR_MODE", "false") | |
| import spaces # noqa: E402, F401 - must patch torch before model modules are imported | |
| from fastapi import HTTPException | |
| from fastapi.responses import FileResponse | |
| from gradio import Server | |
| from gradio.data_classes import FileData | |
| from gradio.utils import get_upload_folder | |
| from PIL import Image, ImageOps | |
| from src.demo.hf_runtime import ( | |
| InfiniSplatRuntime, | |
| export_browser_viewer, | |
| export_filtered_gaussian_ply, | |
| export_standalone_viewer, | |
| prepare_viewer_template, | |
| ) | |
| GPU_DURATION_SECONDS = 6 | |
| REQUEST_CACHE_SECONDS = 3600 | |
| CACHE_CLEANUP_INTERVAL_SECONDS = 3600 | |
| SUPPORTED_IMAGE_SUFFIXES = {".jpg", ".jpeg", ".png", ".webp"} | |
| EXAMPLE_FILENAMES = [ | |
| "painting_room.jpg", | |
| "summer_room.jpg", | |
| "animate_room.jpg", | |
| "meerkat.jpg", | |
| "bedroom.jpg", | |
| "my_bedroom.JPG", | |
| "beggar_home.jpg", | |
| "cave_ai.jpg", | |
| "eth3d_courtyard.png", | |
| "dragon_ball.jpg", | |
| "flower_room.jpg", | |
| "maksim-shutov-unsplash.jpg", | |
| "ghibli_realroom.jpg", | |
| "ghibli_room.jpg", | |
| "pexels-masi.jpg", | |
| "gym.png", | |
| "living_room.jpg", | |
| "scannetpp_fe94fc30cf.JPG", | |
| "loft_room.jpg", | |
| "old_livingroom.png", | |
| "sofa_ai.jpg", | |
| ] | |
| EXAMPLE_SPANS = [4, 4, 4, 3, 4, 5, 4, 5, 3, 3, 3, 5, 4, 5, 3, 5, 4, 5, 3, 3, 5] | |
| REPO_ROOT = Path(__file__).resolve().parent | |
| FRONTEND_ROOT = REPO_ROOT / "src/demo/server_frontend" | |
| UPLOAD_ROOT = Path(get_upload_folder()).resolve() | |
| OUTPUT_ROOT = UPLOAD_ROOT / "infinisplat" | |
| OUTPUT_ROOT.mkdir(parents=True, exist_ok=True) | |
| EXAMPLE_ROOT = Path( | |
| os.environ.get("INFINISPLAT_EXAMPLE_DIR", REPO_ROOT / "examples/data/rgb_demo") | |
| ).resolve() | |
| EXAMPLE_THUMBNAIL_ROOT = OUTPUT_ROOT / "_example_thumbnails" | |
| def _log_timing(stage: str, request_dir: Path, started_at: float, **metrics: int) -> None: | |
| payload = { | |
| "stage": stage, | |
| "request": request_dir.name, | |
| "seconds": round(time.perf_counter() - started_at, 3), | |
| **metrics, | |
| } | |
| print(f"INFINISPLAT_TIMING {json.dumps(payload, sort_keys=True)}", flush=True) | |
| def _cleanup_request_directories(max_age_seconds: int | None) -> None: | |
| if not OUTPUT_ROOT.is_dir(): | |
| return | |
| cutoff = None if max_age_seconds is None else time.time() - max_age_seconds | |
| for request_dir in OUTPUT_ROOT.iterdir(): | |
| if request_dir.is_symlink() or not request_dir.is_dir(): | |
| continue | |
| try: | |
| if uuid.UUID(hex=request_dir.name).hex != request_dir.name: | |
| continue | |
| except ValueError: | |
| continue | |
| if cutoff is not None and request_dir.stat().st_mtime > cutoff: | |
| continue | |
| with contextlib.suppress(FileNotFoundError): | |
| shutil.rmtree(request_dir) | |
| async def _cleanup_loop() -> None: | |
| while True: | |
| await asyncio.sleep(CACHE_CLEANUP_INTERVAL_SECONDS) | |
| await asyncio.to_thread(_cleanup_request_directories, REQUEST_CACHE_SECONDS) | |
| async def _lifespan(_app: Server): | |
| await asyncio.to_thread(_cleanup_request_directories, REQUEST_CACHE_SECONDS) | |
| cleanup_task = asyncio.create_task(_cleanup_loop()) | |
| try: | |
| yield | |
| finally: | |
| cleanup_task.cancel() | |
| with contextlib.suppress(asyncio.CancelledError): | |
| await cleanup_task | |
| await asyncio.to_thread(_cleanup_request_directories, None) | |
| def _prepare_examples() -> list[dict[str, str | int]]: | |
| missing = [name for name in EXAMPLE_FILENAMES if not (EXAMPLE_ROOT / name).is_file()] | |
| if missing: | |
| raise FileNotFoundError(f"Missing example images in {EXAMPLE_ROOT}: {missing}") | |
| EXAMPLE_THUMBNAIL_ROOT.mkdir(parents=True, exist_ok=True) | |
| examples = [] | |
| for index, (filename, span) in enumerate(zip(EXAMPLE_FILENAMES, EXAMPLE_SPANS)): | |
| source = EXAMPLE_ROOT / filename | |
| thumbnail = EXAMPLE_THUMBNAIL_ROOT / f"{index:02d}-{source.stem}.webp" | |
| if not thumbnail.is_file() or thumbnail.stat().st_mtime < source.stat().st_mtime: | |
| with Image.open(source) as image: | |
| image = ImageOps.exif_transpose(image).convert("RGB") | |
| image.thumbnail((480, 320), Image.Resampling.LANCZOS) | |
| image.save(thumbnail, format="WEBP", quality=82, method=6) | |
| examples.append( | |
| { | |
| "index": index, | |
| "name": source.stem.replace("_", " "), | |
| "filename": filename, | |
| "span": span, | |
| "thumbnail_url": f"/infinisplat/examples/{index}/thumbnail", | |
| "source_url": f"/infinisplat/examples/{index}/source", | |
| } | |
| ) | |
| return examples | |
| def _file_url(path: Path) -> str: | |
| return f"/gradio_api/file={quote(str(path.resolve()), safe='/')}" | |
| def _request_dir(request_id: str) -> Path: | |
| try: | |
| normalized = uuid.UUID(hex=request_id).hex | |
| except (ValueError, AttributeError) as error: | |
| raise ValueError("Invalid reconstruction request ID.") from error | |
| if normalized != request_id: | |
| raise ValueError("Invalid reconstruction request ID.") | |
| request_dir = (OUTPUT_ROOT / normalized).resolve() | |
| if request_dir.parent != OUTPUT_ROOT: | |
| raise ValueError("Invalid reconstruction request directory.") | |
| if not request_dir.is_dir(): | |
| raise FileNotFoundError("The reconstruction result has expired.") | |
| request_dir.touch() | |
| return request_dir | |
| def _uploaded_image_path(image: FileData | dict | str) -> Path: | |
| if isinstance(image, str): | |
| path_value = image | |
| elif isinstance(image, dict): | |
| path_value = image.get("path") | |
| else: | |
| path_value = getattr(image, "path", None) | |
| if not path_value: | |
| raise ValueError("Please upload an image.") | |
| image_path = Path(path_value).resolve() | |
| if not image_path.is_file() or image_path.suffix.lower() not in SUPPORTED_IMAGE_SUFFIXES: | |
| raise ValueError("Please upload a JPG, PNG, or WebP image.") | |
| try: | |
| image_path.relative_to(UPLOAD_ROOT) | |
| except ValueError as error: | |
| raise ValueError("The uploaded image is outside the Gradio upload directory.") from error | |
| return image_path | |
| runtime = InfiniSplatRuntime.load() | |
| viewer_template = prepare_viewer_template(OUTPUT_ROOT) | |
| examples = _prepare_examples() | |
| app = Server( | |
| title="InfiniSplat", | |
| description="Implicit Gaussian decoding for large-baseline monocular view synthesis.", | |
| lifespan=_lifespan, | |
| ) | |
| def _run_reconstruction(image_path: Path, artifact_path: Path) -> str: | |
| """Run one GPU reconstruction and persist only CPU-resident tensors.""" | |
| return str(runtime.infer_to_artifact(image_path=image_path, artifact_path=artifact_path)) | |
| def reconstruct(image: FileData) -> dict[str, str | int]: | |
| """Reconstruct a 3D Gaussian scene from one RGB image.""" | |
| _cleanup_request_directories(REQUEST_CACHE_SECONDS) | |
| request_dir = OUTPUT_ROOT / uuid.uuid4().hex | |
| request_dir.mkdir(parents=True, exist_ok=False) | |
| started_at = time.perf_counter() | |
| artifact = Path( | |
| _run_reconstruction( | |
| _uploaded_image_path(image), | |
| request_dir / "gaussians.pt", | |
| ) | |
| ) | |
| _log_timing( | |
| "gpu_reconstruct", | |
| request_dir, | |
| started_at, | |
| artifact_bytes=artifact.stat().st_size, | |
| ) | |
| return {"request_id": request_dir.name, "artifact_size": artifact.stat().st_size} | |
| def export_ply(request_id: str, z_offset: float = 0.0) -> dict[str, str | int]: | |
| """Export the reconstructed scene as a filtered Gaussian PLY.""" | |
| if not math.isfinite(z_offset) or abs(z_offset) > 100: | |
| raise ValueError("The export Z offset must be between -100 and 100.") | |
| request_dir = _request_dir(request_id) | |
| artifact = request_dir / "gaussians.pt" | |
| if not artifact.is_file(): | |
| raise FileNotFoundError("The reconstruction artifact is unavailable.") | |
| started_at = time.perf_counter() | |
| scene_ply = export_filtered_gaussian_ply( | |
| artifact_path=artifact, | |
| output_dir=request_dir, | |
| z_offset=z_offset, | |
| ) | |
| artifact.unlink() | |
| _log_timing("ply_export", request_dir, started_at, ply_bytes=scene_ply.stat().st_size) | |
| return {"url": _file_url(scene_ply), "size": scene_ply.stat().st_size} | |
| def export_viewer(request_id: str) -> dict[str, str | int]: | |
| """Create an optimized browser viewer for a reconstructed scene.""" | |
| request_dir = _request_dir(request_id) | |
| scene_ply = request_dir / "scene.ply" | |
| if not scene_ply.is_file(): | |
| raise FileNotFoundError("The PLY export is unavailable.") | |
| started_at = time.perf_counter() | |
| exported = export_browser_viewer(scene_ply=scene_ply, viewer_template=viewer_template) | |
| _log_timing( | |
| "browser_viewer", | |
| request_dir, | |
| started_at, | |
| sog_bytes=exported.scene_sog.stat().st_size, | |
| viewer_html_bytes=exported.viewer_html.stat().st_size, | |
| ) | |
| return {"url": _file_url(exported.viewer_html), "size": exported.viewer_html.stat().st_size} | |
| def export_html(request_id: str) -> dict[str, str | int]: | |
| """Bundle the browser viewer into one standalone HTML file.""" | |
| request_dir = _request_dir(request_id) | |
| viewer_html = request_dir / "viewer.html" | |
| if not viewer_html.is_file(): | |
| raise FileNotFoundError("The browser viewer is unavailable.") | |
| started_at = time.perf_counter() | |
| standalone = export_standalone_viewer( | |
| viewer_html=viewer_html, | |
| viewer_template=viewer_template, | |
| ) | |
| _log_timing("standalone_html", request_dir, started_at, html_bytes=standalone.stat().st_size) | |
| return {"url": _file_url(standalone), "size": standalone.stat().st_size} | |
| def get_viewer_template() -> dict[str, str]: | |
| """Return the shared viewer shell used to warm browser assets.""" | |
| return {"url": _file_url(viewer_template.viewer_html)} | |
| async def homepage() -> FileResponse: | |
| return FileResponse(FRONTEND_ROOT / "index.html", media_type="text/html") | |
| async def frontend_asset(filename: str) -> FileResponse: | |
| assets = { | |
| "styles.css": (FRONTEND_ROOT / "styles.css", "text/css"), | |
| "app.js": (FRONTEND_ROOT / "app.js", "text/javascript"), | |
| } | |
| if filename not in assets: | |
| raise HTTPException(status_code=404, detail="Asset not found") | |
| path, media_type = assets[filename] | |
| return FileResponse(path, media_type=media_type) | |
| async def list_examples() -> list[dict[str, str | int]]: | |
| return examples | |
| def _example(index: int) -> dict[str, str | int]: | |
| if index < 0 or index >= len(examples): | |
| raise HTTPException(status_code=404, detail="Example not found") | |
| return examples[index] | |
| async def example_thumbnail(index: int) -> FileResponse: | |
| example = _example(index) | |
| filename = f"{index:02d}-{Path(str(example['filename'])).stem}.webp" | |
| return FileResponse(EXAMPLE_THUMBNAIL_ROOT / filename, media_type="image/webp") | |
| async def example_source(index: int) -> FileResponse: | |
| example = _example(index) | |
| return FileResponse(EXAMPLE_ROOT / str(example["filename"])) | |
| demo = app | |
| if __name__ == "__main__": | |
| app.launch( | |
| server_name="0.0.0.0", | |
| server_port=int(os.environ.get("PORT", "7860")), | |
| allowed_paths=[str(OUTPUT_ROOT)], | |
| max_file_size="20mb", | |
| show_error=True, | |
| ssr_mode=False, | |
| footer_links=[], | |
| ) | |