""" TripoSR + Dust3R - CPU Tab 1: TripoSR (Single Image -> 3D) - AleenDG's exact code Tab 2: Dust3R (Multi-view -> Point Cloud) """ import os import sys import subprocess import logging import tempfile import time from pathlib import Path import gradio as gr import numpy as np import rembg import torch from PIL import Image from functools import partial from tsr.system import TSR from tsr.utils import remove_background, resize_foreground, to_gradio_3d_orientation # ============================================================ # DEVICE SETUP # ============================================================ if torch.cuda.is_available(): device = "cuda:0" else: device = "cpu" d = os.environ.get("DEVICE", None) if d != None: device = d print(f"Using device: {device}", flush=True) # ============================================================ # TRIPOSR MODEL (AleenDG's exact code) # ============================================================ model = TSR.from_pretrained( "stabilityai/TripoSR", config_name="config.yaml", weight_name="model.ckpt", ) model.renderer.set_chunk_size(131072) model.to(device) rembg_session = rembg.new_session() def check_input_image(input_image): if input_image is None: raise gr.Error("No image uploaded!") def preprocess(input_image, do_remove_background, foreground_ratio): def fill_background(image): image = np.array(image).astype(np.float32) / 255.0 image = image[:, :, :3] * image[:, :, 3:4] + (1 - image[:, :, 3:4]) * 0.5 image = Image.fromarray((image * 255.0).astype(np.uint8)) return image if do_remove_background: image = input_image.convert("RGB") image = remove_background(image, rembg_session) image = resize_foreground(image, foreground_ratio) image = fill_background(image) else: image = input_image if image.mode == "RGBA": image = fill_background(image) return image def generate(image, mc_resolution=256): scene_codes = model(image, device=device) mesh = model.extract_mesh(scene_codes, resolution=mc_resolution)[0] mesh = to_gradio_3d_orientation(mesh) mesh_path = tempfile.NamedTemporaryFile(suffix=".glb", delete=False) mesh.export(mesh_path.name) return mesh_path.name # ============================================================ # DUST3R (Multi-view to Point Cloud) # ============================================================ def setup_dust3r(): """Setup Dust3R repository.""" dust3r_path = Path("/tmp/dust3r") if not dust3r_path.exists(): print("Cloning Dust3R...", flush=True) subprocess.run( ["git", "clone", "--recursive", "--depth", "1", "https://github.com/naver/dust3r", str(dust3r_path)], capture_output=True ) if str(dust3r_path) not in sys.path: sys.path.insert(0, str(dust3r_path)) return dust3r_path class Timer: def __init__(self): self.logs = [] self.start_time = None def start(self): self.start_time = time.time() self.logs = [] self.log("Started") def log(self, msg): elapsed = time.time() - self.start_time if self.start_time else 0 entry = f"[{elapsed:6.2f}s] {msg}" self.logs.append(entry) print(entry, flush=True) def get_logs(self): return "\n".join(self.logs) timer = Timer() def dust3r_reconstruct(images, image_size=512, min_conf_thr=3.0): """Reconstruct 3D point cloud from multiple images using Dust3R.""" timer.start() if not images or len(images) < 2: return None, "Please upload at least 2 images" if len(images) > 5: images = images[:5] timer.log("Limited to 5 images") timer.log(f"Running on {device}") try: import gc setup_dust3r() from dust3r.model import AsymmetricCroCo3DStereo from dust3r.inference import inference from dust3r.image_pairs import make_pairs from dust3r.utils.image import load_images from dust3r.cloud_opt import global_aligner, GlobalAlignerMode from dust3r.demo import get_3D_model_from_scene temp_paths = [] for i, img in enumerate(images): path = tempfile.mktemp(suffix=".png") if hasattr(img, 'name'): import shutil shutil.copy(img.name, path) elif isinstance(img, str): import shutil shutil.copy(img, path) else: Image.fromarray(np.array(img)).save(path) temp_paths.append(path) timer.log("Loading Dust3R model...") dust3r_model = AsymmetricCroCo3DStereo.from_pretrained("naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt") dust3r_model = dust3r_model.to(device) dust3r_model.eval() timer.log("Model loaded") imgs = load_images(temp_paths, size=image_size) pairs = make_pairs(imgs, scene_graph="complete", prefilter=None, symmetrize=True) timer.log(f"Processing {len(pairs)} pairs...") timer.log("Running inference...") with torch.no_grad(): output = inference(pairs, dust3r_model, device, batch_size=1) timer.log("Inference complete") mode = GlobalAlignerMode.PointCloudOptimizer if len(imgs) > 2 else GlobalAlignerMode.PairViewer scene = global_aligner(output, device=device, mode=mode) if mode == GlobalAlignerMode.PointCloudOptimizer: scene.compute_global_alignment(init="mst", niter=300, schedule="linear", lr=0.01) timer.log("Alignment complete") # Export point cloud outdir = tempfile.mkdtemp() timer.log("Exporting point cloud...") output_path = get_3D_model_from_scene( outdir, silent=True, scene=scene, min_conf_thr=min_conf_thr, as_pointcloud=True, mask_sky=False, clean_depth=True, transparent_cams=True, cam_size=0.05 ) timer.log(f"Point cloud saved: {output_path}") del dust3r_model gc.collect() timer.log("COMPLETE") return output_path, timer.get_logs() except Exception as e: import traceback timer.log(f"ERROR: {e}") timer.log(traceback.format_exc()) return None, timer.get_logs() # ============================================================ # GRADIO UI (AleenDG's exact layout for TripoSR) # ============================================================ with gr.Blocks() as demo: gr.Markdown("# TripoSR + Dust3R (CPU)") with gr.Tabs(): # TAB 1: TripoSR (AleenDG's exact layout) with gr.TabItem("TripoSR (Image to 3D)"): with gr.Row(variant="panel"): with gr.Column(): with gr.Row(): input_image = gr.Image( label="Input Image", image_mode="RGBA", sources="upload", type="pil", elem_id="content_image", ) processed_image = gr.Image(label="Processed Image", interactive=False) with gr.Row(): with gr.Group(): do_remove_background = gr.Checkbox( label="Remove Background", value=True ) foreground_ratio = gr.Slider( label="Foreground Ratio", minimum=0.5, maximum=1.0, value=0.85, step=0.05, ) mc_resolution = gr.Slider( label="Mesh Resolution", minimum=64, maximum=320, value=256, step=32, ) with gr.Row(): submit = gr.Button("Generate", elem_id="generate", variant="primary") with gr.Column(): output_model = gr.Model3D( label="Output Model (GLB)", interactive=False, ) submit.click(fn=check_input_image, inputs=[input_image]).success( fn=preprocess, inputs=[input_image, do_remove_background, foreground_ratio], outputs=[processed_image], ).success( fn=generate, inputs=[processed_image, mc_resolution], outputs=[output_model], ) def run_example(image_pil): preprocessed = preprocess(image_pil, True, 0.85) mesh_path = generate(preprocessed, 256) return preprocessed, mesh_path gr.Examples( examples=["examples/front-view-robot.png"], inputs=[input_image], outputs=[processed_image, output_model], fn=run_example, cache_examples=True, cache_mode="lazy", label="Example" ) # TAB 2: Dust3R with gr.TabItem("Dust3R (Multi-view to Point Cloud)"): gr.Markdown("Upload 2-5 images from different viewpoints") with gr.Row(variant="panel"): with gr.Column(): dust3r_images = gr.File( label="Upload 2-5 Images", file_count="multiple", file_types=["image"] ) with gr.Accordion("Settings", open=False): dust3r_size = gr.Radio([224, 512], value=512, label="Image Size") dust3r_conf = gr.Slider(1.0, 20.0, value=3.0, step=0.1, label="Min Confidence") dust3r_btn = gr.Button("Reconstruct", variant="primary") with gr.Column(): dust3r_output = gr.Model3D(label="Point Cloud") dust3r_logs = gr.Textbox(label="Logs", lines=10, interactive=False) dust3r_btn.click( fn=dust3r_reconstruct, inputs=[dust3r_images, dust3r_size, dust3r_conf], outputs=[dust3r_output, dust3r_logs], ) gr.Examples( examples=[ [["examples/front-view-robot.png", "examples/front-left-robot.png", "examples/front-right-robot.png"]], ], inputs=[dust3r_images], label="Example: Robot (3 views)" ) gr.Markdown("---\n[TripoSR](https://huggingface.co/stabilityai/TripoSR) | [Dust3R](https://github.com/naver/dust3r)") demo.queue(max_size=10) demo.launch(server_name="0.0.0.0", server_port=7860)