Team Ai
Apppublic

Gosula/Stable_diffusion_model

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
stablediffusion.py197 linesDownload Raw Back to root
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