Compatible with diffusers0.12.1

Link: https://code.alibaba-inc.com/Ali-MaaS/MaaS-lib/codereview/11479564
This commit is contained in:
ly103369
2023-01-29 14:47:48 +00:00
committed by yingda.chen
parent 5a01eca834
commit 14afbfd39a

View File

@@ -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