otroivan/Whitebox-Style-Transfer-Editing
0
1import base642import datetime3import os4import sys5from io import BytesIO6from pathlib import Path7import numpy as np8import requests9import torch10import torch.nn.functional as F11from PIL import Image12 13PACKAGE_PARENT = 'wise'14SCRIPT_DIR = os.path.dirname(os.path.realpath(os.path.join(os.getcwd(), os.path.expanduser(__file__))))15sys.path.append(os.path.normpath(os.path.join(SCRIPT_DIR, PACKAGE_PARENT)))16 17import streamlit as st18from streamlit.logger import get_logger19from st_click_detector import click_detector20import streamlit.components.v1 as components21from streamlit.source_util import get_pages22from streamlit_extras.switch_page_button import switch_page23 24from demo_config import HUGGING_FACE25from parameter_optimization.parametric_styletransfer import single_optimize26from parameter_optimization.parametric_styletransfer import CONFIG as ST_CONFIG27from parameter_optimization.strotss_org import strotss, pil_resize_long_edge_to28import helpers.session_state as session_state29from helpers import torch_to_np, np_to_torch30from effects import get_default_settings, MinimalPipelineEffect 31 32st.set_page_config(layout="wide")33BASE_URL = "https://ivpg.hpi3d.de/wise/wise-demo/images/"34LOGGER = get_logger(__name__)35 36effect_type = "minimal_pipeline"37 38if "click_counter" not in st.session_state:39 st.session_state.click_counter = 140 41if "action" not in st.session_state:42 st.session_state["action"] = ""43 44content_urls = [45 {46 "name": "Portrait", "id": "portrait",47 "src": BASE_URL + "/content/portrait.jpeg"48 },49 {50 "name": "Tuebingen", "id": "tubingen",51 "src": BASE_URL + "/content/tubingen.jpeg"52 },53 {54 "name": "Colibri", "id": "colibri",55 "src": BASE_URL + "/content/colibri.jpeg"56 }57]58 59style_urls = [60 {61 "name": "Starry Night, Van Gogh", "id": "starry_night",62 "src": BASE_URL + "/style/starry_night.jpg"63 },64 {65 "name": "The Scream, Edward Munch", "id": "the_scream",66 "src": BASE_URL + "/style/the_scream.jpg"67 },68 {69 "name": "The Great Wave, Ukiyo-e", "id": "wave",70 "src": BASE_URL + "/style/wave.jpg"71 },72 {73 "name": "Woman with Hat, Henry Matisse", "id": "woman_with_hat",74 "src": BASE_URL + "/style/woman_with_hat.jpg"75 }76]77 78 79def last_image_clicked(type="content", action=None, ):80 kw = "last_image_clicked" + "_" + type81 if action:82 session_state.get(**{kw: action})83 elif kw not in session_state.get():84 return None85 else:86 return session_state.get()[kw]87 88 89@st.cache90def _retrieve_from_id(clicked, urls):91 src = [x["src"] for x in urls if x["id"] == clicked][0]92 img = Image.open(requests.get(src, stream=True).raw)93 return img, src94 95 96def store_img_from_id(clicked, urls, imgtype):97 img, src = _retrieve_from_id(clicked, urls)98 session_state.get(**{f"{imgtype}_im": img, f"{imgtype}_render_src": src, f"{imgtype}_id": clicked})99 100 101def img_choice_panel(imgtype, urls, default_choice, expanded):102 with st.expander(f"Select {imgtype} image:", expanded=expanded):103 html_code = '<div class="column" style="display: flex; flex-wrap: wrap; padding: 0 4px;">'104 for url in urls:105 html_code += f"<a href='#' id='{url['id']}' style='padding: 0px 5px'><img height='160px' style='margin-top: 8px;' src='{url['src']}'></a>"106 html_code += "</div>"107 clicked = click_detector(html_code)108 109 if not clicked and st.session_state["action"] not in ("uploaded", "switch_page_from_local_edits", "switch_page_from_presets", "slider_change", "reset"): # default val110 store_img_from_id(default_choice, urls, imgtype)111 112 st.write("OR: ")113 114 with st.form(imgtype + "-form", clear_on_submit=True):115 uploaded_im = st.file_uploader(f"Load {imgtype} image:", type=["png", "jpg"], )116 upload_pressed = st.form_submit_button("Upload")117 118 if upload_pressed and uploaded_im is not None:119 img = Image.open(uploaded_im)120 buffered = BytesIO()121 img.save(buffered, format="JPEG")122 encoded = base64.b64encode(buffered.getvalue()).decode()123 # session_state.get(uploaded_im=img, content_render_src=f"data:image/jpeg;base64,{encoded}")124 session_state.get(**{f"{imgtype}_im": img, f"{imgtype}_render_src": f"data:image/jpeg;base64,{encoded}",125 f"{imgtype}_id": "uploaded"})126 st.session_state["action"] = "uploaded"127 st.write("uploaded.")128 129 last_clicked = last_image_clicked(type=imgtype)130 print("last_clicked", last_clicked, "clicked", clicked, "action", st.session_state["action"] )131 if not upload_pressed and clicked != "": # trigger when no file uploaded132 if last_clicked != clicked: # only activate when content was actually clicked133 store_img_from_id(clicked, urls, imgtype)134 last_image_clicked(type=imgtype, action=clicked)135 st.session_state["action"] = "clicked"136 st.session_state.click_counter += 1 # hack to get page to reload at top137 138 state = session_state.get()139 st.sidebar.write(f'Selected {imgtype} image:')140 st.sidebar.markdown(f'<img src="{state[f"{imgtype}_render_src"]}" width=240px></img>', unsafe_allow_html=True)141 142 143def optimize(effect, preset, result_image_placeholder):144 content = st.session_state["Content_im"]145 style = st.session_state["Style_im"]146 result_image_placeholder.text("<- Custom content/style needs to be style transferred")147 optimize_button = st.sidebar.button("Optimize Style Transfer")148 if optimize_button:149 if HUGGING_FACE:150 result_image_placeholder.warning("NST optimization is currently disabled in this HuggingFace Space because it takes ~5min to optimize. To try it out, please clone the repo and change the huggingface variable in demo_config.py")151 st.stop()152 153 result_image_placeholder.text("Executing NST to create reference image..")154 base_dir = f"result/{datetime.datetime.now().strftime(r'%Y-%m-%d %H.%Mh %Ss')}"155 os.makedirs(base_dir)156 with st.spinner(text="Running NST"):157 reference = strotss(pil_resize_long_edge_to(content, 1024),158 pil_resize_long_edge_to(style, 1024), content_weight=16.0,159 device=torch.device("cuda"), space="uniform")160 progress_bar = result_image_placeholder.progress(0.0)161 ref_save_path = os.path.join(base_dir, "reference.jpg")162 content_save_path = os.path.join(base_dir, "content.jpg")163 resize_to = 720164 reference = pil_resize_long_edge_to(reference, resize_to)165 reference.save(ref_save_path)166 content.save(content_save_path)167 ST_CONFIG["n_iterations"] = 300168 with st.spinner(text="Optimizing parameters.."):169 vp, content_img_cuda = single_optimize(effect, preset, "l1", content_save_path, str(ref_save_path),170 write_video=False, base_dir=base_dir,171 iter_callback=lambda i: progress_bar.progress(172 float(i) / ST_CONFIG["n_iterations"]))173 return content_img_cuda.detach(), vp.cuda().detach()174 else:175 if not "result_vp" in st.session_state:176 st.stop()177 else:178 return st.session_state["effect_input"], st.session_state["result_vp"]179 180 181@st.cache(hash_funcs={MinimalPipelineEffect: id})182def create_effect():183 effect, preset, param_set = get_default_settings(effect_type)184 effect.enable_checkpoints()185 effect.cuda()186 return effect, preset187 188 189def load_visual_params(vp_path: str, img_org: Image, org_cuda: torch.Tensor, effect) -> torch.Tensor:190 if Path(vp_path).exists():191 vp = torch.load(vp_path).detach().clone()192 vp = F.interpolate(vp, (img_org.height, img_org.width))193 if len(effect.vpd.vp_ranges) == vp.shape[1]:194 return vp195 # use preset and save it196 vp = effect.vpd.preset_tensor(preset, org_cuda, add_local_dims=True)197 torch.save(vp, vp_path)198 return vp199 200 201# @st.cache(hash_funcs={torch.Tensor: id})202@st.experimental_memo203def load_params(content_id, style_id):#, effect):204 preoptim_param_path = os.path.join("precomputed", effect_type, content_id, style_id)205 img_org = Image.open(os.path.join(preoptim_param_path, "input.png"))206 content_cuda = np_to_torch(img_org).cuda()207 vp_path = os.path.join(preoptim_param_path, "vp.pt")208 vp = load_visual_params(vp_path, img_org, content_cuda, effect)209 return content_cuda, vp210 211 212def render_effect(effect, content_cuda, vp):213 with torch.no_grad():214 result_cuda = effect(content_cuda, vp)215 img_res = Image.fromarray((torch_to_np(result_cuda) * 255.0).astype(np.uint8))216 return img_res217 218 219result_container = st.container()220coll1, coll2 = result_container.columns([3,2])221coll1.header("Result")222coll2.header("Global Edits")223result_image_placeholder = coll1.empty()224result_image_placeholder.markdown("## loading..")225 226img_choice_panel("Content", content_urls, "portrait", expanded=True)227img_choice_panel("Style", style_urls, "starry_night", expanded=True)228 229state = session_state.get()230content_id = state["Content_id"]231style_id = state["Style_id"]232 233effect, preset = create_effect()234 235print("content id, style id", content_id, style_id )236if st.session_state["action"] == "uploaded":237 content_img, _vp = optimize(effect, preset, result_image_placeholder)238elif st.session_state["action"] in ("switch_page_from_local_edits", "switch_page_from_presets", "slider_change") or \239 content_id == "uploaded" or style_id == "uploaded":240 print("restore param")241 _vp = st.session_state["result_vp"]242 content_img = st.session_state["effect_input"]243else:244 print("load_params")245 content_img, _vp = load_params(content_id, style_id)#, effect)246 247vp = torch.clone(_vp)248 249 250def reset_params(means, names):251 for i, name in enumerate(names):252 st.session_state["slider_" + name] = means[i]253 254def on_slider():255 st.session_state["action"] = "slider_change"256 257 258with coll2:259 show_params_names = [ 'bumpScale', "bumpOpacity", "contourOpacity"]260 display_means = []261 def create_slider(name):262 mean = torch.mean(vp[:, effect.vpd.name2idx[name]]).item()263 display_mean = mean + 0.5264 display_means.append(display_mean)265 if "slider_" + name not in st.session_state or st.session_state["action"] != "slider_change": 266 st.session_state["slider_" + name] = display_mean267 slider = st.slider(f"Mean {name}: ", 0.0, 1.0, step=0.05, key="slider_" + name, on_change=on_slider)268 vp[:, effect.vpd.name2idx[name]] += slider - display_mean269 vp.clamp_(-0.5, 0.5)270 271 for name in show_params_names:272 create_slider(name)273 274 others_idx = set(range(len(effect.vpd.vp_ranges))) - set([effect.vpd.name2idx[name] for name in show_params_names])275 others_names = [effect.vpd.vp_ranges[i][0] for i in sorted(list(others_idx))]276 other_param = st.selectbox("Other parameters: ", others_names)277 create_slider(other_param)278 279 280 reset_button = st.button("Reset Parameters", on_click=reset_params, args=(display_means, show_params_names))281 if reset_button:282 st.session_state["action"] = "reset"283 st.experimental_rerun()284 285 edit_locally_btn = st.button("Edit Local Parameter Maps")286 if edit_locally_btn:287 switch_page('️ local edits')288 289 apply_presets = st.button("Paint Presets")290 if apply_presets:291 switch_page("Apply_preset")292 293img_res = render_effect(effect, content_img, vp)294 295st.session_state["result_vp"] = vp296st.session_state["effect_input"] = content_img297st.session_state["last_result"] = img_res298 299with coll1:300 # width = int(img_res.width * 500 / img_res.height)301 result_image_placeholder.image(img_res)#, width=width)302 303# a bit hacky way to return focus to top of page after clicking on images304components.html(305 f"""306 <p>{st.session_state.click_counter}</p>307 <script>308 window.parent.document.querySelector('section.main').scrollTo(0, 0);309 </script>310 """,311 height=0312)313 