File size: 4,340 Bytes
a28da82
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1f619c4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6b2d8f0
a28da82
1f619c4
 
 
 
 
 
a28da82
 
 
 
 
 
 
 
 
 
 
1f619c4
 
 
 
a28da82
 
1f619c4
 
 
 
 
a28da82
 
 
c397958
 
 
 
 
 
 
 
 
 
 
a28da82
 
 
1f619c4
c397958
a28da82
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
import streamlit as st
import torch
import torchvision.transforms as T
from torchvision.models.detection import fasterrcnn_resnet50_fpn
from PIL import Image, ImageDraw, ImageFont

# Load a pre-trained Faster R-CNN model
@st.cache_resource
def load_model():
    model = fasterrcnn_resnet50_fpn(pretrained=True)
    model.eval()
    return model

model = load_model()

# Full COCO class list for reference
COCO_INSTANCE_CATEGORY_NAMES = [
    '__background__', 'person', 'bicycle', 'car', 'motorcycle', 'airplane', 'bus',
    'train', 'truck', 'boat', 'traffic light', 'fire hydrant', 'stop sign',
    'parking meter', 'bench', 'bird', 'cat', 'dog', 'horse', 'sheep', 'cow',
    'elephant', 'bear', 'zebra', 'giraffe', 'backpack', 'umbrella', 'handbag',
    'tie', 'suitcase', 'frisbee', 'skis', 'snowboard', 'sports ball', 'kite',
    'baseball bat', 'baseball glove', 'skateboard', 'surfboard', 'tennis racket',
    'bottle', 'wine glass', 'cup', 'fork', 'knife', 'spoon', 'bowl', 'banana',
    'apple', 'sandwich', 'orange', 'broccoli', 'carrot', 'hot dog', 'pizza',
    'donut', 'cake', 'chair', 'couch', 'potted plant', 'bed', 'dining table',
    'toilet', 'tv', 'laptop', 'mouse', 'remote', 'keyboard', 'cell phone',
    'microwave', 'oven', 'toaster', 'sink', 'refrigerator', 'book', 'clock',
    'vase', 'scissors', 'teddy bear', 'hair drier', 'toothbrush'
]

# Office-related object categories
OFFICE_INSTANCE_CATEGORY_NAMES = [
    '__background__',
    'person', 'chair', 'couch', 'potted plant', 'dining table', 'tv', 'laptop',
    'mouse', 'remote', 'keyboard', 'cell phone', 'book', 'clock', 'vase',
    'scissors', 'teddy bear', 'hair drier', 'toothbrush', 'bottle', 'cup',
    'fork', 'knife', 'spoon', 'bowl', 'backpack', 'handbag', 'tie', 'suitcase',
    'microwave', 'oven', 'toaster', 'sink', 'refrigerator', 'bench', 'umbrella'
]

# Define a transformation to convert the image to tensor
transform = T.Compose([
    T.ToTensor(),
])

# Function to draw bounding boxes on the image
def draw_boxes(image, predictions, threshold=0.5):
    draw = ImageDraw.Draw(image)
    font = ImageFont.load_default()

    # Map valid indices to office-related categories
    office_indices = {i: COCO_INSTANCE_CATEGORY_NAMES[i] for i, name in enumerate(COCO_INSTANCE_CATEGORY_NAMES) if name in OFFICE_INSTANCE_CATEGORY_NAMES}

    for idx, box in enumerate(predictions[0]['boxes']):
        score = predictions[0]['scores'][idx].item()
        label_index = predictions[0]['labels'][idx].item()

        # Only process office-related labels
        if label_index in office_indices and score >= threshold:
            label = office_indices[label_index]
            box = box.tolist()
            draw.rectangle(box, outline="red", width=3)
            text = f"{label}: {int(score * 100)}%"
            text_size = font.getbbox(text)  # Calculate text size
            text_width = text_size[2] - text_size[0]
            text_height = text_size[3] - text_size[1]
            text_position = (box[0], box[1] - text_height)

            # Draw a filled rectangle for text background
            draw.rectangle(
                [box[0], box[1] - text_height, box[0] + text_width, box[1]],
                fill="red"
            )
            draw.text(text_position, text, fill="white", font=font)

# Streamlit interface
def main():
    st.title("Office Object Detection with Faster R-CNN")
    st.write("Upload an image and detect office-related objects using Faster R-CNN.")

    # File uploader
    uploaded_image = st.file_uploader("Upload an image...", type=["jpg", "jpeg", "png"])

    if uploaded_image is not None:
        image = Image.open(uploaded_image).convert("RGB")
        st.image(image, caption="Uploaded Image", use_container_width=True)

        if st.button("Analyze"):
            with st.spinner("Analyzing..."):
                image_tensor = transform(image).unsqueeze(0)  # Add batch dimension

                # Perform object detection
                with torch.no_grad():
                    predictions = model(image_tensor)

                # Draw bounding boxes
                image_with_boxes = image.copy()
                draw_boxes(image_with_boxes, predictions)

                st.image(image_with_boxes, caption="Detected Objects", use_container_width=True)

if __name__ == "__main__":
    main()