CoolFace
Apppublic

CodeIcarus/image-captioning-segmentation

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
app.py65 linesDownload Raw Back to root
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