room-object-segmentation / app_broken.py
krishnanandiraju's picture
Deploy room segmentation app with caching
d581790 verified
Raw
History Blame Contribute Delete
5.96 kB
"""
Room Object Segmentation with Grounding DINO + SAM2
Zero-shot detection and segmentation for indoor scenes
Version: 2.0 - Simplified for persistent GPU
"""
import gradio as gr
import numpy as np
import torch
from PIL import Image
from transformers import AutoProcessor, AutoModelForZeroShotObjectDetection
from sam2.build_sam import build_sam2_hf
from sam2.sam2_image_predictor import SAM2ImagePredictor
import supervision as sv
import hashlib
# Model IDs
GDINO_ID = "IDEA-Research/grounding-dino-tiny"
SAM2_ID = "facebook/sam2.1-hiera-small"
# Set device
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"🎯 Using device: {device}")
# Load models at startup
print("Loading Grounding DINO...")
gdino_processor = AutoProcessor.from_pretrained(GDINO_ID)
gdino_model = AutoModelForZeroShotObjectDetection.from_pretrained(GDINO_ID).to(device)
print("✓ Grounding DINO loaded")
print("Loading SAM2...")
sam2_model = build_sam2_hf(SAM2_ID, device=device)
sam2_predictor = SAM2ImagePredictor(sam2_model)
print("✓ SAM2 loaded")
print("🚀 All models ready!")
def image_hash(image):
"""Generate hash for caching."""
if isinstance(image, str):
with open(image, 'rb') as f:
return hashlib.sha256(f.read()).hexdigest()[:16]
elif isinstance(image, np.ndarray):
return hashlib.sha256(image.tobytes()).hexdigest()[:16]
return None
def run_pipeline(image_path, text_prompt, box_threshold, text_threshold):
"""Run detection + segmentation pipeline."""
# Load image
image = Image.open(image_path).convert("RGB")
img_np = np.array(image)
# Grounding DINO detection
inputs = gdino_processor(images=image, text=text_prompt, return_tensors="pt").to(device)
with torch.no_grad():
outputs = gdino_model(**inputs)
results = gdino_processor.post_process_grounded_object_detection(
outputs,
inputs.input_ids,
box_threshold=box_threshold,
text_threshold=text_threshold,
target_sizes=[image.size[::-1]]
)[0]
# Convert to numpy
boxes = results["boxes"].cpu().numpy()
labels = results["labels"]
scores = results["scores"].cpu().numpy()
if len(boxes) == 0:
return image, {
"num_detections": 0,
"message": "No objects detected. Try lowering the thresholds."
}
# SAM2 segmentation
sam2_predictor.set_image(img_np)
masks, iou_predictions, _ = sam2_predictor.predict(
box=boxes,
multimask_output=False
)
# Handle SAM2 scores
if hasattr(iou_predictions, 'cpu'):
sam_scores = iou_predictions.cpu().numpy()
else:
sam_scores = np.array(iou_predictions) if isinstance(iou_predictions, list) else iou_predictions
if len(sam_scores.shape) > 1:
sam_scores = sam_scores[:, 0]
# Create detections
detections = sv.Detections(
xyxy=boxes,
mask=masks[:, 0, :, :] if len(masks.shape) == 4 else masks,
class_id=np.arange(len(boxes)),
confidence=scores
)
# Annotate image
box_annotator = sv.BoxAnnotator()
label_annotator = sv.LabelAnnotator()
mask_annotator = sv.MaskAnnotator()
annotated = mask_annotator.annotate(scene=img_np.copy(), detections=detections)
annotated = box_annotator.annotate(scene=annotated, detections=detections)
annotated = label_annotator.annotate(
scene=annotated,
detections=detections,
labels=[f"{labels[i]} {scores[i]:.2f}" for i in range(len(labels))]
)
# Build JSON response
result = {
"num_detections": len(boxes),
"detections": [
{
"id": i,
"label": labels[i],
"confidence": float(scores[i]),
"sam_score": float(sam_scores[i]) if i < len(sam_scores) else None,
"bbox": boxes[i].tolist(),
"area": float(np.sum(masks[i, 0, :, :] if len(masks.shape) == 4 else masks[i]))
}
for i in range(len(boxes))
],
"image_shape": list(img_np.shape)
}
return Image.fromarray(annotated), result
def segment(image, text_prompt, box_threshold, text_threshold):
"""Main entry point."""
if image is None:
return None, {"error": "No image provided"}
try:
print(f"Processing: {text_prompt[:50]}...")
# Run pipeline
annotated_img, result = run_pipeline(image, text_prompt, box_threshold, text_threshold)
return annotated_img, result
except Exception as e:
import traceback
error_msg = f"Error: {str(e)}\n{traceback.format_exc()}"
print(error_msg)
return None, {"error": error_msg}
# Gradio interface
demo = gr.Interface(
fn=segment,
inputs=[
gr.Image(type="filepath", label="Upload room image"),
gr.Textbox(
value="chair . table . sofa . lamp . bed . cabinet . shelf . desk",
label="Objects to detect (separate with ' . ')",
placeholder="chair . table . sofa"
),
gr.Slider(0.0, 1.0, value=0.35, label="Box Threshold", info="Higher = fewer but more confident boxes"),
gr.Slider(0.0, 1.0, value=0.25, label="Text Threshold", info="Higher = stricter text matching"),
],
outputs=[
gr.Image(label="Annotated Image"),
gr.JSON(label="Detection Results"),
],
title="🏠 Room Object Segmentation",
description="Zero-shot object detection and segmentation using Grounding DINO + SAM 2.1",
examples=[
[None, "chair . table . sofa . lamp", 0.35, 0.25],
],
api_name="segment",
)
if __name__ == "__main__":
demo.launch(show_error=True)