mikekuniavsky's picture
Fix result extraction: use result[glb] trimesh.Trimesh directly
3bc6675
Raw History Blame Contribute Delete
8.16 kB
"""
SAM 3D Objects MCP Server
Image → 3D Object (GLB)
Automatic object detection with SAM2 + 3D reconstruction with SAM 3D Objects.
"""
import os
import sys
import subprocess
import tempfile
import uuid
from pathlib import Path
# ---------------------------------------------------------------------------
# kaolin shim — provide stub symbols so sam-3d-objects imports do not crash.
# No kaolin wheel exists for torch 2.10+cu128; the symbols used are either
# dead imports (Camera, IpyTurntableVisualizer) or trivial shape checkers.
# ---------------------------------------------------------------------------
try:
import kaolin # noqa: F401 — real kaolin installed, nothing to do
except ModuleNotFoundError:
import types, torch as _torch
def _check_tensor(tensor, shape, throw=True):
if not _torch.is_tensor(tensor):
if throw: raise TypeError(f"Expected tensor, got {type(tensor)}")
return False
if len(tensor.shape) != len(shape):
if throw: raise ValueError(f"Expected {len(shape)}D, got {len(tensor.shape)}D")
return False
for a, e in zip(tensor.shape, shape):
if e is not None and a != e:
if throw: raise ValueError(f"Shape {tensor.shape} != {shape}")
return False
return True
_k = types.ModuleType("kaolin"); sys.modules["kaolin"] = _k
_kv = types.ModuleType("kaolin.visualize"); sys.modules["kaolin.visualize"] = _kv
_kv.IpyTurntableVisualizer = type("IpyTurntableVisualizer", (), {})
_kr = types.ModuleType("kaolin.render"); sys.modules["kaolin.render"] = _kr
_krc = types.ModuleType("kaolin.render.camera"); sys.modules["kaolin.render.camera"] = _krc
_krc.Camera = type("Camera", (), {})
_krc.CameraExtrinsics = type("CameraExtrinsics", (), {})
_krc.PinholeIntrinsics = type("PinholeIntrinsics", (), {})
_ku = types.ModuleType("kaolin.utils"); sys.modules["kaolin.utils"] = _ku
_kut = types.ModuleType("kaolin.utils.testing"); sys.modules["kaolin.utils.testing"] = _kut
_kut.check_tensor = _check_tensor
print("kaolin shim installed (no real kaolin wheel for this torch/CUDA)")
import gradio as gr
import numpy as np
import spaces
from huggingface_hub import snapshot_download, login
from PIL import Image
# Login with HF_TOKEN if available
if os.environ.get("HF_TOKEN"):
login(token=os.environ.get("HF_TOKEN"))
# Set CUDA_HOME for sam-3d-objects (expects conda but we're not using it)
os.environ.setdefault("LIDRA_SKIP_INIT", "1")
if "CUDA_HOME" not in os.environ:
os.environ["CUDA_HOME"] = "/usr/local/cuda"
if "CONDA_PREFIX" not in os.environ:
os.environ["CONDA_PREFIX"] = "/usr/local"
# Clone sam-3d-objects repo if not exists
SAM3D_PATH = Path("/home/user/app/sam-3d-objects")
if not SAM3D_PATH.exists():
print("Cloning sam-3d-objects repository...")
subprocess.run([
"git", "clone",
"https://github.com/facebookresearch/sam-3d-objects.git",
str(SAM3D_PATH)
], check=True)
# Add both repo root and notebook folder to path
sys.path.insert(0, str(SAM3D_PATH))
sys.path.insert(0, str(SAM3D_PATH / "notebook"))
# Global models
SAM3D_MODEL = None
SAM2_GENERATOR = None
def load_sam2():
"""Load SAM2 automatic mask generator"""
global SAM2_GENERATOR
if SAM2_GENERATOR is not None:
return SAM2_GENERATOR
from sam2.automatic_mask_generator import SAM2AutomaticMaskGenerator
print("Loading SAM2 model...")
SAM2_GENERATOR = SAM2AutomaticMaskGenerator.from_pretrained("facebook/sam2-hiera-large")
print("✓ SAM2 loaded")
return SAM2_GENERATOR
def load_sam3d():
"""Load SAM 3D Objects model"""
global SAM3D_MODEL
if SAM3D_MODEL is not None:
return SAM3D_MODEL
import torch
print("Loading SAM 3D Objects model...")
# Download checkpoints
checkpoint_dir = snapshot_download(
repo_id="facebook/sam-3d-objects",
token=os.environ.get("HF_TOKEN")
)
# Import from notebook/inference.py
from inference import Inference
# Config path in the repo
config_path = str(Path(checkpoint_dir) / "checkpoints" / "pipeline.yaml")
SAM3D_MODEL = Inference(config_path, compile=False)
# Point to downloaded checkpoints
print("✓ SAM 3D Objects loaded")
return SAM3D_MODEL
@spaces.GPU(duration=600)
def reconstruct_objects(image: np.ndarray):
"""
Automatically detect and reconstruct 3D objects from image.
Args:
image: Input RGB image
Returns:
tuple: (glb_path, preview_image, status)
"""
if image is None:
return None, None, "❌ No image provided"
try:
import torch
import trimesh
from PIL import Image as PILImage
# Load models
generator = load_sam2()
inference = load_sam3d()
# Convert to PIL if needed
if isinstance(image, np.ndarray):
pil_image = PILImage.fromarray(image)
else:
pil_image = image
image = np.array(pil_image)
# Auto-detect all objects with SAM2
print("Detecting objects...")
masks = generator.generate(image)
if not masks or len(masks) == 0:
return None, image, "⚠️ No objects detected"
# Sort by area, take largest object
masks = sorted(masks, key=lambda x: x['area'], reverse=True)
best_mask = masks[0]['segmentation']
# Create preview with mask overlay
preview = image.copy()
preview[best_mask] = (preview[best_mask] * 0.5 + np.array([0, 255, 0]) * 0.5).astype(np.uint8)
# Convert mask to PIL
# Run 3D reconstruction
print("Reconstructing 3D...")
result = inference(image=image, mask=best_mask)
if result is None:
return None, preview, "⚠️ 3D reconstruction failed"
# Export as GLB
output_dir = tempfile.mkdtemp()
glb_path = f"{output_dir}/object_{uuid.uuid4().hex[:8]}.glb"
# result["glb"] is a trimesh.Trimesh already built by postprocess_slat_output
glb_mesh = result.get("glb")
if glb_mesh is not None:
glb_mesh.export(glb_path)
else:
return None, preview, "⚠️ Could not extract 3D data (no glb in result)"
return glb_path, preview, f"✓ Detected {len(masks)} objects, reconstructed largest"
except Exception as e:
import traceback
traceback.print_exc()
return None, None, f"❌ Error: {e}"
# Gradio Interface
with gr.Blocks(title="SAM 3D Objects MCP") as demo:
gr.Markdown("""
# 📦 SAM 3D Objects MCP Server
**Image → 3D Object (GLB)**
Automatically detects objects and reconstructs the largest one in 3D.
""")
with gr.Row():
with gr.Column():
input_image = gr.Image(label="Input Image", type="numpy")
btn = gr.Button("🚀 Detect & Reconstruct", variant="primary", size="lg")
with gr.Column():
preview = gr.Image(label="Detected Object", type="numpy", interactive=False)
status = gr.Textbox(label="Status")
with gr.Row():
with gr.Column():
output_model = gr.Model3D(label="3D Preview")
with gr.Column():
output_file = gr.File(label="Download GLB")
btn.click(
reconstruct_objects,
inputs=[input_image],
outputs=[output_model, preview, status]
)
output_model.change(lambda x: x, inputs=[output_model], outputs=[output_file])
gr.Markdown("""
---
### MCP Server
```json
{
"mcpServers": {
"sam3d-objects": {
"url": "https://mikekuniavsky-sam3d-objects-mcp.hf.space/gradio_api/mcp/sse"
}
}
}
```
""")
# Pre-download SAM 3D checkpoints at startup so first inference isn't slow
print("Pre-downloading SAM 3D Objects checkpoints...")
_checkpoint_dir = snapshot_download(
repo_id="facebook/sam-3d-objects",
token=os.environ.get("HF_TOKEN")
)
print(f"✓ Checkpoints ready at {_checkpoint_dir}")
if __name__ == "__main__":
demo.launch(mcp_server=True)