atmptdie/cycleGAN
0
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 