CoolFace
Apppublic

otroivan/Whitebox-Style-Transfer-Editing

sourceHugging Facemitupdated 4y agoView on Hugging Face
0likes
Whitebox_style_transfer.py313 linesDownload Raw Back to root
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