apapiu/transformer_diffusion
1
1import gradio as gr2from PIL import Image3import requests4 5from tld.denoiser import Denoiser6from tld.diffusion import DiffusionGenerator7 8from diffusers import AutoencoderKL, AutoencoderTiny9from tqdm import tqdm10import clip11import torch12import numpy as np13import torchvision.utils as vutils14import torchvision.transforms as transforms15from torch.utils.data import DataLoader, TensorDataset16from PIL import Image17 18device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")19to_pil = transforms.ToPILImage()20 21 22def download_file(url, filename):23 24 with requests.get(url, stream=True) as r:25 r.raise_for_status() 26 with open(filename, 'wb') as f:27 for chunk in r.iter_content(chunk_size=8192): 28 f.write(chunk)29 30@torch.no_grad()31def encode_text(label, model):32 text_tokens = clip.tokenize(label, truncate=True).to(device)33 text_encoding = model.encode_text(text_tokens)34 return text_encoding.cpu()35 36def generate_image_from_text(prompt, class_guidance=6, seed=11, num_imgs=1, img_size = 32):37 38 n_iter = 1539 nrow = int(np.sqrt(num_imgs))40 41 cur_prompts = [prompt]*num_imgs42 labels = encode_text(cur_prompts, clip_model)43 out, out_latent = diffuser.generate(labels=labels,44 num_imgs=num_imgs,45 class_guidance=class_guidance,46 seed=seed,47 n_iter=n_iter,48 exponent=1,49 scale_factor=8,50 sharp_f=0,51 bright_f=052 )53 54 out = to_pil((vutils.make_grid((out+1)/2, nrow=nrow, padding=4)).float().clip(0, 1))55 56 out.save(f'{prompt}_cfg:{class_guidance}_seed:{seed}.png')57 58 print("Images Generated and Saved. They will shortly output below.")59 return out60 61###config:62vae_scale_factor = 863img_size = 3264model_dtype = torch.float3265 66file_url = "https://huggingface.co/apapiu/small_ldt/resolve/main/state_dict_378000.pth"67local_filename = "state_dict_378000.pth"68download_file(file_url, local_filename)69 70 71denoiser = Denoiser(image_size=32, noise_embed_dims=256, patch_size=2,72 embed_dim=768, dropout=0, n_layers=12)73 74 75state_dict = torch.load('state_dict_378000.pth', map_location=torch.device('cpu'))76 77denoiser = denoiser.to(model_dtype)78denoiser.load_state_dict(state_dict)79denoiser = denoiser.to(device)80 81vae = AutoencoderKL.from_pretrained("madebyollin/sdxl-vae-fp16-fix",82 torch_dtype=model_dtype).to(device)83 84clip_model, preprocess = clip.load("ViT-L/14")85clip_model = clip_model.to(device)86 87diffuser = DiffusionGenerator(denoiser, vae, device, model_dtype)88 89# Define the Gradio interface90iface = gr.Interface(91 fn=generate_image_from_text, # The function to generate the image92 inputs=["text", "slider"],93 outputs="image",94 title="Text-to-Image Generator",95 description="Enter a text prompt to generate an image."96)97 98# Launch the app99iface.launch()