leowajda/diffusion_model
1
1import gradio as gr2import numpy as np3import tensorflow as tf4from huggingface_hub import from_pretrained_keras5from diffusion_sampler import DiffusionSampler6 7devices = '\n'.join([f'- {device.name}' for device in tf.config.list_physical_devices('GPU')]) or 'No GPU devices found.'8print(f"GPUs available: {devices}")9 10scheduler_button = gr.Radio(11 choices=["Linear", "Cosine"],12 label="Noise Scheduler",13 value="Linear",14 info="""15 Decides whether to employ a model trained with a linear scheduler, 16 as proposed by Jonathan Ho et al. in 'Denoising Diffusion Probabilistic Models',17 or the cosine variant introduced by Alex Nichol et al. in 'Improved Denoising Diffusion Probabilistic Models'.18 """,19)20 21sampling_button = gr.Radio(22 choices=["DDPM", "DDIM"],23 label="Sampling Procedure",24 value="DDPM",25 info="""26 Selects either the stocasthic sampling procedure described by Jonathan Ho et al. in 'Denoising Diffusion Probabilistic Models',27 or the implicit variant proposed by Jiaming Song et al. in 'Denoising Diffusion Implicit Models'.28 For the latter, it is also necessary to specify the sub-sequence strategy and the number of sampling steps.29 """,30)31 32subsequence_button = gr.Radio(33 choices=["Linear", "Quadratic"],34 label="Sub-Sequence",35 value="Linear",36 info="""37 Specific to DDIM sampling, this parameter chooses the procedure 38 for forming the sub-sequence employed during the sampling process.39 """,40)41 42ema_button = gr.Checkbox(43 value=True,44 label="Exponential Moving Average",45 info="""46 Whether to invoke the network with the applied exponential moving average on the model's weights.47 """48)49 50images_button = gr.Number(51 label="Number of images to generate",52 value=10,53 precision=0,54 minimum=1,55 maximum=64,56 info="""57 The number of images to be generated.58 Larger batch sizes result in longer inference times.59 """60)61 62step_button = gr.Slider(63 minimum=700,64 value=1_000,65 maximum=1_000,66 randomize=True,67 label="Number of sampling steps",68 info="""69 Relevant exclusively to DDIM sampling, this parameter determines the number of steps to be utilized during sampling.70 The default value is set to 1000 in the case of DDPM sampling.71 """72)73 74gallery = gr.Gallery(75 columns=4,76 allow_preview=False,77 show_download_button=False,78 show_share_button=False,79 label="""80 Generated Flowers81 """82)83 84diffusion_models = {85 "linear":86 DiffusionSampler(87 model=from_pretrained_keras("leowajda/linear_diffusion", cache_dir="cache"),88 ema_model=from_pretrained_keras("leowajda/linear_diffusion_ema", cache_dir="cache"),89 noise_scheduler="linear"90 ),91 92 "cosine":93 DiffusionSampler(94 model=from_pretrained_keras("leowajda/cosine_diffusion", cache_dir="cache"),95 ema_model=from_pretrained_keras("leowajda/cosine_diffusion_ema", cache_dir="cache"),96 noise_scheduler="cosine"97 )98}99 100 101def call_model(102 model_to_call: str,103 sample_strategy: str = "ddim",104 step_strategy: str = "uniform",105 ema: bool = True,106 steps: int = 1_000,107 num_images: int = 0,108):109 diffusion_model = diffusion_models[model_to_call.lower()]110 images = diffusion_model.generate_images(111 num_images=num_images,112 steps=steps,113 sample_strategy=sample_strategy.lower(),114 step_strategy=step_strategy.lower(),115 ema=ema,116 )117 118 return images.numpy().astype(np.uint8)119 120 121demo = gr.Interface(122 fn=call_model,123 inputs=[scheduler_button, sampling_button, subsequence_button, ema_button, step_button, images_button],124 outputs=gallery,125 cache_examples=True,126 title="""Unconditional Image Generation Through Denoising Diffusion Implicit Models""",127 examples=[128 ["Linear", "DDPM", "Linear", True, 1_000, 25],129 ["Linear", "DDPM", "Linear", False, 1_000, 25],130 ["Cosine", "DDPM", "Linear", False, 750, 25],131 ["Cosine", "DDPM", "Linear", True, 750, 25],132 ["Linear", "DDIM", "Linear", True, 750, 25],133 ],134 description=""" 135 <p align="center">136 Supervisor: <strong>Wojciech Oronowicz – Jaśkowiak, PhD</strong>137  138 Author: <strong>Leonardo Wajda</strong>139  140 Specialization: <strong>Intelligent Data Processing Systems</strong>141 </p>142 """,143)144 145demo.queue(default_concurrency_limit=None).launch()146 