CoolFace
Apppublic

veltre/GroundingSAM

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
3likes
app.py81 linesDownload Raw Back to root
1from transformers import AutoProcessor, AutoModelForZeroShotObjectDetection2import torch3from transformers import SamModel, SamProcessor4import spaces5import numpy as np6device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')7 8sam_model = SamModel.from_pretrained("facebook/sam-vit-base").to("cuda")9sam_processor = SamProcessor.from_pretrained("facebook/sam-vit-base")10 11model_id = "IDEA-Research/grounding-dino-base"12 13dino_processor = AutoProcessor.from_pretrained(model_id)14dino_model = AutoModelForZeroShotObjectDetection.from_pretrained(model_id).to(device)15 16def infer_dino(img, text_queries, score_threshold):17  queries=""18  for query in text_queries:19    queries += f"{query}. "20 21  width, height = img.shape[:2]22 23  target_sizes=[(width, height)]24  inputs = dino_processor(text=queries, images=img, return_tensors="pt").to(device)25 26  with torch.no_grad():27    outputs = dino_model(**inputs)28    outputs.logits = outputs.logits.cpu()29    outputs.pred_boxes = outputs.pred_boxes.cpu()30    results = dino_processor.post_process_grounded_object_detection(outputs=outputs, input_ids=inputs.input_ids,31                                                                  box_threshold=score_threshold,32                                                                  target_sizes=target_sizes)33  return results34 35 36@spaces.GPU37def query_image(img, text_queries, dino_threshold):38  text_queries = text_queries39  text_queries = text_queries.split(",")40  dino_output = infer_dino(img, text_queries, dino_threshold)41  result_labels=[]42  for pred in dino_output:43    boxes = pred["boxes"].cpu()44    scores = pred["scores"].cpu()45    labels = pred["labels"]46    box = [torch.round(pred["boxes"][0], decimals=2), torch.round(pred["boxes"][1], decimals=2), 47        torch.round(pred["boxes"][2], decimals=2), torch.round(pred["boxes"][3], decimals=2)]48    for box, score, label in zip(boxes, scores, labels):49      if label != "":50        inputs = sam_processor(51                img,52                input_boxes=[[[box]]],53                return_tensors="pt"54            ).to("cuda")55 56        with torch.no_grad():57            outputs = sam_model(**inputs)58 59        mask = sam_processor.image_processor.post_process_masks(60            outputs.pred_masks.cpu(),61            inputs["original_sizes"].cpu(),62            inputs["reshaped_input_sizes"].cpu()63        )[0][0][0].numpy()64        mask = mask[np.newaxis, ...]65        result_labels.append((mask, label))66  return img, result_labels67 68import gradio as gr69 70description = "This Space combines [GroundingDINO](https://huggingface.co/IDEA-Research/grounding-dino-base), a bleeding-edge zero-shot object detection model with [SAM](https://huggingface.co/facebook/sam-vit-base), the state-of-the-art mask generation model. SAM normally doesn't accept text input. Combining SAM with OWLv2 makes SAM text promptable. Try the example or input an image and comma separated candidate labels to segment."71demo = gr.Interface(72    query_image,73    inputs=[gr.Image(label="Image Input"), gr.Textbox(label = "Candidate Labels"), gr.Slider(0, 1, value=0.05, label="Confidence Threshold for GroundingDINO")],74    outputs="annotatedimage",75    title="GroundingDINO 🤝 SAM for Zero-shot Segmentation",76    description=description,77    examples=[78        ["./cats.png", "cat, fishnet", 0.16],["./bee.jpg", "bee, flower", 0.16]79    ],80)81demo.launch(debug=True)