Team Ai
Apppublic

spark-nlp/VisionEncoderDecoderForImageCaptioning

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
1likes
Demo.py123 linesDownload Raw Back to root
1import streamlit as st2import sparknlp3import os4import pandas as pd5 6from sparknlp.base import *7from sparknlp.annotator import *8from pyspark.ml import Pipeline9from sparknlp.pretrained import PretrainedPipeline10 11# Page configuration12st.set_page_config(13    layout="wide", 14    initial_sidebar_state="auto"15)16 17# CSS for styling18st.markdown("""19    <style>20        .main-title {21            font-size: 36px;22            color: #4A90E2;23            font-weight: bold;24            text-align: center;25        }26        .section {27            background-color: #f9f9f9;28            padding: 10px;29            border-radius: 10px;30            margin-top: 10px;31        }32        .section p, .section ul {33            color: #666666;34        }35    </style>36""", unsafe_allow_html=True)37 38@st.cache_resource39def init_spark():40    return sparknlp.start()41 42@st.cache_resource43def create_pipeline(model):44    imageAssembler = ImageAssembler() \45        .setInputCol("image") \46        .setOutputCol("image_assembler")47 48    imageCaptioning = VisionEncoderDecoderForImageCaptioning \49        .pretrained("image_captioning_vit_gpt2") \50        .setBeamSize(2) \51        .setDoSample(False) \52        .setInputCols(["image_assembler"]) \53        .setOutputCol("caption")54 55    pipeline = Pipeline(stages=[imageAssembler, imageCaptioning])56    return pipeline57 58def fit_data(pipeline, data):59    empty_df = spark.createDataFrame([['']]).toDF('text')60    model = pipeline.fit(empty_df)61    light_pipeline = LightPipeline(model)62    annotations_result = light_pipeline.fullAnnotateImage(data)63    return annotations_result[0]['caption'][0].result64 65def save_uploadedfile(uploadedfile):66    filepath = os.path.join(IMAGE_FILE_PATH, uploadedfile.name)67    with open(filepath, "wb") as f:68        if hasattr(uploadedfile, 'getbuffer'):69            f.write(uploadedfile.getbuffer())70        else:71            f.write(uploadedfile.read())72        73# Sidebar content74model_list = ['image_captioning_vit_gpt2']75model = st.sidebar.selectbox(76    "Choose the pretrained model",77    model_list,78    help="For more info about the models visit: https://sparknlp.org/models"79)80 81# Set up the page layout82st.markdown(f'<div class="main-title">VisionEncoderDecoder For Image Captioning</div>', unsafe_allow_html=True)83# st.markdown(f'<div class="section"><p>{sub_title}</p></div>', unsafe_allow_html=True)84 85# Reference notebook link in sidebar86link = """87<a href="https://colab.research.google.com/github/JohnSnowLabs/spark-nlp/blob/master/examples/python/annotation/image/VisionEncoderDecoderForImageCaptioning.ipynb">88    <img src="https://colab.research.google.com/assets/colab-badge.svg" style="zoom: 1.3" alt="Open In Colab"/>89</a>90"""91st.sidebar.markdown('Reference notebook:')92st.sidebar.markdown(link, unsafe_allow_html=True)93 94# Load examples95IMAGE_FILE_PATH = f"inputs"96image_files = sorted([file for file in os.listdir(IMAGE_FILE_PATH) if file.split('.')[-1]=='png' or file.split('.')[-1]=='jpg' or file.split('.')[-1]=='JPEG' or file.split('.')[-1]=='jpeg'])97 98img_options = st.selectbox("Select an image", image_files)99uploadedfile = st.file_uploader("Try it for yourself!")100 101if uploadedfile:102    file_details = {"FileName":uploadedfile.name,"FileType":uploadedfile.type}103    save_uploadedfile(uploadedfile)104    selected_image = f"{IMAGE_FILE_PATH}/{uploadedfile.name}"105elif img_options:106    selected_image = f"{IMAGE_FILE_PATH}/{img_options}"107 108st.subheader('Classified Image')109 110image_size = st.slider('Image Size', 400, 1000, value=400, step = 100)111 112try:113    st.image(f"{IMAGE_FILE_PATH}/{selected_image}", width=image_size)114except:115    st.image(selected_image, width=image_size)116 117st.subheader('Classification')118 119spark = init_spark()120Pipeline = create_pipeline(model)121output = fit_data(Pipeline, selected_image)122 123st.markdown(f'This document has been classified as  : **{output}**')