CodeIcarus/image-captioning-segmentation
0
1import streamlit as st2from PIL import Image3import torch4import torchvision.transforms as T5from torchvision.models.detection import maskrcnn_resnet50_fpn6from transformers import BlipProcessor, BlipForConditionalGeneration7import numpy as np8 9# ====== Load Captioning Model ======10processor = BlipProcessor.from_pretrained("Salesforce/blip-image-captioning-base")11model = BlipForConditionalGeneration.from_pretrained("Salesforce/blip-image-captioning-base")12 13# ====== Load Pretrained Segmentation Model ======14seg_model = maskrcnn_resnet50_fpn(pretrained=True)15seg_model.eval()16 17# ====== Image Transform ======18transform = T.Compose([T.ToTensor()])19 20# ====== Streamlit UI Config ======21st.set_page_config(page_title="Image Captioning & Segmentation", layout="wide")22st.title("๐ผ๏ธ Image Captioning + ๐ฏ Segmentation")23st.markdown("Upload an image to generate a caption and visualize object segmentation.")24 25uploaded_file = st.file_uploader("๐ค Upload an Image", type=["png", "jpg", "jpeg"])26 27if uploaded_file is not None:28 image = Image.open(uploaded_file).convert("RGB")29 st.image(image, caption="Original Image", width=400)30 31 col1, col2 = st.columns(2)32 33 with st.spinner("โณ Running captioning and segmentation..."):34 35 # ====== Generate Caption ======36 inputs = processor(images=image, return_tensors="pt")37 output = model.generate(**inputs)38 caption = processor.decode(output[0], skip_special_tokens=True)39 40 # ====== Run Segmentation ======41 img_tensor = transform(image).unsqueeze(0)42 with torch.no_grad():43 pred = seg_model(img_tensor)[0]44 45 # ====== Draw Segmentation Masks ======46 def draw_masks(img, prediction, max_masks=5):47 img_np = np.array(img).copy()48 masks = prediction['masks']49 for i in range(min(len(masks), max_masks)):50 mask = masks[i, 0].mul(255).byte().cpu().numpy()51 red_mask = np.zeros_like(img_np)52 red_mask[:, :, 0] = mask # Apply mask to red channel53 img_np = np.where(red_mask > 0, 0.5 * img_np + 0.5 * red_mask, img_np)54 return Image.fromarray(img_np.astype(np.uint8))55 56 segmented_image = draw_masks(image, pred)57 58 # ====== Display Output ======59 col1.subheader("๐ Caption:")60 col1.markdown(f"**`{caption}`**")61 62 col2.subheader("๐ฏ Segmented Image:")63 col2.image(segmented_image, use_container_width=True)64 65 