Team Ai
Apppublic

multimodalart/diffusion

sourceHugging Facemitupdated 4y agoView on Hugging Face
10likes
app.py173 linesDownload Raw Back to root
1import gc2import math3import sys4 5#from IPython import display6import torch7from torch import nn8from torch.nn import functional as F9from torchvision import transforms10from torchvision import utils as tv_utils11from torchvision.transforms import functional as TF12import gradio as gr13from git.repo.base import Repo14from os.path import exists as path_exists15 16if not (path_exists(f"v-diffusion-pytorch")):17    Repo.clone_from("https://github.com/crowsonkb/v-diffusion-pytorch", "v-diffusion-pytorch")18if not (path_exists(f"CLIP")):19    Repo.clone_from("https://github.com/openai/CLIP", "CLIP")20sys.path.append('v-diffusion-pytorch')21 22from huggingface_hub import hf_hub_download23 24from CLIP import clip25from diffusion import get_model, sampling, utils26 27class MakeCutouts(nn.Module):28    def __init__(self, cut_size, cutn, cut_pow=1.):29        super().__init__()30        self.cut_size = cut_size31        self.cutn = cutn32        self.cut_pow = cut_pow33 34    def forward(self, input):35        sideY, sideX = input.shape[2:4]36        max_size = min(sideX, sideY)37        min_size = min(sideX, sideY, self.cut_size)38        cutouts = []39        for _ in range(self.cutn):40            size = int(torch.rand([])**self.cut_pow * (max_size - min_size) + min_size)41            offsetx = torch.randint(0, sideX - size + 1, ())42            offsety = torch.randint(0, sideY - size + 1, ())43            cutout = input[:, :, offsety:offsety + size, offsetx:offsetx + size]44            cutout = F.adaptive_avg_pool2d(cutout, self.cut_size)45            cutouts.append(cutout)46        return torch.cat(cutouts)47 48def spherical_dist_loss(x, y):49    x = F.normalize(x, dim=-1)50    y = F.normalize(y, dim=-1)51    return (x - y).norm(dim=-1).div(2).arcsin().pow(2).mul(2)52    53cc12m_model = hf_hub_download(repo_id="multimodalart/crowsonkb-v-diffusion-cc12m-1-cfg", filename="cc12m_1_cfg.pth")54#cc12m_small_model = hf_hub_download(repo_id="multimodalart/crowsonkb-v-diffusion-cc12m-1-cfg", filename="cc12m_1.pth")55model = get_model('cc12m_1_cfg')()56_, side_y, side_x = model.shape57model.load_state_dict(torch.load(cc12m_model, map_location='cpu'))58model = model.half().cuda().eval().requires_grad_(False)59 60#model_small = get_model('cc12m_1')()61#model_small.load_state_dict(torch.load(cc12m_model, map_location='cpu'))62#model_small = model_small.half().cuda().eval().requires_grad_(False)63 64clip_model = clip.load(model.clip_model, jit=False, device='cuda')[0]65clip_model.eval().requires_grad_(False)66normalize = transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073],67                                     std=[0.26862954, 0.26130258, 0.27577711])68make_cutouts = MakeCutouts(clip_model.visual.input_resolution, 16, 1.)69gc.collect()70torch.cuda.empty_cache()71 72def run_all(prompt, steps, n_images, weight, clip_guided):73    gc.collect()74    torch.cuda.empty_cache()75    import random76    seed = int(random.randint(0, 2147483647))77    target_embed = clip_model.encode_text(clip.tokenize(prompt).to('cuda')).float()#.cuda()78    79    if(clip_guided):80        n_images = 181        steps = steps*582        clip_guidance_scale = weight*10083        prompts = [prompt]84        target_embeds, weights = [], []85        def parse_prompt(prompt):86            if prompt.startswith('http://') or prompt.startswith('https://'):87                vals = prompt.rsplit(':', 2)88                vals = [vals[0] + ':' + vals[1], *vals[2:]]89            else:90                vals = prompt.rsplit(':', 1)91            vals = vals + ['', '1'][len(vals):]92            return vals[0], float(vals[1])93        94        for prompt in prompts:95            txt, weight = parse_prompt(prompt)96            target_embeds.append(clip_model.encode_text(clip.tokenize(txt).to('cuda')).float())97            weights.append(weight)98        99        target_embeds = torch.cat(target_embeds)100        weights = torch.tensor(weights, device='cuda')101        if weights.sum().abs() < 1e-3:102            raise RuntimeError('The weights must not sum to 0.')103        weights /= weights.sum().abs()104        clip_embed = F.normalize(target_embeds.mul(weights[:, None]).sum(0, keepdim=True), dim=-1)105        clip_embed = target_embed.repeat([n_images, 1])106    107    def cfg_model_fn(x, t):108        """The CFG wrapper function."""109        n = x.shape[0]110        x_in = x.repeat([2, 1, 1, 1])111        t_in = t.repeat([2])112        clip_embed_repeat = target_embed.repeat([n, 1])113        clip_embed_in = torch.cat([torch.zeros_like(clip_embed_repeat), clip_embed_repeat])114        v_uncond, v_cond = model(x_in, t_in, clip_embed_in).chunk(2, dim=0)115        v = v_uncond + (v_cond - v_uncond) * weight116        return v   117    def make_cond_model_fn(model, cond_fn):118        def cond_model_fn(x, t, **extra_args):119            with torch.enable_grad():120                x = x.detach().requires_grad_()121                v = model(x, t, **extra_args)122                alphas, sigmas = utils.t_to_alpha_sigma(t)123                pred = x * alphas[:, None, None, None] - v * sigmas[:, None, None, None]124                cond_grad = cond_fn(x, t, pred, **extra_args).detach()125                v = v.detach() - cond_grad * (sigmas[:, None, None, None] / alphas[:, None, None, None])126            return v127        return cond_model_fn128    def cond_fn(x, t, pred, clip_embed):129        if min(pred.shape[2:4]) < 256:130            pred = F.interpolate(pred, scale_factor=2, mode='bilinear', align_corners=False)131        clip_in = normalize(make_cutouts((pred + 1) / 2))132        image_embeds = clip_model.encode_image(clip_in).view([16, x.shape[0], -1])133        losses = spherical_dist_loss(image_embeds, clip_embed[None])134        loss = losses.mean(0).sum() * clip_guidance_scale135        grad = -torch.autograd.grad(loss, x)[0]136        return grad137    138    torch.manual_seed(seed)139    x = torch.randn([n_images, 3, side_y, side_x], device='cuda')140    t = torch.linspace(1, 0, steps + 1, device='cuda')[:-1]141    if model.min_t == 0:142        step_list = utils.get_spliced_ddpm_cosine_schedule(t)143    else:144        step_list = utils.get_ddpm_schedule(t)145    if(not clip_guided):146        outs = sampling.plms_sample(cfg_model_fn, x, step_list, {})#, callback=display_callback)147    else:148        extra_args = {'clip_embed': clip_embed}149        cond_fn_ = cond_fn150        model_fn = make_cond_model_fn(model, cond_fn_)151        outs = sampling.plms_sample(model_fn, x, step_list, extra_args)152    images_out = []153    for i, out in enumerate(outs):154        images_out.append(utils.to_pil_image(out))155    return(images_out)156    157 158##################### START GRADIO HERE ############################159gallery = gr.outputs.Carousel(label="Individual images",components=["image"])160iface = gr.Interface(161    fn=run_all, 162    inputs=[163    gr.inputs.Textbox(label="Prompt - try adding increments to your prompt such as 'oil on canvas', 'a painting', 'a book cover'",default="an eerie alien forest"),164    gr.inputs.Slider(label="Steps - more steps can increase quality but will take longer to generate",default=40,maximum=80,minimum=1,step=1),165    gr.inputs.Slider(label="Number of images in parallel", default=2, maximum=4, minimum=1, step=1),166    gr.inputs.Slider(label="Weight - how closely the image should resemble the prompt", default=5, maximum=15, minimum=0, step=1),167    gr.inputs.Checkbox(label="CLIP Guided - improves coherence with complex prompts, makes it slower (with CLIP Guidance only one image is generated)"),168    ], 169    outputs=gallery,170    title="Generate images from text with V-Diffusion",171    description="<div>By typing a prompt and pressing submit you can generate images based on this prompt. <a href='https://github.com/crowsonkb/v-diffusion-pytorch' target='_blank'>V-Diffusion</a> is diffusion text-to-image model created by <a href='https://twitter.com/RiversHaveWings' target='_blank'>Katherine Crowson</a> and <a href='https://twitter.com/jd_pressman'>JDP</a>, trained on the <a href='https://github.com/google-research-datasets/conceptual-12m'>CC12M dataset</a>. The UI to the model was assembled by <a style='color: rgb(99, 102, 241);font-weight:bold' href='https://twitter.com/multimodalart' target='_blank'>@multimodalart</a>, keep up with the <a style='color: rgb(99, 102, 241);' href='https://multimodal.art/news' target='_blank'>latest multimodal ai art news here</a> and consider <a style='color: rgb(99, 102, 241);' href='https://www.patreon.com/multimodalart' target='_blank'>supporting us on Patreon</a></div>",172    )173iface.launch(enable_queue=True)