Team Ai
Apppublic

intelli-zen/detr_cppe5_object_detection

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
main.py269 linesDownload Raw Back to root
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