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 29d agoView on Hugging Face
9likes22kdownloads
interpolate_stable_diffusion.py526 linesDownload Raw Back to v0.23.0
1import inspect2import time3from pathlib import Path4from typing import Callable, List, Optional, Union5 6import numpy as np7import torch8from transformers import CLIPImageProcessor, CLIPTextModel, CLIPTokenizer9 10from diffusers import DiffusionPipeline11from diffusers.configuration_utils import FrozenDict12from diffusers.models import AutoencoderKL, UNet2DConditionModel13from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput14from diffusers.pipelines.stable_diffusion.safety_checker import StableDiffusionSafetyChecker15from diffusers.schedulers import DDIMScheduler, LMSDiscreteScheduler, PNDMScheduler16from diffusers.utils import deprecate, logging17 18 19logger = logging.get_logger(__name__)  # pylint: disable=invalid-name20 21 22def slerp(t, v0, v1, DOT_THRESHOLD=0.9995):23    """helper function to spherically interpolate two arrays v1 v2"""24 25    if not isinstance(v0, np.ndarray):26        inputs_are_torch = True27        input_device = v0.device28        v0 = v0.cpu().numpy()29        v1 = v1.cpu().numpy()30 31    dot = np.sum(v0 * v1 / (np.linalg.norm(v0) * np.linalg.norm(v1)))32    if np.abs(dot) > DOT_THRESHOLD:33        v2 = (1 - t) * v0 + t * v134    else:35        theta_0 = np.arccos(dot)36        sin_theta_0 = np.sin(theta_0)37        theta_t = theta_0 * t38        sin_theta_t = np.sin(theta_t)39        s0 = np.sin(theta_0 - theta_t) / sin_theta_040        s1 = sin_theta_t / sin_theta_041        v2 = s0 * v0 + s1 * v142 43    if inputs_are_torch:44        v2 = torch.from_numpy(v2).to(input_device)45 46    return v247 48 49class StableDiffusionWalkPipeline(DiffusionPipeline):50    r"""51    Pipeline for text-to-image generation using Stable Diffusion.52 53    This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods the54    library implements for all the pipelines (such as downloading or saving, running on a particular device, etc.)55 56    Args:57        vae ([`AutoencoderKL`]):58            Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.59        text_encoder ([`CLIPTextModel`]):60            Frozen text-encoder. Stable Diffusion uses the text portion of61            [CLIP](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel), specifically62            the [clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14) variant.63        tokenizer (`CLIPTokenizer`):64            Tokenizer of class65            [CLIPTokenizer](https://huggingface.co/docs/transformers/v4.21.0/en/model_doc/clip#transformers.CLIPTokenizer).66        unet ([`UNet2DConditionModel`]): Conditional U-Net architecture to denoise the encoded image latents.67        scheduler ([`SchedulerMixin`]):68            A scheduler to be used in combination with `unet` to denoise the encoded image latents. Can be one of69            [`DDIMScheduler`], [`LMSDiscreteScheduler`], or [`PNDMScheduler`].70        safety_checker ([`StableDiffusionSafetyChecker`]):71            Classification module that estimates whether generated images could be considered offensive or harmful.72            Please, refer to the [model card](https://huggingface.co/CompVis/stable-diffusion-v1-4) for details.73        feature_extractor ([`CLIPImageProcessor`]):74            Model that extracts features from generated images to be used as inputs for the `safety_checker`.75    """76 77    def __init__(78        self,79        vae: AutoencoderKL,80        text_encoder: CLIPTextModel,81        tokenizer: CLIPTokenizer,82        unet: UNet2DConditionModel,83        scheduler: Union[DDIMScheduler, PNDMScheduler, LMSDiscreteScheduler],84        safety_checker: StableDiffusionSafetyChecker,85        feature_extractor: CLIPImageProcessor,86    ):87        super().__init__()88 89        if hasattr(scheduler.config, "steps_offset") and scheduler.config.steps_offset != 1:90            deprecation_message = (91                f"The configuration file of this scheduler: {scheduler} is outdated. `steps_offset`"92                f" should be set to 1 instead of {scheduler.config.steps_offset}. Please make sure "93                "to update the config accordingly as leaving `steps_offset` might led to incorrect results"94                " in future versions. If you have downloaded this checkpoint from the Hugging Face Hub,"95                " it would be very nice if you could open a Pull request for the `scheduler/scheduler_config.json`"96                " file"97            )98            deprecate("steps_offset!=1", "1.0.0", deprecation_message, standard_warn=False)99            new_config = dict(scheduler.config)100            new_config["steps_offset"] = 1101            scheduler._internal_dict = FrozenDict(new_config)102 103        if safety_checker is None:104            logger.warning(105                f"You have disabled the safety checker for {self.__class__} by passing `safety_checker=None`. Ensure"106                " that you abide to the conditions of the Stable Diffusion license and do not expose unfiltered"107                " results in services or applications open to the public. Both the diffusers team and Hugging Face"108                " strongly recommend to keep the safety filter enabled in all public facing circumstances, disabling"109                " it only for use-cases that involve analyzing network behavior or auditing its results. For more"110                " information, please have a look at https://github.com/huggingface/diffusers/pull/254 ."111            )112 113        self.register_modules(114            vae=vae,115            text_encoder=text_encoder,116            tokenizer=tokenizer,117            unet=unet,118            scheduler=scheduler,119            safety_checker=safety_checker,120            feature_extractor=feature_extractor,121        )122 123    def enable_attention_slicing(self, slice_size: Optional[Union[str, int]] = "auto"):124        r"""125        Enable sliced attention computation.126 127        When this option is enabled, the attention module will split the input tensor in slices, to compute attention128        in several steps. This is useful to save some memory in exchange for a small speed decrease.129 130        Args:131            slice_size (`str` or `int`, *optional*, defaults to `"auto"`):132                When `"auto"`, halves the input to the attention heads, so attention will be computed in two steps. If133                a number is provided, uses as many slices as `attention_head_dim // slice_size`. In this case,134                `attention_head_dim` must be a multiple of `slice_size`.135        """136        if slice_size == "auto":137            # half the attention head size is usually a good trade-off between138            # speed and memory139            slice_size = self.unet.config.attention_head_dim // 2140        self.unet.set_attention_slice(slice_size)141 142    def disable_attention_slicing(self):143        r"""144        Disable sliced attention computation. If `enable_attention_slicing` was previously invoked, this method will go145        back to computing attention in one step.146        """147        # set slice_size = `None` to disable `attention slicing`148        self.enable_attention_slicing(None)149 150    @torch.no_grad()151    def __call__(152        self,153        prompt: Optional[Union[str, List[str]]] = None,154        height: int = 512,155        width: int = 512,156        num_inference_steps: int = 50,157        guidance_scale: float = 7.5,158        negative_prompt: Optional[Union[str, List[str]]] = None,159        num_images_per_prompt: Optional[int] = 1,160        eta: float = 0.0,161        generator: Optional[torch.Generator] = None,162        latents: Optional[torch.FloatTensor] = None,163        output_type: Optional[str] = "pil",164        return_dict: bool = True,165        callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,166        callback_steps: int = 1,167        text_embeddings: Optional[torch.FloatTensor] = None,168        **kwargs,169    ):170        r"""171        Function invoked when calling the pipeline for generation.172 173        Args:174            prompt (`str` or `List[str]`, *optional*, defaults to `None`):175                The prompt or prompts to guide the image generation. If not provided, `text_embeddings` is required.176            height (`int`, *optional*, defaults to 512):177                The height in pixels of the generated image.178            width (`int`, *optional*, defaults to 512):179                The width in pixels of the generated image.180            num_inference_steps (`int`, *optional*, defaults to 50):181                The number of denoising steps. More denoising steps usually lead to a higher quality image at the182                expense of slower inference.183            guidance_scale (`float`, *optional*, defaults to 7.5):184                Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).185                `guidance_scale` is defined as `w` of equation 2. of [Imagen186                Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >187                1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,188                usually at the expense of lower image quality.189            negative_prompt (`str` or `List[str]`, *optional*):190                The prompt or prompts not to guide the image generation. Ignored when not using guidance (i.e., ignored191                if `guidance_scale` is less than `1`).192            num_images_per_prompt (`int`, *optional*, defaults to 1):193                The number of images to generate per prompt.194            eta (`float`, *optional*, defaults to 0.0):195                Corresponds to parameter eta (η) in the DDIM paper: https://arxiv.org/abs/2010.02502. Only applies to196                [`schedulers.DDIMScheduler`], will be ignored for others.197            generator (`torch.Generator`, *optional*):198                A [torch generator](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make generation199                deterministic.200            latents (`torch.FloatTensor`, *optional*):201                Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image202                generation. Can be used to tweak the same generation with different prompts. If not provided, a latents203                tensor will ge generated by sampling using the supplied random `generator`.204            output_type (`str`, *optional*, defaults to `"pil"`):205                The output format of the generate image. Choose between206                [PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.207            return_dict (`bool`, *optional*, defaults to `True`):208                Whether or not to return a [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] instead of a209                plain tuple.210            callback (`Callable`, *optional*):211                A function that will be called every `callback_steps` steps during inference. The function will be212                called with the following arguments: `callback(step: int, timestep: int, latents: torch.FloatTensor)`.213            callback_steps (`int`, *optional*, defaults to 1):214                The frequency at which the `callback` function will be called. If not specified, the callback will be215                called at every step.216            text_embeddings (`torch.FloatTensor`, *optional*, defaults to `None`):217                Pre-generated text embeddings to be used as inputs for image generation. Can be used in place of218                `prompt` to avoid re-computing the embeddings. If not provided, the embeddings will be generated from219                the supplied `prompt`.220 221        Returns:222            [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] or `tuple`:223            [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] if `return_dict` is True, otherwise a `tuple.224            When returning a tuple, the first element is a list with the generated images, and the second element is a225            list of `bool`s denoting whether the corresponding generated image likely represents "not-safe-for-work"226            (nsfw) content, according to the `safety_checker`.227        """228 229        if height % 8 != 0 or width % 8 != 0:230            raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")231 232        if (callback_steps is None) or (233            callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0)234        ):235            raise ValueError(236                f"`callback_steps` has to be a positive integer but is {callback_steps} of type"237                f" {type(callback_steps)}."238            )239 240        if text_embeddings is None:241            if isinstance(prompt, str):242                batch_size = 1243            elif isinstance(prompt, list):244                batch_size = len(prompt)245            else:246                raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")247 248            # get prompt text embeddings249            text_inputs = self.tokenizer(250                prompt,251                padding="max_length",252                max_length=self.tokenizer.model_max_length,253                return_tensors="pt",254            )255            text_input_ids = text_inputs.input_ids256 257            if text_input_ids.shape[-1] > self.tokenizer.model_max_length:258                removed_text = self.tokenizer.batch_decode(text_input_ids[:, self.tokenizer.model_max_length :])259                print(260                    "The following part of your input was truncated because CLIP can only handle sequences up to"261                    f" {self.tokenizer.model_max_length} tokens: {removed_text}"262                )263                text_input_ids = text_input_ids[:, : self.tokenizer.model_max_length]264            text_embeddings = self.text_encoder(text_input_ids.to(self.device))[0]265        else:266            batch_size = text_embeddings.shape[0]267 268        # duplicate text embeddings for each generation per prompt, using mps friendly method269        bs_embed, seq_len, _ = text_embeddings.shape270        text_embeddings = text_embeddings.repeat(1, num_images_per_prompt, 1)271        text_embeddings = text_embeddings.view(bs_embed * num_images_per_prompt, seq_len, -1)272 273        # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)274        # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`275        # corresponds to doing no classifier free guidance.276        do_classifier_free_guidance = guidance_scale > 1.0277        # get unconditional embeddings for classifier free guidance278        if do_classifier_free_guidance:279            uncond_tokens: List[str]280            if negative_prompt is None:281                uncond_tokens = [""] * batch_size282            elif type(prompt) is not type(negative_prompt):283                raise TypeError(284                    f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="285                    f" {type(prompt)}."286                )287            elif isinstance(negative_prompt, str):288                uncond_tokens = [negative_prompt]289            elif batch_size != len(negative_prompt):290                raise ValueError(291                    f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"292                    f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"293                    " the batch size of `prompt`."294                )295            else:296                uncond_tokens = negative_prompt297 298            max_length = self.tokenizer.model_max_length299            uncond_input = self.tokenizer(300                uncond_tokens,301                padding="max_length",302                max_length=max_length,303                truncation=True,304                return_tensors="pt",305            )306            uncond_embeddings = self.text_encoder(uncond_input.input_ids.to(self.device))[0]307 308            # duplicate unconditional embeddings for each generation per prompt, using mps friendly method309            seq_len = uncond_embeddings.shape[1]310            uncond_embeddings = uncond_embeddings.repeat(1, num_images_per_prompt, 1)311            uncond_embeddings = uncond_embeddings.view(batch_size * num_images_per_prompt, seq_len, -1)312 313            # For classifier free guidance, we need to do two forward passes.314            # Here we concatenate the unconditional and text embeddings into a single batch315            # to avoid doing two forward passes316            text_embeddings = torch.cat([uncond_embeddings, text_embeddings])317 318        # get the initial random noise unless the user supplied it319 320        # Unlike in other pipelines, latents need to be generated in the target device321        # for 1-to-1 results reproducibility with the CompVis implementation.322        # However this currently doesn't work in `mps`.323        latents_shape = (batch_size * num_images_per_prompt, self.unet.config.in_channels, height // 8, width // 8)324        latents_dtype = text_embeddings.dtype325        if latents is None:326            if self.device.type == "mps":327                # randn does not work reproducibly on mps328                latents = torch.randn(latents_shape, generator=generator, device="cpu", dtype=latents_dtype).to(329                    self.device330                )331            else:332                latents = torch.randn(latents_shape, generator=generator, device=self.device, dtype=latents_dtype)333        else:334            if latents.shape != latents_shape:335                raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {latents_shape}")336            latents = latents.to(self.device)337 338        # set timesteps339        self.scheduler.set_timesteps(num_inference_steps)340 341        # Some schedulers like PNDM have timesteps as arrays342        # It's more optimized to move all timesteps to correct device beforehand343        timesteps_tensor = self.scheduler.timesteps.to(self.device)344 345        # scale the initial noise by the standard deviation required by the scheduler346        latents = latents * self.scheduler.init_noise_sigma347 348        # prepare extra kwargs for the scheduler step, since not all schedulers have the same signature349        # eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.350        # eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502351        # and should be between [0, 1]352        accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())353        extra_step_kwargs = {}354        if accepts_eta:355            extra_step_kwargs["eta"] = eta356 357        for i, t in enumerate(self.progress_bar(timesteps_tensor)):358            # expand the latents if we are doing classifier free guidance359            latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents360            latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)361 362            # predict the noise residual363            noise_pred = self.unet(latent_model_input, t, encoder_hidden_states=text_embeddings).sample364 365            # perform guidance366            if do_classifier_free_guidance:367                noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)368                noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)369 370            # compute the previous noisy sample x_t -> x_t-1371            latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs).prev_sample372 373            # call the callback, if provided374            if callback is not None and i % callback_steps == 0:375                step_idx = i // getattr(self.scheduler, "order", 1)376                callback(step_idx, t, latents)377 378        latents = 1 / 0.18215 * latents379        image = self.vae.decode(latents).sample380 381        image = (image / 2 + 0.5).clamp(0, 1)382 383        # we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16384        image = image.cpu().permute(0, 2, 3, 1).float().numpy()385 386        if self.safety_checker is not None:387            safety_checker_input = self.feature_extractor(self.numpy_to_pil(image), return_tensors="pt").to(388                self.device389            )390            image, has_nsfw_concept = self.safety_checker(391                images=image, clip_input=safety_checker_input.pixel_values.to(text_embeddings.dtype)392            )393        else:394            has_nsfw_concept = None395 396        if output_type == "pil":397            image = self.numpy_to_pil(image)398 399        if not return_dict:400            return (image, has_nsfw_concept)401 402        return StableDiffusionPipelineOutput(images=image, nsfw_content_detected=has_nsfw_concept)403 404    def embed_text(self, text):405        """takes in text and turns it into text embeddings"""406        text_input = self.tokenizer(407            text,408            padding="max_length",409            max_length=self.tokenizer.model_max_length,410            truncation=True,411            return_tensors="pt",412        )413        with torch.no_grad():414            embed = self.text_encoder(text_input.input_ids.to(self.device))[0]415        return embed416 417    def get_noise(self, seed, dtype=torch.float32, height=512, width=512):418        """Takes in random seed and returns corresponding noise vector"""419        return torch.randn(420            (1, self.unet.config.in_channels, height // 8, width // 8),421            generator=torch.Generator(device=self.device).manual_seed(seed),422            device=self.device,423            dtype=dtype,424        )425 426    def walk(427        self,428        prompts: List[str],429        seeds: List[int],430        num_interpolation_steps: Optional[int] = 6,431        output_dir: Optional[str] = "./dreams",432        name: Optional[str] = None,433        batch_size: Optional[int] = 1,434        height: Optional[int] = 512,435        width: Optional[int] = 512,436        guidance_scale: Optional[float] = 7.5,437        num_inference_steps: Optional[int] = 50,438        eta: Optional[float] = 0.0,439    ) -> List[str]:440        """441        Walks through a series of prompts and seeds, interpolating between them and saving the results to disk.442 443        Args:444            prompts (`List[str]`):445                List of prompts to generate images for.446            seeds (`List[int]`):447                List of seeds corresponding to provided prompts. Must be the same length as prompts.448            num_interpolation_steps (`int`, *optional*, defaults to 6):449                Number of interpolation steps to take between prompts.450            output_dir (`str`, *optional*, defaults to `./dreams`):451                Directory to save the generated images to.452            name (`str`, *optional*, defaults to `None`):453                Subdirectory of `output_dir` to save the generated images to. If `None`, the name will454                be the current time.455            batch_size (`int`, *optional*, defaults to 1):456                Number of images to generate at once.457            height (`int`, *optional*, defaults to 512):458                Height of the generated images.459            width (`int`, *optional*, defaults to 512):460                Width of the generated images.461            guidance_scale (`float`, *optional*, defaults to 7.5):462                Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).463                `guidance_scale` is defined as `w` of equation 2. of [Imagen464                Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >465                1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,466                usually at the expense of lower image quality.467            num_inference_steps (`int`, *optional*, defaults to 50):468                The number of denoising steps. More denoising steps usually lead to a higher quality image at the469                expense of slower inference.470            eta (`float`, *optional*, defaults to 0.0):471                Corresponds to parameter eta (η) in the DDIM paper: https://arxiv.org/abs/2010.02502. Only applies to472                [`schedulers.DDIMScheduler`], will be ignored for others.473 474        Returns:475            `List[str]`: List of paths to the generated images.476        """477        if not len(prompts) == len(seeds):478            raise ValueError(479                f"Number of prompts and seeds must be equalGot {len(prompts)} prompts and {len(seeds)} seeds"480            )481 482        name = name or time.strftime("%Y%m%d-%H%M%S")483        save_path = Path(output_dir) / name484        save_path.mkdir(exist_ok=True, parents=True)485 486        frame_idx = 0487        frame_filepaths = []488        for prompt_a, prompt_b, seed_a, seed_b in zip(prompts, prompts[1:], seeds, seeds[1:]):489            # Embed Text490            embed_a = self.embed_text(prompt_a)491            embed_b = self.embed_text(prompt_b)492 493            # Get Noise494            noise_dtype = embed_a.dtype495            noise_a = self.get_noise(seed_a, noise_dtype, height, width)496            noise_b = self.get_noise(seed_b, noise_dtype, height, width)497 498            noise_batch, embeds_batch = None, None499            T = np.linspace(0.0, 1.0, num_interpolation_steps)500            for i, t in enumerate(T):501                noise = slerp(float(t), noise_a, noise_b)502                embed = torch.lerp(embed_a, embed_b, t)503 504                noise_batch = noise if noise_batch is None else torch.cat([noise_batch, noise], dim=0)505                embeds_batch = embed if embeds_batch is None else torch.cat([embeds_batch, embed], dim=0)506 507                batch_is_ready = embeds_batch.shape[0] == batch_size or i + 1 == T.shape[0]508                if batch_is_ready:509                    outputs = self(510                        latents=noise_batch,511                        text_embeddings=embeds_batch,512                        height=height,513                        width=width,514                        guidance_scale=guidance_scale,515                        eta=eta,516                        num_inference_steps=num_inference_steps,517                    )518                    noise_batch, embeds_batch = None, None519 520                    for image in outputs["images"]:521                        frame_filepath = str(save_path / f"frame_{frame_idx:06d}.png")522                        image.save(frame_filepath)523                        frame_filepaths.append(frame_filepath)524                        frame_idx += 1525        return frame_filepaths526