Gosula/Stable_diffusion_model
0
1from base64 import b64encode2from utils import *3import numpy4import torch5from diffusers import AutoencoderKL, LMSDiscreteScheduler, UNet2DConditionModel6from huggingface_hub import notebook_login7 8# For video display:9from IPython.display import HTML10from matplotlib import pyplot as plt11from pathlib import Path12from PIL import Image13from torch import autocast14from torchvision import transforms as tfms15from tqdm.auto import tqdm16from transformers import CLIPTextModel, CLIPTokenizer, logging17import os18import shutil19from device import torch_device,vae,text_encoder,unet,tokenizer,scheduler,token_emb_layer,pos_emb_layer,position_embeddings20torch.manual_seed(1)21if not (Path.home()/'.cache/huggingface'/'token').exists(): notebook_login()22 23# Supress some unnecessary warnings when loading the CLIPTextModel24logging.set_verbosity_error()25 26# Set device27 28 29def generate_distorted_image(pil_image,vae):30 # View a noised version31 encoded = pil_to_latent(pil_image)32 noise = torch.randn_like(encoded) # Random noise33 34 sampling_step = 5 # Equivalent to step 10 out of 15 in the schedule above35 # encoded_and_noised = scheduler.add_noise(encoded, noise, timestep) # Diffusers 0.3 and below36 encoded_and_noised = scheduler.add_noise(encoded, noise, timesteps=torch.tensor([scheduler.timesteps[sampling_step]]))37 return latents_to_pil(encoded_and_noised)[0] # Display38 39def set_timesteps(scheduler, num_inference_steps):40 scheduler.set_timesteps(num_inference_steps)41 scheduler.timesteps = scheduler.timesteps.to(torch.float32)42 43 44# Some settings45def generate_image(prompt,concept_embed,num_inference_steps=50,color_postprocessing=False,noised_image=False,loss_scale=10,seed=42):46 height = 512 # default height of Stable Diffusion47 width = 512 # default width of Stable Diffusion48 num_inference_steps = num_inference_steps # Number of denoising steps49 guidance_scale = 7.5 # Scale for classifier-free guidance50 generator = torch.manual_seed(seed) # Seed generator to create the inital latent noise51 batch_size = 152 # Define the directory name53 directory_name = "steps"54 55 # Check if the directory exists, and if so, delete it56 if os.path.exists(directory_name):57 shutil.rmtree(directory_name)58 59 #Create the directory60 os.makedirs(directory_name)61 # Prep text62 #text_input = tokenizer(prompt, padding="max_length", max_length=tokenizer.model_max_length, truncation=True, return_tensors="pt")63# with torch.no_grad():64# text_embeddings = text_encoder(text_input.input_ids.to(torch_device))[0]65 66 text_input = tokenizer(prompt, padding="max_length", max_length=tokenizer.model_max_length, truncation=True, return_tensors="pt")67 input_ids = text_input.input_ids.to(torch_device)68 custom_style_token=tokenizer.encode("cs",add_special_token=False)[0]69 # Get token embeddings70 token_embeddings = token_emb_layer(input_ids)71 embed_key=list(concept_embed.keys())[0]72 # The new embedding. In this case just the input embedding of token 2368...73 replacement_token_embedding = concept_embed[embed_key]74 token_embeddings[0,torch.where(input_ids[0]==custom_style_token)]=replacement_token_embedding.to(torch_device)75 # Combine with pos embs76 input_embeddings = token_embeddings + position_embeddings77 78 # Feed through to get final output embs79 modified_output_embeddings = get_output_embeds(input_embeddings)80 81 max_length = text_input.input_ids.shape[-1]82 uncond_input = tokenizer(83 [""] * batch_size, padding="max_length", max_length=max_length, return_tensors="pt"84 )85 with torch.no_grad():86 uncond_embeddings = text_encoder(uncond_input.input_ids.to(torch_device))[0]87 text_embeddings = torch.cat([uncond_embeddings, modified_output_embeddings])88 89 # minor fix to ensure MPS compatibility, fixed in diffusers PR 392590 91 set_timesteps(scheduler,num_inference_steps)92 93 # Prep latents94 latents = torch.randn(95 (batch_size, unet.in_channels, height // 8, width // 8),96 generator=generator,97 )98 latents = latents.to(torch_device)99 latents = latents * scheduler.init_noise_sigma # Scaling (previous versions did latents = latents * self.scheduler.sigmas[0]100 101 # Loop102 with autocast("cuda"): # will fallback to CPU if no CUDA; no autocast for MPS103 for i, t in tqdm(enumerate(scheduler.timesteps), total=len(scheduler.timesteps)):104 # expand the latents if we are doing classifier-free guidance to avoid doing two forward passes.105 latent_model_input = torch.cat([latents] * 2)106 sigma = scheduler.sigmas[i]107 # Scale the latents (preconditioning):108 # latent_model_input = latent_model_input / ((sigma**2 + 1) ** 0.5) # Diffusers 0.3 and below109 latent_model_input = scheduler.scale_model_input(latent_model_input, t)110 111 # predict the noise residual112 with torch.no_grad():113 noise_pred = unet(latent_model_input, t, encoder_hidden_states=text_embeddings).sample114 115 # perform guidance116 noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)117 noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)118 119 # compute the previous noisy sample x_t -> x_t-1120 # latents = scheduler.step(noise_pred, i, latents)["prev_sample"] # Diffusers 0.3 and below121 122 #latents = torch.tensor(initial_latents, requires_grad=True)123 ### ADDITIONAL GUIDANCE ###124 # Requires grad on the latents125 if color_postprocessing:126 latents = latents.detach().requires_grad_()127 128 # Get the predicted x0:129 latents_x0 = latents - sigma * noise_pred130 131 # Decode to image space132 denoised_images = vae.decode((1 / 0.18215) * latents_x0).sample / 2 + 0.5133 #denoised_images = vae.decode((1 / 0.18215) * latents_x0) / 2 + 0.5 # (0, 1)134 135 # Calculate loss136 loss = orange_loss(denoised_images) * loss_scale137 #loss = color_loss(denoised_images,postporcessing_color) * color_loss_scale138 if i%10==0:139 print(i, 'loss:', loss.item())140 141 # Get gradient142 cond_grad = -torch.autograd.grad(loss, latents)[0]143 144 # Modify the latents based on this gradient145 latents = latents.detach() + cond_grad * sigma**2146 147 148 ### And saving as before ###149 # Get the predicted x0:150 latents_x0 = latents - sigma * noise_pred151 im_t0 = latents_to_pil(latents_x0)[0]152 153 # And the previous noisy sample x_t -> x_t-1154 latents = scheduler.step(noise_pred, t, latents)["prev_sample"]155 im_next = latents_to_pil(latents)[0]156 157 # Combine the two images and save for later viewing158 im = Image.new('RGB', (1024, 512))159 im.paste(im_next, (0, 0))160 im.paste(im_t0, (512, 0))161 im.save(f'steps/{i:04}.jpeg')162 163 else:164 latents = scheduler.step(noise_pred, t, latents).prev_sample165 166 167 if noised_image:168 output = generate_distorted_image(latents_to_pil(latents)[0],vae)169 else:170 output = latents_to_pil(latents)[0]171 172 return output173def get_output_embeds(input_embeddings):174 # CLIP's text model uses causal mask, so we prepare it here:175 bsz, seq_len = input_embeddings.shape[:2]176 causal_attention_mask = text_encoder.text_model._build_causal_attention_mask(bsz, seq_len, dtype=input_embeddings.dtype)177 178 # Getting the output embeddings involves calling the model with passing output_hidden_states=True179 # so that it doesn't just return the pooled final predictions:180 encoder_outputs = text_encoder.text_model.encoder(181 inputs_embeds=input_embeddings,182 attention_mask=None, # We aren't using an attention mask so that can be None183 causal_attention_mask=causal_attention_mask.to(torch_device),184 output_attentions=None,185 output_hidden_states=True, # We want the output embs not the final output186 return_dict=None,187 )188 189 # We're interested in the output hidden state only190 output = encoder_outputs[0]191 192 # There is a final layer norm we need to pass these through193 output = text_encoder.text_model.final_layer_norm(output)194 195 # And now they're ready!196 return output197 