diff --git a/modelscope/pipelines/multi_modal/diffusers_wrapped/stable_diffusion/chinese_stable_diffusion_pipeline.py b/modelscope/pipelines/multi_modal/diffusers_wrapped/stable_diffusion/chinese_stable_diffusion_pipeline.py index fa8e1f50..d1627962 100644 --- a/modelscope/pipelines/multi_modal/diffusers_wrapped/stable_diffusion/chinese_stable_diffusion_pipeline.py +++ b/modelscope/pipelines/multi_modal/diffusers_wrapped/stable_diffusion/chinese_stable_diffusion_pipeline.py @@ -6,7 +6,7 @@ # and publicly available at # https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py -from typing import Any, Dict, List, Union +from typing import Any, Dict, List, Optional, Union import cv2 import numpy as np @@ -136,8 +136,15 @@ class _DiffuersChineseStableDiffusionPipeline(StableDiffusionPipeline): feature_extractor=feature_extractor, requires_safety_checker=requires_safety_checker) - def _encode_prompt(self, prompt, device, num_images_per_prompt, - do_classifier_free_guidance, negative_prompt): + def _encode_prompt( + self, + prompt, + device, + num_images_per_prompt, + do_classifier_free_guidance, + negative_prompt=None, + prompt_embeds: Optional[torch.FloatTensor] = None, + negative_prompt_embeds: Optional[torch.FloatTensor] = None): r""" Encodes the prompt into text encoder hidden states. @@ -153,27 +160,43 @@ class _DiffuersChineseStableDiffusionPipeline(StableDiffusionPipeline): negative_prompt (`str` or `List[str]`): The prompt or prompts not to guide the image generation. Ignored when not using guidance (i.e., ignored if `guidance_scale` is less than `1`). + prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + negative_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt + weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input + argument. """ - batch_size = len(prompt) if isinstance(prompt, list) else 1 + if prompt is not None and isinstance(prompt, str): + batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] - text_inputs = self.tokenizer( - text=prompt, - padding='max_length', - truncation=True, - max_length=52, - return_tensors='pt') - text_inputs = {k: v.to(device) for k, v in text_inputs.items()} - text_embeddings = self.text_encoder(**text_inputs) - text_embeddings = text_embeddings[0] + if prompt_embeds is None: + text_inputs = self.tokenizer( + text=prompt, + padding='max_length', + truncation=True, + max_length=52, + return_tensors='pt') + text_inputs = {k: v.to(device) for k, v in text_inputs.items()} + prompt_embeds = self.text_encoder(**text_inputs) + prompt_embeds = prompt_embeds[0] + + prompt_embeds = prompt_embeds.to( + dtype=self.text_encoder.dtype, device=device) # duplicate text embeddings for each generation per prompt, using mps friendly method - bs_embed, seq_len, _ = text_embeddings.shape - text_embeddings = text_embeddings.repeat(1, num_images_per_prompt, 1) - text_embeddings = text_embeddings.view( - bs_embed * num_images_per_prompt, seq_len, -1) + bs_embed, seq_len, _ = prompt_embeds.shape + prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1) + prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, + seq_len, -1) # get unconditional embeddings for classifier free guidance - if do_classifier_free_guidance: + if do_classifier_free_guidance and negative_prompt_embeds is None: uncond_tokens: List[str] if negative_prompt is None: uncond_tokens = [''] * batch_size @@ -198,19 +221,24 @@ class _DiffuersChineseStableDiffusionPipeline(StableDiffusionPipeline): max_length=52, return_tensors='pt') uncond_input = {k: v.to(device) for k, v in uncond_input.items()} - uncond_embeddings = self.text_encoder(**uncond_input) - uncond_embeddings = uncond_embeddings[0] + negative_prompt_embeds = self.text_encoder(**uncond_input) + negative_prompt_embeds = negative_prompt_embeds[0] + if do_classifier_free_guidance: # duplicate unconditional embeddings for each generation per prompt, using mps friendly method - seq_len = uncond_embeddings.shape[1] - uncond_embeddings = uncond_embeddings.repeat( + seq_len = negative_prompt_embeds.shape[1] + + negative_prompt_embeds = negative_prompt_embeds.to( + dtype=self.text_encoder.dtype, device=device) + + negative_prompt_embeds = negative_prompt_embeds.repeat( 1, num_images_per_prompt, 1) - uncond_embeddings = uncond_embeddings.view( + negative_prompt_embeds = negative_prompt_embeds.view( batch_size * num_images_per_prompt, seq_len, -1) # For classifier free guidance, we need to do two forward passes. # Here we concatenate the unconditional and text embeddings into a single batch # to avoid doing two forward passes - text_embeddings = torch.cat([uncond_embeddings, text_embeddings]) + prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds]) - return text_embeddings + return prompt_embeds