CoolFace
Datasetpublic

diffusers/community-pipelines-mirror

Community Pipeline Examples For more information about community pipelines, please have a look at this issue. Community pipeline examples consist pipelines that have been added by the community. Please have a look at the following tables to get an overview of all community examples. Click on the Code Example to get a copy-and-paste ready code example that you can try out. If a community pipeline doesn't work as expected, please open an issue and ping the author on it. Please… See the full description on the dataset page: https://huggingface.co/datasets/diffusers/community-pipelines-mirror.

sourceHugging Faceupdated 27d agoView on Hugging Face
9likes22kdownloads
clip_guided_stable_diffusion.py338 linesDownload Raw Back to root
1import inspect2from typing import List, Optional, Union3 4import torch5from torch import nn6from torch.nn import functional as F7from torchvision import transforms8from transformers import CLIPImageProcessor, CLIPModel, CLIPTextModel, CLIPTokenizer9 10from diffusers import (11    AutoencoderKL,12    DDIMScheduler,13    DPMSolverMultistepScheduler,14    LMSDiscreteScheduler,15    PNDMScheduler,16    UNet2DConditionModel,17)18from diffusers.pipelines.pipeline_utils import DiffusionPipeline, StableDiffusionMixin19from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion import StableDiffusionPipelineOutput20 21 22class MakeCutouts(nn.Module):23    def __init__(self, cut_size, cut_power=1.0):24        super().__init__()25 26        self.cut_size = cut_size27        self.cut_power = cut_power28 29    def forward(self, pixel_values, num_cutouts):30        sideY, sideX = pixel_values.shape[2:4]31        max_size = min(sideX, sideY)32        min_size = min(sideX, sideY, self.cut_size)33        cutouts = []34        for _ in range(num_cutouts):35            size = int(torch.rand([]) ** self.cut_power * (max_size - min_size) + min_size)36            offsetx = torch.randint(0, sideX - size + 1, ())37            offsety = torch.randint(0, sideY - size + 1, ())38            cutout = pixel_values[:, :, offsety : offsety + size, offsetx : offsetx + size]39            cutouts.append(F.adaptive_avg_pool2d(cutout, self.cut_size))40        return torch.cat(cutouts)41 42 43def spherical_dist_loss(x, y):44    x = F.normalize(x, dim=-1)45    y = F.normalize(y, dim=-1)46    return (x - y).norm(dim=-1).div(2).arcsin().pow(2).mul(2)47 48 49def set_requires_grad(model, value):50    for param in model.parameters():51        param.requires_grad = value52 53 54class CLIPGuidedStableDiffusion(DiffusionPipeline, StableDiffusionMixin):55    """CLIP guided stable diffusion based on the amazing repo by @crowsonkb and @Jack00056    - https://github.com/Jack000/glid-3-xl57    - https://github.dev/crowsonkb/k-diffusion58    """59 60    def __init__(61        self,62        vae: AutoencoderKL,63        text_encoder: CLIPTextModel,64        clip_model: CLIPModel,65        tokenizer: CLIPTokenizer,66        unet: UNet2DConditionModel,67        scheduler: Union[PNDMScheduler, LMSDiscreteScheduler, DDIMScheduler, DPMSolverMultistepScheduler],68        feature_extractor: CLIPImageProcessor,69    ):70        super().__init__()71        self.register_modules(72            vae=vae,73            text_encoder=text_encoder,74            clip_model=clip_model,75            tokenizer=tokenizer,76            unet=unet,77            scheduler=scheduler,78            feature_extractor=feature_extractor,79        )80 81        self.normalize = transforms.Normalize(mean=feature_extractor.image_mean, std=feature_extractor.image_std)82        self.cut_out_size = (83            feature_extractor.size84            if isinstance(feature_extractor.size, int)85            else feature_extractor.size["shortest_edge"]86        )87        self.make_cutouts = MakeCutouts(self.cut_out_size)88 89        set_requires_grad(self.text_encoder, False)90        set_requires_grad(self.clip_model, False)91 92    def freeze_vae(self):93        set_requires_grad(self.vae, False)94 95    def unfreeze_vae(self):96        set_requires_grad(self.vae, True)97 98    def freeze_unet(self):99        set_requires_grad(self.unet, False)100 101    def unfreeze_unet(self):102        set_requires_grad(self.unet, True)103 104    @torch.enable_grad()105    def cond_fn(106        self,107        latents,108        timestep,109        index,110        text_embeddings,111        noise_pred_original,112        text_embeddings_clip,113        clip_guidance_scale,114        num_cutouts,115        use_cutouts=True,116    ):117        latents = latents.detach().requires_grad_()118 119        latent_model_input = self.scheduler.scale_model_input(latents, timestep)120 121        # predict the noise residual122        noise_pred = self.unet(latent_model_input, timestep, encoder_hidden_states=text_embeddings).sample123 124        if isinstance(self.scheduler, (PNDMScheduler, DDIMScheduler, DPMSolverMultistepScheduler)):125            alpha_prod_t = self.scheduler.alphas_cumprod[timestep]126            beta_prod_t = 1 - alpha_prod_t127            # compute predicted original sample from predicted noise also called128            # "predicted x_0" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf129            pred_original_sample = (latents - beta_prod_t ** (0.5) * noise_pred) / alpha_prod_t ** (0.5)130 131            fac = torch.sqrt(beta_prod_t)132            sample = pred_original_sample * (fac) + latents * (1 - fac)133        elif isinstance(self.scheduler, LMSDiscreteScheduler):134            sigma = self.scheduler.sigmas[index]135            sample = latents - sigma * noise_pred136        else:137            raise ValueError(f"scheduler type {type(self.scheduler)} not supported")138 139        sample = 1 / self.vae.config.scaling_factor * sample140        image = self.vae.decode(sample).sample141        image = (image / 2 + 0.5).clamp(0, 1)142 143        if use_cutouts:144            image = self.make_cutouts(image, num_cutouts)145        else:146            image = transforms.Resize(self.cut_out_size)(image)147        image = self.normalize(image).to(latents.dtype)148 149        image_embeddings_clip = self.clip_model.get_image_features(image)150        image_embeddings_clip = image_embeddings_clip / image_embeddings_clip.norm(p=2, dim=-1, keepdim=True)151 152        if use_cutouts:153            dists = spherical_dist_loss(image_embeddings_clip, text_embeddings_clip)154            dists = dists.view([num_cutouts, sample.shape[0], -1])155            loss = dists.sum(2).mean(0).sum() * clip_guidance_scale156        else:157            loss = spherical_dist_loss(image_embeddings_clip, text_embeddings_clip).mean() * clip_guidance_scale158 159        grads = -torch.autograd.grad(loss, latents)[0]160 161        if isinstance(self.scheduler, LMSDiscreteScheduler):162            latents = latents.detach() + grads * (sigma**2)163            noise_pred = noise_pred_original164        else:165            noise_pred = noise_pred_original - torch.sqrt(beta_prod_t) * grads166        return noise_pred, latents167 168    @torch.no_grad()169    def __call__(170        self,171        prompt: Union[str, List[str]],172        height: Optional[int] = 512,173        width: Optional[int] = 512,174        num_inference_steps: Optional[int] = 50,175        guidance_scale: Optional[float] = 7.5,176        num_images_per_prompt: Optional[int] = 1,177        eta: float = 0.0,178        clip_guidance_scale: Optional[float] = 100,179        clip_prompt: Optional[Union[str, List[str]]] = None,180        num_cutouts: Optional[int] = 4,181        use_cutouts: Optional[bool] = True,182        generator: Optional[torch.Generator] = None,183        latents: Optional[torch.Tensor] = None,184        output_type: Optional[str] = "pil",185        return_dict: bool = True,186    ):187        if isinstance(prompt, str):188            batch_size = 1189        elif isinstance(prompt, list):190            batch_size = len(prompt)191        else:192            raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")193 194        if height % 8 != 0 or width % 8 != 0:195            raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")196 197        # get prompt text embeddings198        text_input = self.tokenizer(199            prompt,200            padding="max_length",201            max_length=self.tokenizer.model_max_length,202            truncation=True,203            return_tensors="pt",204        )205        text_embeddings = self.text_encoder(text_input.input_ids.to(self.device))[0]206        # duplicate text embeddings for each generation per prompt207        text_embeddings = text_embeddings.repeat_interleave(num_images_per_prompt, dim=0)208 209        if clip_guidance_scale > 0:210            if clip_prompt is not None:211                clip_text_input = self.tokenizer(212                    clip_prompt,213                    padding="max_length",214                    max_length=self.tokenizer.model_max_length,215                    truncation=True,216                    return_tensors="pt",217                ).input_ids.to(self.device)218            else:219                clip_text_input = text_input.input_ids.to(self.device)220            text_embeddings_clip = self.clip_model.get_text_features(clip_text_input)221            text_embeddings_clip = text_embeddings_clip / text_embeddings_clip.norm(p=2, dim=-1, keepdim=True)222            # duplicate text embeddings clip for each generation per prompt223            text_embeddings_clip = text_embeddings_clip.repeat_interleave(num_images_per_prompt, dim=0)224 225        # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)226        # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`227        # corresponds to doing no classifier free guidance.228        do_classifier_free_guidance = guidance_scale > 1.0229        # get unconditional embeddings for classifier free guidance230        if do_classifier_free_guidance:231            max_length = text_input.input_ids.shape[-1]232            uncond_input = self.tokenizer([""], padding="max_length", max_length=max_length, return_tensors="pt")233            uncond_embeddings = self.text_encoder(uncond_input.input_ids.to(self.device))[0]234            # duplicate unconditional embeddings for each generation per prompt235            uncond_embeddings = uncond_embeddings.repeat_interleave(num_images_per_prompt, dim=0)236 237            # For classifier free guidance, we need to do two forward passes.238            # Here we concatenate the unconditional and text embeddings into a single batch239            # to avoid doing two forward passes240            text_embeddings = torch.cat([uncond_embeddings, text_embeddings])241 242        # get the initial random noise unless the user supplied it243 244        # Unlike in other pipelines, latents need to be generated in the target device245        # for 1-to-1 results reproducibility with the CompVis implementation.246        # However this currently doesn't work in `mps`.247        latents_shape = (batch_size * num_images_per_prompt, self.unet.config.in_channels, height // 8, width // 8)248        latents_dtype = text_embeddings.dtype249        if latents is None:250            if self.device.type == "mps":251                # randn does not work reproducibly on mps252                latents = torch.randn(latents_shape, generator=generator, device="cpu", dtype=latents_dtype).to(253                    self.device254                )255            else:256                latents = torch.randn(latents_shape, generator=generator, device=self.device, dtype=latents_dtype)257        else:258            if latents.shape != latents_shape:259                raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {latents_shape}")260            latents = latents.to(self.device)261 262        # set timesteps263        accepts_offset = "offset" in set(inspect.signature(self.scheduler.set_timesteps).parameters.keys())264        extra_set_kwargs = {}265        if accepts_offset:266            extra_set_kwargs["offset"] = 1267 268        self.scheduler.set_timesteps(num_inference_steps, **extra_set_kwargs)269 270        # Some schedulers like PNDM have timesteps as arrays271        # It's more optimized to move all timesteps to correct device beforehand272        timesteps_tensor = self.scheduler.timesteps.to(self.device)273 274        # scale the initial noise by the standard deviation required by the scheduler275        latents = latents * self.scheduler.init_noise_sigma276 277        # prepare extra kwargs for the scheduler step, since not all schedulers have the same signature278        # eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.279        # eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502280        # and should be between [0, 1]281        accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())282        extra_step_kwargs = {}283        if accepts_eta:284            extra_step_kwargs["eta"] = eta285 286        # check if the scheduler accepts generator287        accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())288        if accepts_generator:289            extra_step_kwargs["generator"] = generator290 291        for i, t in enumerate(self.progress_bar(timesteps_tensor)):292            # expand the latents if we are doing classifier free guidance293            latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents294            latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)295 296            # predict the noise residual297            noise_pred = self.unet(latent_model_input, t, encoder_hidden_states=text_embeddings).sample298 299            # perform classifier free guidance300            if do_classifier_free_guidance:301                noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)302                noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)303 304            # perform clip guidance305            if clip_guidance_scale > 0:306                text_embeddings_for_guidance = (307                    text_embeddings.chunk(2)[1] if do_classifier_free_guidance else text_embeddings308                )309                noise_pred, latents = self.cond_fn(310                    latents,311                    t,312                    i,313                    text_embeddings_for_guidance,314                    noise_pred,315                    text_embeddings_clip,316                    clip_guidance_scale,317                    num_cutouts,318                    use_cutouts,319                )320 321            # compute the previous noisy sample x_t -> x_t-1322            latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs).prev_sample323 324        # scale and decode the image latents with vae325        latents = 1 / self.vae.config.scaling_factor * latents326        image = self.vae.decode(latents).sample327 328        image = (image / 2 + 0.5).clamp(0, 1)329        image = image.cpu().permute(0, 2, 3, 1).numpy()330 331        if output_type == "pil":332            image = self.numpy_to_pil(image)333 334        if not return_dict:335            return (image, None)336 337        return StableDiffusionPipelineOutput(images=image, nsfw_content_detected=None)338