CoolFace
Apppublic

atmptdie/cycleGAN

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
streamlit_app.py120 linesDownload Raw Back to src
1import streamlit as st2 3import os4import logging5import typing as tp6from dataclasses import dataclass7 8import cycle_gan9import config10 11"""12# Welcome to Streamlit!13 14Edit `/streamlit_app.py` to customize this app to your heart's desire :heart:.15If you have any questions, checkout our [documentation](https://docs.streamlit.io) and [community16forums](https://discuss.streamlit.io).17 18In the meantime, below is an example of what you can do with just a few lines of code:19"""20 21 22@dataclass23class LoadedData:24    model: cycle_gan.CycleGAN25    datasetA: cycle_gan.ImageDatasetNoLabel26    datasetB: cycle_gan.ImageDatasetNoLabel27    de_normalize_a: tp.Callable28    de_normalize_b: tp.Callable29 30 31@st.cache_resource32def load_data(model_type: config.SavedModel):33    weights_path = os.path.join(os.getcwd(), model_type.value.weights_path)34    imgs_dir = os.path.join(os.getcwd(), model_type.value.imgs_path)35 36    model = cycle_gan.load_model(weights_path)37 38    tranforms_a, tranforms_b, de_normalize_a, de_normalize_b = config.get_transforms(39        model_type,40    )41 42    datasetA = cycle_gan.ImageDatasetNoLabel(os.path.join(imgs_dir, "A"), tranforms_a)43    datasetB = cycle_gan.ImageDatasetNoLabel(os.path.join(imgs_dir, "B"), tranforms_b)44 45    return LoadedData(46        model=model,47        datasetA=datasetA,48        datasetB=datasetB,49        de_normalize_a=de_normalize_a,50        de_normalize_b=de_normalize_b,51    )52 53 54def transform_image(generator, dataset, index):55    img = dataset[index].unsqueeze(0)56    fake_img = generator(img)[0]57 58    return fake_img59 60 61def main():62    st.set_page_config(page_title="CycleGAN Style Transfer", layout="wide")63    st.title("CycleGAN Demo")64 65    st.sidebar.header("Controls")66    model_choice = st.sidebar.selectbox(67        "Select Variant", ["Van Gogh ↔ Photo", "Samoyed ↔ Newfoundland"]68    )69 70    if "Samoyed" in model_choice:71        logging.info("loading dogs model...")72        data = load_data(config.SavedModel.DOGS)73    else:74        logging.info("loading vangogh model...")75        data = load_data(config.SavedModel.VANGOGH)76 77    logging.info("model loaded")78 79    direction = st.sidebar.selectbox("Select Variant", ["A -> B", "B -> A"])80 81    selected_dataset = data.datasetA if direction[0] == "A" else data.datasetB82    denorm_left = data.de_normalize_a if direction[0] == "A" else data.de_normalize_b83    denorm_right = data.de_normalize_b if direction[0] == "A" else data.de_normalize_a84    selected_generator = data.model.G_AB if direction[0] == "A" else data.model.G_BA85 86    logging.info(f"selected direction: {direction}")87 88    images = {f"image_{i}": i for i in range(len(selected_dataset))}89    logging.info(f"have {len(images)} available images")90 91    col1, col2 = st.columns(2)92    selected_image = st.sidebar.selectbox(93        "Select an image to transform:", list(images.keys())94    )95    index = images[selected_image]96    logging.info(f"selected image: {selected_image}")97 98    logging.info("showing original image...")99    col1.image(100        denorm_left(selected_dataset[index]),101        caption=f"Original image",102        use_container_width=True,103        clamp=True,104    )105 106    logging.info("transforming image...")107    new_image = transform_image(selected_generator, selected_dataset, index)108    logging.info("image transformed")109 110    logging.info("showing transformed image...")111    col2.image(112        denorm_right(new_image),113        caption=f"Transformed image",114        use_container_width=True,115        clamp=True,116    )117 118 119main()120