bala1802/StableDiffusionModel
0
1import torch2from tqdm.auto import tqdm3from diffusers import LMSDiscreteScheduler4 5import config6 7def construct_text_embeddings(pipe, prompt):8 text_input = pipe.tokenizer(prompt, padding='max_length', 9 max_length = pipe.tokenizer.model_max_length, truncation= True, 10 return_tensors="pt")11 uncond_input = pipe.tokenizer([""] * config.BATCH_SIZE, padding="max_length", 12 max_length= text_input.input_ids.shape[-1], 13 return_tensors="pt")14 with torch.no_grad():15 text_input_embeddings = pipe.text_encoder(text_input.input_ids.to(config.DEVICE))[0]16 with torch.no_grad():17 uncond_embeddings = pipe.text_encoder(uncond_input.input_ids.to(config.DEVICE))[0]18 19 text_embeddings = torch.cat([uncond_embeddings, text_input_embeddings])20 return text_embeddings21 22def initialize_latent(seed_number, pipe, scheduler):23 generator = torch.manual_seed(seed_number)24 latent = torch.randn((config.BATCH_SIZE, pipe.unet.config.in_channels, 25 config.HEIGHT//8, config.WIDTH//8), 26 generator = generator).to(torch.float16)27 latent = latent.to(config.DEVICE)28 latent = latent * scheduler.init_noise_sigma29 return latent30 31def run_prediction(pipe, text_embeddings, scheduler, latent, loss_function=None):32 for i, t in tqdm(enumerate(scheduler.timesteps), total = len(scheduler.timesteps)):33 latent_model_input = torch.cat([latent] * 2)34 sigma = scheduler.sigmas[i]35 latent_model_input = scheduler.scale_model_input(latent_model_input, t)36 37 with torch.no_grad():38 noise_pred = pipe.unet(latent_model_input.to(torch.float16), t, encoder_hidden_states=text_embeddings)["sample"]39 noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)40 noise_pred = noise_pred_uncond + config.GUIDANCE_SCALE * (noise_pred_text - noise_pred_uncond)41 42 if loss_function and i%5 == 0:43 latent = latent.detach().requires_grad_()44 latent_x0 = latent - sigma * noise_pred45 46 denoised_images = pipe.vae.decode((1/ 0.18215) * latent_x0).sample / 2 + 0.5 # range(0,1)47 48 loss = loss_function(denoised_images) * config.LOSS_SCALE49 print(f"loss {loss}")50 51 cond_grad = torch.autograd.grad(loss, latent)[0]52 latent = latent.detach() - cond_grad * sigma**253 54 latent = scheduler.step(noise_pred,t, latent).prev_sample55 56 return latent57 58def generate_images(pipe, seed_number, prompt, loss_function=None):59 60 scheduler = LMSDiscreteScheduler(beta_start = 0.00085, 61 beta_end = 0.012, 62 beta_schedule = "scaled_linear", 63 num_train_timesteps = 1000)64 scheduler.set_timesteps(config.NUM_INFERENCE_STEPS)65 scheduler.timesteps = scheduler.timesteps.to(torch.float32)66 67 text_embeddings = construct_text_embeddings(pipe=pipe, prompt=prompt)68 latent = initialize_latent(seed_number=seed_number, pipe=pipe, scheduler=scheduler)69 latent = run_prediction(pipe=pipe, text_embeddings=text_embeddings, 70 scheduler=scheduler, latent=latent, 71 loss_function=loss_function)72 73 return latent