intelli-zen/detr_cppe5_object_detection
0
1#!/usr/bin/python32# -*- coding: utf-8 -*-3import argparse4import io5import json6import os7import re8from typing import Dict, List9 10from project_settings import project_path11 12os.environ["HUGGINGFACE_HUB_CACHE"] = (project_path / "cache/huggingface/hub").as_posix()13 14import gradio as gr15import matplotlib.pyplot as plt16import numpy as np17from PIL import Image18import requests19import torch20from transformers.models.auto.processing_auto import AutoImageProcessor21from transformers.models.auto.feature_extraction_auto import AutoFeatureExtractor22from transformers.models.auto.modeling_auto import AutoModelForObjectDetection23import validators24 25from project_settings import project_path26 27 28# colors for visualization29COLORS = [30 [0.000, 0.447, 0.741],31 [0.850, 0.325, 0.098],32 [0.929, 0.694, 0.125],33 [0.494, 0.184, 0.556],34 [0.466, 0.674, 0.188],35 [0.301, 0.745, 0.933]36]37 38 39def get_original_image(url_input):40 if validators.url(url_input):41 image = Image.open(requests.get(url_input, stream=True).raw)42 return image43 44 45def figure2image(fig):46 buf = io.BytesIO()47 fig.savefig(buf)48 buf.seek(0)49 pil_image = Image.open(buf)50 base_width = 75051 width_percent = base_width / float(pil_image.size[0])52 height_size = (float(pil_image.size[1]) * float(width_percent))53 height_size = int(height_size)54 pil_image = pil_image.resize((base_width, height_size), Image.Resampling.LANCZOS)55 return pil_image56 57 58def non_max_suppression(boxes, scores, threshold):59 """Apply non-maximum suppression at test time to avoid detecting too many60 overlapping bounding boxes for a given object.61 Args:62 boxes: array of [xmin, ymin, xmax, ymax]63 scores: array of scores associated with each box.64 threshold: IoU threshold65 Return:66 keep: indices of the boxes to keep67 """68 x1 = boxes[:, 0]69 y1 = boxes[:, 1]70 x2 = boxes[:, 2]71 y2 = boxes[:, 3]72 73 areas = (x2 - x1 + 1) * (y2 - y1 + 1)74 order = scores.argsort()[::-1] # get boxes with more confidence first75 76 keep = []77 while order.size > 0:78 i = order[0] # pick max confidence box79 keep.append(i)80 81 xx1 = np.maximum(x1[i], x1[order[1:]])82 yy1 = np.maximum(y1[i], y1[order[1:]])83 xx2 = np.minimum(x2[i], x2[order[1:]])84 yy2 = np.minimum(y2[i], y2[order[1:]])85 86 w = np.maximum(0.0, xx2 - xx1 + 1) # maximum width87 h = np.maximum(0.0, yy2 - yy1 + 1) # maximum height88 inter = w * h89 90 ovr = inter / (areas[i] + areas[order[1:]] - inter)91 inds = np.where(ovr <= threshold)[0]92 order = order[inds + 1]93 94 return keep95 96 97def draw_boxes(image, boxes, scores, labels, threshold: float,98 idx_to_label: Dict[int, str] = None, labels_to_show: str = None):99 if isinstance(labels_to_show, str):100 if len(labels_to_show.strip()) == 0:101 labels_to_show = None102 else:103 labels_to_show = labels_to_show.split(",")104 labels_to_show = [label.strip().lower() for label in labels_to_show]105 labels_to_show = None if len(labels_to_show) == 0 else labels_to_show106 107 plt.figure(figsize=(50, 50))108 plt.imshow(image)109 110 if idx_to_label is not None:111 labels = [idx_to_label[x] for x in labels]112 113 axis = plt.gca()114 colors = COLORS * len(boxes)115 for score, (xmin, ymin, xmax, ymax), label, color in zip(scores, boxes, labels, colors):116 if labels_to_show is not None and label.lower() not in labels_to_show:117 continue118 if score < threshold:119 continue120 axis.add_patch(plt.Rectangle((xmin, ymin), xmax - xmin, ymax - ymin, fill=False, color=color, linewidth=10))121 axis.text(xmin, ymin, f"{label}: {score:0.2f}", fontsize=60, bbox=dict(facecolor="yellow", alpha=0.8))122 plt.axis("off")123 124 return figure2image(plt.gcf())125 126 127def detr_object_detection(url_input: str,128 image_input: Image,129 pretrained_model_name_or_path: str = "qgyd2021/detr_cppe5_object_detection",130 threshold: float = 0.5,131 iou_threshold: float = 0.5,132 labels_to_show: str = None,133 ):134 # feature_extractor = AutoFeatureExtractor.from_pretrained(pretrained_model_name_or_path)135 model = AutoModelForObjectDetection.from_pretrained(pretrained_model_name_or_path)136 image_processor = AutoImageProcessor.from_pretrained(pretrained_model_name_or_path)137 138 # image139 if validators.url(url_input):140 image = get_original_image(url_input)141 elif image_input:142 image = image_input143 else:144 raise AssertionError("at least one `url_input` and `image_input`")145 image_size = torch.tensor([tuple(reversed(image.size))])146 147 # infer148 # inputs = feature_extractor(images=image, return_tensors="pt")149 inputs = image_processor(images=image, return_tensors="pt")150 outputs = model.forward(**inputs)151 152 processed_outputs = image_processor.post_process_object_detection(153 outputs, threshold=threshold, target_sizes=image_size)154 # processed_outputs = feature_extractor.post_process(outputs, target_sizes=image_size)155 processed_outputs = processed_outputs[0]156 157 # draw box158 boxes = processed_outputs["boxes"].detach().numpy()159 scores = processed_outputs["scores"].detach().numpy()160 labels = processed_outputs["labels"].detach().numpy()161 162 keep = non_max_suppression(boxes, scores, threshold=iou_threshold)163 boxes = boxes[keep]164 scores = scores[keep]165 labels = labels[keep]166 167 viz_image: Image = draw_boxes(168 image, boxes, scores, labels,169 threshold=threshold,170 idx_to_label=model.config.id2label,171 labels_to_show=labels_to_show172 )173 return viz_image174 175 176def main():177 178 title = "## Detr Cppe5 Object Detection"179 180 description = """181 reference:182 https://huggingface.co/docs/transformers/tasks/object_detection183 184 """185 186 example_urls = [187 *[188 [189 "https://huggingface.co/datasets/intelli-zen/cppe-5/resolve/main/data/images/{}.png".format(idx),190 "intelli-zen/detr_cppe5_object_detection",191 0.5, 0.6, None192 ] for idx in range(1001, 1030)193 ]194 ]195 196 example_images = [197 [198 "data/2lnWoly.jpg",199 "intelli-zen/detr_cppe5_object_detection",200 0.5, 0.6, None201 ]202 ]203 204 with gr.Blocks() as blocks:205 gr.Markdown(value=title)206 gr.Markdown(value=description)207 208 model_name = gr.components.Dropdown(209 choices=[210 "intelli-zen/detr_cppe5_object_detection",211 ],212 value="intelli-zen/detr_cppe5_object_detection",213 label="model_name",214 )215 threshold_slider = gr.components.Slider(216 minimum=0, maximum=1.0,217 step=0.01, value=0.5,218 label="Threshold"219 )220 iou_threshold_slider = gr.components.Slider(221 minimum=0, maximum=1.0,222 step=0.1, value=0.5,223 label="IOU Threshold"224 )225 classes_to_detect = gr.Textbox(placeholder="e.g. person, truck (split by , comma).",226 label="labels to show")227 228 with gr.Tabs():229 with gr.TabItem("Image URL"):230 with gr.Row():231 with gr.Column():232 url_input = gr.Textbox(lines=1, label="Enter valid image URL here..")233 original_image = gr.Image()234 url_input.change(get_original_image, url_input, original_image)235 with gr.Column():236 img_output_from_url = gr.Image()237 238 url_but = gr.Button("Detect")239 240 with gr.Row():241 gr.Examples(examples=example_urls,242 inputs=[url_input, model_name, threshold_slider, iou_threshold_slider],243 examples_per_page=5,244 )245 246 with gr.TabItem("Image Upload"):247 with gr.Row():248 img_input = gr.Image(type="pil")249 img_output_from_upload = gr.Image()250 251 img_but = gr.Button("Detect")252 253 with gr.Row():254 gr.Examples(examples=example_images,255 inputs=[img_input, model_name, threshold_slider, iou_threshold_slider],256 examples_per_page=5,257 )258 259 inputs = [url_input, img_input, model_name, threshold_slider, iou_threshold_slider, classes_to_detect]260 url_but.click(detr_object_detection, inputs=inputs, outputs=[img_output_from_url], queue=True)261 img_but.click(detr_object_detection, inputs=inputs, outputs=[img_output_from_upload], queue=True)262 263 blocks.queue().launch()264 return265 266 267if __name__ == '__main__':268 main()269 