objects / app.py
reinformator's picture
Update app.py
1f619c4 verified
Raw History Blame Contribute Delete
4.34 kB
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()