Team Ai
Apppublic

taskswithcode/salient-object-detection

sourceHugging Facemitupdated 4y agoView on Hugging Face
13likes
app.py234 linesDownload Raw Back to root
1import time2import sys3import streamlit as st4import string5import os6from io import StringIO 7import pdb8import json9import torch10import requests11import socket12from streamlit_image_select import image_select13 14 15 16 17 18use_case = {"1":"Image background removal - (upload any picture and remove background)","2":"Masking foreground for downstream inpainting task"}19mask_types = {20"rgba - makes background white":"rgba",21"green - makes the background green":"green",22"blur - blurs background":"blur",23"map - makes the foreground white and rest black ":"map"24}25 26 27 28APP_NAME = "hf/salient_object_detection"29INFO_URL = "https://www.taskswithcode.com/stats/"30TMP_DIR="tmp_dir"31TMP_SEED = 132 33 34 35        36 37def get_views(action):38    ret_val = 039    #return "{:,}".format(ret_val)40    hostname = socket.gethostname()41    ip_address = socket.gethostbyname(hostname)42    if ("view_count" not in st.session_state):43        try:44           app_info = {'name': APP_NAME,"action":action,"host":hostname,"ip":ip_address}45           res = requests.post(INFO_URL, json = app_info).json()46           print(res)47           data = res["count"]48        except Exception as e:49           data = 050           print(f"Exception in get_views - uncached case. view count not cached: {str(e)}")51        ret_val = data52        st.session_state["view_count"] = data53    else:54        ret_val = st.session_state["view_count"]55        if (action != "init"):56           try:57               app_info = {'name': APP_NAME,"action":action,"host":hostname,"ip":ip_address}58               print(app_info)59               res = requests.post(INFO_URL, json = app_info).json()60           except Exception as e:61                print(f"Exception in get_views - Non init case. view count not cached: {str(e)}")62    return "{:,}".format(ret_val)63        64 65 66 67def construct_model_info_for_display(model_names,api_info):68    options_arr  = []69    #markdown_str = f"<div style=\"font-size:16px; color: #2f2f2f; text-align: left\"><br/><b>Models evaluated ({len(model_names)})</b><br/></div>"70    markdown_str = f"<div style=\"font-size:16px; color: #2f2f2f; text-align: left\"><br/><b>Model evaluated </b><br/></div>"71    markdown_str += f"<div style=\"font-size:2px; color: #2f2f2f; text-align: left\"><br/></div>"72    for node in model_names:73        options_arr .append(node["name"])74        if (node["mark"] == "True"):75            markdown_str += f"<div style=\"font-size:16px; color: #5f5f5f; text-align: left\">&nbsp;•&nbsp;Model:&nbsp;<a href=\'{node['paper_url']}\' target='_blank'>{node['name']}</a><br/>&nbsp;&nbsp;&nbsp;&nbsp;Code released by:&nbsp;<a href=\'{node['orig_author_url']}\' target='_blank'>{node['orig_author']}</a><br/>&nbsp;&nbsp;&nbsp;&nbsp;Model info:&nbsp;<a href=\'{node['sota_info']['sota_link']}\' target='_blank'>{node['sota_info']['task']}</a></div>"76            if ("Note" in node):77                markdown_str += f"<div style=\"font-size:16px; color: #a91212; text-align: left\">&nbsp;&nbsp;&nbsp;&nbsp;{node['Note']}<a href=\'{node['alt_url']}\' target='_blank'>link</a></div>"78            markdown_str += "<div style=\"font-size:16px; color: #5f5f5f; text-align: left\"><br/></div>"79 80    81    markdown_str += f"<div style=\"font-size:16px; color: #2f2f2f; text-align: left\"><b>{api_info['desc']}</b><br/></div>"82    for method in api_info["methods"]:83        lang = method["lang"]84        example = open(method["usage"]).read()85        markdown_str += f"<div style=\"font-size:16px; color: #5f5f5f; text-align: center\"><b>{lang} usage</b></div>"86        markdown_str += f"<div style=\"font-size:14px; color: #bfbfbf; text-align: left\">{example}<br/></div>"87        88    markdown_str += "<div style=\"font-size:12px; color: #9f9f9f; text-align: left\"><b><br/>Note:</b><br/>•&nbsp;Uploaded files are loaded into non-persistent memory for the duration of the computation. They are not cached</div>"89    markdown_str += "<div style=\"font-size:12px; color: #9f9f9f; text-align: left\"><br/><a href=\'https://github.com/taskswithcode/salient_object_detection_app.git\' target='_blank'>Github code</a> for this app</div>"90 91    return options_arr,markdown_str92 93 94def init_page():95    st.set_page_config(page_title='TWC - Image foreground masking or background removal with state-of-the-art models', page_icon="logo.jpg", layout='centered', initial_sidebar_state='auto',96            menu_items={97             'About': 'This app was created by taskswithcode. http://taskswithcode.com'98             99              })100    col,pad = st.columns([85,15])101 102    with col:103        st.image("long_form_logo_with_icon.png")104 105 106def run_test(config,input_file_name,display_area,uploaded_file,mask_type):107    global TMP_SEED108    display_area.text("Processing request...")109    try:110        if (uploaded_file is None):111            file_data = open(input_file_name, "rb")112            r = requests.post(config["SERVER_ADDRESS"], data={"mask":mask_type}, files={"test":file_data})113        else:114            file_data = uploaded_file.read()115            file_name = f"{TMP_DIR}/{TMP_SEED}_{str(time.time()).replace('.','_')}_{uploaded_file.name}"116            TMP_SEED += 1117            with open(file_name,"wb") as fp:118                fp.write(file_data)119            file_data = open(file_name, "rb")120            r = requests.post(config["SERVER_ADDRESS"], data={"mask":mask_type}, files={"test":file_data})121            os.remove(file_name)122        print("Servers response:",r.status_code,len(r.content))123        if (r.status_code == 200):124            size = "{:,}".format(len(r.content))125            return {"response":r.content,"size":size}126        else:127            return {"error":f"API request failed {r.status_code}"}128    except Exception as e:129        st.error("Some error occurred during prediction" + str(e))130        #st.stop()131        return {"error":f"Exception in performing image masking: {str(e)}"}132    return {} 133 134 135    136 137def display_results(results,response_info,mask):138    main_sent = f"<div style=\"font-size:14px; color: #2f2f2f; text-align: left\">{response_info}<br/><br/></div>"139    body_sent = []140    download_data = {}141    main_sent = main_sent + "\n" + '\n'.join(body_sent)142    st.markdown(main_sent,unsafe_allow_html=True)143    st.image(results["response"], caption=f'Output of Image background removal with mask: {mask}')144    st.session_state["download_ready"]  = results["response"]145    get_views("submit")146 147 148def init_session():149    init_page()150    st.session_state["model_name"] = "insprynet"151    st.session_state["download_ready"] = None    152    st.session_state["model_name"] = "ss_test"153    st.session_state["file_name"] = "default"154    st.session_state["mask_type"] = "rgba"155 156def app_main(app_mode,example_files,model_name_files,api_info_files,config_file):157  init_session()158  with open(example_files) as fp:159        example_file_names = json.load(fp) 160  with open(model_name_files) as fp:161        model_names = json.load(fp)162  with open(config_file) as fp:163        config = json.load(fp)164  with open(api_info_files) as fp:165        api_info = json.load(fp)166  curr_use_case = use_case[app_mode].split(".")[0]167  curr_use_case = use_case[app_mode].split(".")[0]168  st.markdown("<h5 style='text-align: center;'>Image foreground masking or background removal</h5>", unsafe_allow_html=True)169  st.markdown(f"<div style='color: #4f4f4f; text-align: left'>Image masking using state-of-the-art models for salient object detection(SOD). SOD use cases are<br/>&nbsp;&nbsp;&nbsp;•&nbsp;&nbsp;{use_case['1']}<br/>&nbsp;&nbsp;&nbsp;•&nbsp;&nbsp;{use_case['2']}</div>", unsafe_allow_html=True)170  st.markdown(f"<div style='color: #9f9f9f; text-align: right'>views:&nbsp;{get_views('init')}</div>", unsafe_allow_html=True)171 172 173  try:174      175      176      with st.form('twc_form'):177 178        step1_line = "Upload an image or choose an example image below"179        uploaded_file = st.file_uploader(step1_line, type=["png","jpg","jpeg"])180 181        selected_file_name = image_select("Select image", ["twc_samples/sample1.jpg", "twc_samples/sample2.jpg", "twc_samples/sample3.jpg", "twc_samples/sample4.jpg"])182 183 184        st.write("")185        mask_type = st.selectbox(label=f'Select type of masking',  186                    options = list(dict.keys(mask_types)), index=0,  key = "twc_mask_types")187        mask_type = mask_types[mask_type]188        st.write("")189        submit_button = st.form_submit_button('Run')190        options_arr,markdown_str = construct_model_info_for_display(model_names,api_info)191 192        193        input_status_area = st.empty()194        display_area = st.empty()195        if submit_button:196            start = time.time()197            if uploaded_file is not None:198                st.session_state["file_name"]  = uploaded_file.name199            else:200                st.session_state["file_name"]  = selected_file_name201            st.session_state["mask_type"]  = mask_type202            display_area.empty()203            results = run_test(config,st.session_state["file_name"],display_area,uploaded_file,mask_type)204            with display_area.container():205                if ("error" in results):206                    st.error(results["error"])207                else:208                    device = 'GPU' if torch.cuda.is_available() else 'CPU'209                    response_info = f"Computation time on {device}: {time.time() - start:.2f} secs for image size: {results['size']} bytes"210                    display_results(results,response_info,mask_type)211                    #st.json(results)212      st.download_button(213         label="Download results as png",214         data= st.session_state["download_ready"] if st.session_state["download_ready"] != None else "",215         disabled = False if st.session_state["download_ready"] != None else True,216         file_name= (st.session_state["model_name"] + "_"  + st.session_state["mask_type"] + "_" +  '_'.join(st.session_state["file_name"].split(".")[:-1]) + ".png").replace("/","_"),217         mime='image/png',218         key ="download" 219        )220      221      222 223  except Exception as e:224    st.error("Some error occurred during loading" + str(e))225    #st.stop()  226	227  st.markdown(markdown_str, unsafe_allow_html=True)228  229 230 231if __name__ == "__main__":232   app_main("1","sod_app_examples.json","sod_app_models.json","sod_apis.json","config.json")233 234