Team Ai
Apppublic

bala1802/StableDiffusionModel

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
image_generator.py73 linesDownload Raw Back to root
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