taskswithcode/salient-object-detection
13
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\"> • Model: <a href=\'{node['paper_url']}\' target='_blank'>{node['name']}</a><br/> Code released by: <a href=\'{node['orig_author_url']}\' target='_blank'>{node['orig_author']}</a><br/> Model info: <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\"> {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/>• 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/> • {use_case['1']}<br/> • {use_case['2']}</div>", unsafe_allow_html=True)170 st.markdown(f"<div style='color: #9f9f9f; text-align: right'>views: {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 