Team Ai
Apppublic

ThirdEyeData/Object_Detection

sourceHugging Faceupdated 4y agoView on Hugging Face
2likes
app.py135 linesDownload Raw Back to root
1from detecto import core, utils, visualize2from detecto.visualize import show_labeled_image, plot_prediction_grid3from torchvision import transforms4import matplotlib.pyplot as plt5from tensorflow.keras.utils import img_to_array6import numpy as np7import warnings8from PIL import Image9import streamlit as st10warnings.filterwarnings("ignore", category=UserWarning) 11from tempfile import NamedTemporaryFile12 13import cv214import matplotlib.patches as patches15 16import torch17 18import matplotlib.image as mpimg19import os20 21from detecto.utils import reverse_normalize, normalize_transform, _is_iterable22from torchvision import transforms23 24 25MODEL_PATH = "SD_model_weights.pth"26IMAGE_PATH = "img1.jpeg"27model = core.Model.load(MODEL_PATH, ['cross_arm','pole','tag'])28#warnings.warn(msg)29 30st.title("Object Detection")31image = utils.read_image(IMAGE_PATH) 32predictions = model.predict(image)33labels, boxes, scores = predictions34 35images = ["img1.jpeg","img4.jpeg","img5.jpeg","img6.jpeg"]36with st.sidebar:37    st.write("choose an image")38    st.image(images)39 40 41 42def detect_object(IMAGE_PATH):43    image = utils.read_image(IMAGE_PATH) 44  #  predictions = model.predict(image)45   # labels, boxes, scores = predictions46 47 48    thresh=0.249    filtered_indices=np.where(scores>thresh)50    filtered_scores=scores[filtered_indices]51    filtered_boxes=boxes[filtered_indices]52    num_list = filtered_indices[0].tolist()53    filtered_labels = [labels[i] for i in num_list]54    show_labeled_image(image, filtered_boxes, filtered_labels)55    56    fig1 = show_image(image,filtered_boxes,filtered_labels)57    st.write("Object Detected Image is")58    st.image(fig1)59    #img_array = img_to_array(img)60def show_image(image, boxes, labels=None):61    """Show the image along with the specified boxes around detected objects.62    Also displays each box's label if a list of labels is provided.63    :param image: The image to plot. If the image is a normalized64        torch.Tensor object, it will automatically be reverse-normalized65        and converted to a PIL image for plotting.66    :type image: numpy.ndarray or torch.Tensor67    :param boxes: A torch tensor of size (N, 4) where N is the number68        of boxes to plot, or simply size 4 if N is 1.69    :type boxes: torch.Tensor70    :param labels: (Optional) A list of size N giving the labels of71            each box (labels[i] corresponds to boxes[i]). Defaults to None.72    :type labels: torch.Tensor or None73    **Example**::74        >>> from detecto.core import Model75        >>> from detecto.utils import read_image76        >>> from detecto.visualize import show_labeled_image77        >>> model = Model.load('model_weights.pth', ['tick', 'gate'])78        >>> image = read_image('image.jpg')79        >>> labels, boxes, scores = model.predict(image)80        >>> show_labeled_image(image, boxes, labels)81    """82    fig, ax = plt.subplots(1)83    # If the image is already a tensor, convert it back to a PILImage84    # and reverse normalize it85    if isinstance(image, torch.Tensor):86        image = reverse_normalize(image)87        image = transforms.ToPILImage()(image)88    ax.imshow(image)89    90    # Show a single box or multiple if provided91    if boxes.ndim == 1:92        boxes = boxes.view(1, 4)93 94    if labels is not None and not _is_iterable(labels):95        labels = [labels]96 97    # Plot each box98    for i in range(2):99        box = boxes[i]100        width, height = (box[2] - box[0]).item(), (box[3] - box[1]).item()101        initial_pos = (box[0].item(), box[1].item())102        rect = patches.Rectangle(initial_pos,  width, height, linewidth=1,103                                 edgecolor='r', facecolor='none')104        if labels:105            ax.text(box[0] + 5, box[1] - 5, '{}'.format(labels[i]), color='red')106 107        ax.add_patch(rect)108 109    cp = os.path.abspath(os.getcwd()) + '/foo.png'110    plt.savefig(cp)111    plt.close(fig)112    return cp113    #print(type(plt114 115file = st.file_uploader('Upload an Image',type=(["jpeg","jpg","png"]))116 117if file is None:118    st.write("Please upload an image file")119else:120    image= Image.open(file)121    st.write("Input Image")122    st.image(image,use_column_width = True)123    with NamedTemporaryFile(dir='.', suffix='.jpeg') as f:124        f.write(file.getbuffer())125    #your_function_which_takes_a_path(f.name)126   127        detect_object(f.name)128  129    130    131 132 133 134 135