File size: 6,526 Bytes
cc1e3fb
 
 
 
 
 
 
 
 
 
21605a8
cc1e3fb
 
 
 
 
 
 
 
 
 
 
 
 
9b79ed4
d719eca
cc1e3fb
 
ad87121
 
9b79ed4
cc1e3fb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
15567b5
cc1e3fb
 
 
 
cef7555
cc1e3fb
 
cef7555
f87b5b4
cc1e3fb
9b4f3c3
 
cc1e3fb
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
# -*- coding: utf-8 -*-
"""
Created on Sun Apr  7 05:35:36 2024

@author: aditya.vipradas
"""

# import file system libraries
import os
from pathlib import Path
#from huggingface_hub import notebook_login
from tqdm.auto import tqdm
import gradio as gr

# import image and visualization libraries
from PIL import Image

# import modeling libraries
import torch
from diffusers import AutoencoderKL, LMSDiscreteScheduler, UNet2DConditionModel
from transformers import CLIPTextModel, CLIPTokenizer
from torchvision import transforms as tfms

# check huggingface token
#if not (Path.home()/'.cache/huggingface/token').exists():
#    notebook_login()

# set torch device and suppress duplication warnings
#torch_device = "cuda" if torch.cuda.is_available() else "cpu"
torch_device = "cpu"
#os.environ['HF_HUB_DISABLE_SYMLINKS_WARNING'] = "1"

# load the autoencoder
vae = AutoencoderKL.from_pretrained("CompVis/stable-diffusion-v1-4",
                                    subfolder="vae")

# load the tokenizer and text encoder
tokenizer = CLIPTokenizer.from_pretrained("openai/clip-vit-large-patch14")
text_encoder = CLIPTextModel.from_pretrained("openai/clip-vit-large-patch14")

# load the unet
unet = UNet2DConditionModel.from_pretrained("CompVis/stable-diffusion-v1-4",
                                            subfolder="unet")

# noise scheduler (linear multi-step)
scheduler = LMSDiscreteScheduler(beta_start=0.00085,
                                 beta_end=0.012,
                                 beta_schedule="scaled_linear",
                                 num_train_timesteps=1000)


def latents_to_pil(latents):
    # remove the scaling as mentioned in the documentation
    latents = (1 / 0.18215) * latents
    with torch.no_grad():
        # this generates images with (1, 3, 512, 512) dimensions.
        # permute them to (1, 512, 512, 3) later
        image = vae.decode(latents).sample
    image = (image / 2 + 0.5).clamp(0, 1)
    image = image.detach().cpu().permute(0, 2, 3, 1).numpy()
    images = (image * 255).round().astype("uint8")
    pil_images = [Image.fromarray(image) for image in images]
    return pil_images[0]


def pil_to_latent(input_im):
    with torch.no_grad():
        latent = vae.encode(tfms.ToTensor()(input_im).unsqueeze(0)*2-1)
    return 0.18215 * latent.latent_dist.sample()


# fixed parameters
height = 512
width = 512
batch_size = 1

uncond_input = tokenizer([""] * batch_size,
                         padding="max_length",
                         max_length=tokenizer.model_max_length,
                         return_tensors="pt")


with torch.no_grad():
    uncond_embeddings = text_encoder(uncond_input.input_ids)[0]


def diffusion(image_conditioned, prompt_image, prompt, artist,
              sampling_step, guidance_scale, num_inference_steps, seed):
    generator = torch.manual_seed(seed)
    scheduler.set_timesteps(num_inference_steps+1)
        
    if artist != "":
        prompt = prompt + ", " + artist + " style"
    
    text_input = tokenizer([prompt],
                           padding="max_length",
                           max_length=tokenizer.model_max_length,
                           truncation=True,
                           return_tensors="pt")
    with torch.no_grad():
        text_embeddings = text_encoder(text_input.input_ids)[0]
    text_embeddings = torch.cat([uncond_embeddings, text_embeddings])
    
    noise = torch.randn((batch_size, unet.config.in_channels, height // 8, width // 8), generator=generator)
    if image_conditioned:
        encoded_sketch = pil_to_latent(Image.fromarray(prompt_image).
                                       resize((height, width)))
        encoded_noised = scheduler.add_noise(encoded_sketch, noise,
                                         timesteps=torch.tensor([scheduler.timesteps[sampling_step]]))
    else:
        encoded_noised = noise * scheduler.init_noise_sigma
        
    
    for i, t in enumerate(scheduler.timesteps):
        if i >= sampling_step:
            latent_model_input = torch.cat([encoded_noised] * 2)
            sigma = scheduler.sigmas[i]
            latent_model_input = scheduler.scale_model_input(latent_model_input, t)

            with torch.no_grad():
                noise_pred = unet(latent_model_input, t, encoder_hidden_states = text_embeddings).sample
            
            noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
            noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)

            encoded_noised = scheduler.step(noise_pred, t, encoded_noised).prev_sample
        
        yield latents_to_pil(encoded_noised), f"Generating your image...step {i}/{num_inference_steps}"
        
        if i == num_inference_steps:
            yield latents_to_pil(encoded_noised), "Image generation complete"
            



demo = gr.Interface(
    theme=gr.themes.Soft(),
    title="Image- and Text- Conditioned Stable Diffusion",
    description="Latent stable diffusion implementation with image and text conditioning. Additional artistic rendering options are also provided.\n \
        Potential Applications: \n1. Implement style of one image on another \n2. Convert concepts (drawing, paintings) to realistic images \n3. Guide text prompts with images.",
    fn=diffusion,
    inputs=[gr.Checkbox(value=False, label="Condition on reference image", info="Recommendation: Select sampling step 7 if selected else 0"),
            gr.Image(height=height, width=width, label='Select Reference Image'),
            gr.Textbox(value="A highly realistic and majestic lion with wavy mane, high-definition", info="Enter the text prompt", max_lines=1, label='Prompt'),
            gr.Radio(["Van Gogh", "Johannes Vermeer", "Claude Monet", "Pablo Picasso", "Frida Kahlo", "None"], value="Van Gogh", label="Render", info="Choose your artist"),
            gr.Slider(0, 20, value=0, step=1, label='Sampling step', info="Higher step adds less noise (between 0 and 20)"),
            gr.Slider(0, 14, value=7.5, step=0.5, label='Guidance scale', info="Adherance to prompt (between 0 and 14)"),
            gr.Slider(0, 100, value=70, step=1, label='Number of inference steps', info="Choose between 0 and 100"),
            gr.Slider(0, 100, value=50, step=1, label='Random Seed', info="Change to generate a different image (between 0 and 100)")],
    outputs=[gr.Image(height=height, width=width, label='Stable Diffusion Progress'),
             gr.Textbox(max_lines=1, label='Progress')])

demo.launch()