diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 81391c33e..6ecc973b0 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -217,8 +217,7 @@ def process_diffusers(p: StableDiffusionProcessing): parser = 'Fixed attention' if shared.opts.prompt_attention != 'Fixed attention' and 'StableDiffusion' in model.__class__.__name__: try: - prompt_parser_diffusers.encode_prompts(model, p, prompts, negative_prompts, kwargs.get("num_inference_steps", 1), 0, kwargs.pop("clip_skip", None)) - # prompt_embed, pooled, negative_embed, negative_pooled = , , , , + prompt_parser_diffusers.encode_prompts(model, p, prompts, negative_prompts, kwargs.get("num_inference_steps", 1), kwargs.pop("clip_skip", None)) parser = shared.opts.prompt_attention except Exception as e: shared.log.error(f'Prompt parser encode: {e}') diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 816a476ae..4aae8f5d2 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -4,10 +4,12 @@ import typing import torch from compel import ReturnedEmbeddingsType from compel.embeddings_provider import BaseTextualInversionManager, EmbeddingsProvider +from transformers import PreTrainedTokenizer from modules import shared, prompt_parser, devices debug = shared.log.trace if os.environ.get('SD_PROMPT_DEBUG', None) is not None else lambda *args, **kwargs: None debug('Trace: PROMPT') + CLIP_SKIP_MAPPING = { None: ReturnedEmbeddingsType.LAST_HIDDEN_STATES_NORMALIZED, 1: ReturnedEmbeddingsType.LAST_HIDDEN_STATES_NORMALIZED, @@ -23,15 +25,16 @@ class DiffusersTextualInversionManager(BaseTextualInversionManager): if hasattr(self.pipe, 'embedding_db'): self.pipe.embedding_db.embeddings_used.clear() - # from https://github.com/huggingface/diffusers/blob/705c592ea98ba4e288d837b9cba2767623c78603/src/diffusers/loaders.py#L599 - def maybe_convert_prompt(self, prompt: typing.Union[str, typing.List[str]], tokenizer="PreTrainedTokenizer"): + # code from + # https://github.com/huggingface/diffusers/blob/705c592ea98ba4e288d837b9cba2767623c78603/src/diffusers/loaders.py + def maybe_convert_prompt(self, prompt: typing.Union[str, typing.List[str]], tokenizer: PreTrainedTokenizer): prompts = [prompt] if not isinstance(prompt, typing.List) else prompt prompts = [self._maybe_convert_prompt(p, tokenizer) for p in prompts] if not isinstance(prompt, typing.List): return prompts[0] return prompts - def _maybe_convert_prompt(self, prompt: str, tokenizer="PreTrainedTokenizer"): + def _maybe_convert_prompt(self, prompt: str, tokenizer: PreTrainedTokenizer): tokens = tokenizer.tokenize(prompt) unique_tokens = set(tokens) for token in unique_tokens: @@ -58,7 +61,7 @@ class DiffusersTextualInversionManager(BaseTextualInversionManager): return self.pipe.tokenizer.encode(prompt, add_special_tokens=False) -def get_prompt_schedule(p, prompt, steps): # pylint: disable=unused-argument +def get_prompt_schedule(prompt, steps): t0 = time.time() temp = [] schedule = prompt_parser.get_learned_conditioning_prompt_schedules([prompt], steps)[0] @@ -72,49 +75,46 @@ def get_prompt_schedule(p, prompt, steps): # pylint: disable=unused-argument return temp, len(schedule) > 1 -def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, step: int = 1, clip_skip: typing.Optional[int] = None): # pylint: disable=unused-argument +def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, + clip_skip: typing.Optional[int] = None): if 'StableDiffusion' not in pipe.__class__.__name__ and 'DemoFusion': shared.log.warning(f"Prompt parser not supported: {pipe.__class__.__name__}") return None, None, None, None else: t0 = time.time() - positive_schedule, scheduled = get_prompt_schedule(p, prompts[0], steps) - negative_schedule, neg_scheduled = get_prompt_schedule(p, negative_prompts[0], steps) + positive_schedule, scheduled = get_prompt_schedule(prompts[0], steps) + negative_schedule, neg_scheduled = get_prompt_schedule(negative_prompts[0], steps) p.scheduled_prompt = scheduled or neg_scheduled - p.prompt_embeds = [] - p.positive_pooleds = [] - p.negative_embeds = [] - p.negative_pooleds = [] - cache = {} for i in range(max(len(positive_schedule), len(negative_schedule))): - cached = cache.get(positive_schedule[i % len(positive_schedule)] + negative_schedule[i % len(negative_schedule)], None) - if cached is not None: - prompt_embed, positive_pooled, negative_embed, negative_pooled = cached - else: - prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipe, - positive_schedule[i % len(positive_schedule)], - negative_schedule[i % len(negative_schedule)], - clip_skip) + positive_prompt = positive_schedule[i % len(positive_schedule)] + negative_prompt = negative_schedule[i % len(negative_schedule)] + results = cache.get(positive_prompt + negative_prompt, None) + + if results is None: + results = get_weighted_text_embeddings(pipe, positive_prompt, negative_prompt, clip_skip) + cache[positive_prompt + negative_prompt] = results + + prompt_embed, positive_pooled, negative_embed, negative_pooled = results if prompt_embed is not None: p.prompt_embeds.append(torch.cat([prompt_embed] * len(prompts), dim=0)) if negative_embed is not None: p.negative_embeds.append(torch.cat([negative_embed] * len(negative_prompts), dim=0)) - if positive_pooled is not None and shared.sd_model_type == "sdxl": + if positive_pooled is not None: p.positive_pooleds.append(torch.cat([positive_pooled] * len(prompts), dim=0)) - if negative_pooled is not None and shared.sd_model_type == "sdxl": + if negative_pooled is not None: p.negative_pooleds.append(torch.cat([negative_pooled] * len(negative_prompts), dim=0)) debug(f"Prompt Parser: Elapsed Time {time.time() - t0}") return def get_prompts_with_weights(prompt: str): - manager = DiffusersTextualInversionManager(shared.sd_model, shared.sd_model.tokenizer or shared.sd_model.tokenizer_2) + manager = DiffusersTextualInversionManager(shared.sd_model, + shared.sd_model.tokenizer or shared.sd_model.tokenizer_2) prompt = manager.maybe_convert_prompt(prompt, shared.sd_model.tokenizer or shared.sd_model.tokenizer_2) texts_and_weights = prompt_parser.parse_prompt_attention(prompt) - texts = [t for t, w in texts_and_weights] - text_weights = [w for t, w in texts_and_weights] + texts, text_weights = zip(*texts_and_weights) debug(f'Prompt: weights={texts_and_weights}') return texts, text_weights @@ -129,12 +129,14 @@ def prepare_embedding_providers(pipe, clip_skip): shared.log.warning(f"Prompt parser unsupported: clip_skip={clip_skip}") clip_skip = 2 embedding_type = CLIP_SKIP_MAPPING[clip_skip] - if getattr(pipe, "tokenizer", None) is not None and getattr(pipe, "text_encoder", None) is not None: - embedding = EmbeddingsProvider(tokenizer=pipe.tokenizer, text_encoder=pipe.text_encoder, truncate=False, returned_embeddings_type=embedding_type, device=device) - embeddings_providers.append(embedding) - if getattr(pipe, "tokenizer_2", None) is not None and getattr(pipe, "text_encoder_2", None) is not None: - embedding = EmbeddingsProvider(tokenizer=pipe.tokenizer_2, text_encoder=pipe.text_encoder_2, truncate=False, returned_embeddings_type=embedding_type, device=device) - embeddings_providers.append(embedding) + if hasattr(pipe, "tokenizer") and hasattr(pipe, "text_encoder"): + provider = EmbeddingsProvider(tokenizer=pipe.tokenizer, text_encoder=pipe.text_encoder, truncate=False, + returned_embeddings_type=embedding_type, device=device) + embeddings_providers.append(provider) + if hasattr(pipe, "tokenizer_2") and getattr(pipe, "text_encoder_2"): + provider = EmbeddingsProvider(tokenizer=pipe.tokenizer_2, text_encoder=pipe.text_encoder_2, truncate=False, + returned_embeddings_type=embedding_type, device=device) + embeddings_providers.append(provider) return embeddings_providers @@ -142,12 +144,13 @@ def pad_to_same_length(pipe, embeds): device = pipe.device if str(pipe.device) != 'meta' else devices.device try: # SDXL empty_embed = pipe.encode_prompt("") - except Exception: # SD1.5 + except TypeError: # SD1.5 empty_embed = pipe.encode_prompt("", device, 1, False) - empty_batched = torch.cat([empty_embed[0].to(embeds[0].device)] * embeds[0].shape[0]) max_token_count = max([embed.shape[1] for embed in embeds]) + repeats = max_token_count - min([embed.shape[1] for embed in embeds]) + empty_batched = empty_embed[0].to(embeds[0].device).expand(embeds[0].shape[0], repeats, -1) for i, embed in enumerate(embeds): - while embed.shape[1] < max_token_count: + if embed.shape[1] < max_token_count: embed = torch.cat([embed, empty_batched], dim=1) embeds[i] = embed return embeds @@ -161,12 +164,10 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c neg_prompt = neg_prompt.split("TE2:")[0] ps = [get_prompts_with_weights(p) for p in [prompt, prompt_2]] - positives = [t for t, w in ps] - positive_weights = [w for t, w in ps] + positives, positive_weights = zip(*ps) ns = [get_prompts_with_weights(p) for p in [neg_prompt, neg_prompt_2]] - negatives = [t for t, w in ns] - negative_weights = [w for t, w in ns] - if getattr(pipe, "tokenizer_2", None) is not None and getattr(pipe, "tokenizer", None) is None: + negatives, negative_weights = zip(*ns) + if hasattr(pipe, "tokenizer_2") and not hasattr(pipe, "tokenizer"): positives.pop(0) positive_weights.pop(0) negatives.pop(0) @@ -179,8 +180,8 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c negative_pooled_prompt_embeds = None for i in range(len(embedding_providers)): # add BREAK keyword that splits the prompt into multiple fragments - text = positives[i] - weights = positive_weights[i] + text = list(positives[i]) + weights = list(positive_weights[i]) text.append('BREAK') weights.append(-1) provider_embed = [] @@ -188,13 +189,20 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c pos = text.index('BREAK') debug(f'Prompt: section="{text[:pos]}" len={len(text[:pos])} weights={weights[:pos]}') if len(text[:pos]) > 0: - embed, ptokens = embedding_providers[i].get_embeddings_for_weighted_prompt_fragments(text_batch=[text[:pos]], fragment_weights_batch=[weights[:pos]], device=device, should_return_tokens=True) + embed, ptokens = embedding_providers[i].get_embeddings_for_weighted_prompt_fragments( + text_batch=[text[:pos]], fragment_weights_batch=[weights[:pos]], device=device, + should_return_tokens=True) provider_embed.append(embed) text = text[pos + 1:] weights = weights[pos + 1:] prompt_embeds.append(torch.cat(provider_embed, dim=1)) + debug(f'Prompt: positive unpadded shape = {prompt_embeds[0].shape}') # negative prompt has no keywords - embed, ntokens = embedding_providers[i].get_embeddings_for_weighted_prompt_fragments(text_batch=[negatives[i]], fragment_weights_batch=[negative_weights[i]], device=device, should_return_tokens=True) + embed, ntokens = embedding_providers[i].get_embeddings_for_weighted_prompt_fragments(text_batch=[negatives[i]], + fragment_weights_batch=[ + negative_weights[i]], + device=device, + should_return_tokens=True) negative_prompt_embeds.append(embed) if prompt_embeds[-1].shape[-1] > 768: @@ -212,11 +220,15 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c .argmax(dim=-1), ] else: - pooled_prompt_embeds = embedding_providers[-1].get_pooled_embeddings(texts=[prompt_2], device=device) if prompt_embeds[-1].shape[-1] > 768 else None - negative_pooled_prompt_embeds = embedding_providers[-1].get_pooled_embeddings(texts=[neg_prompt_2], device=device) if negative_prompt_embeds[-1].shape[-1] > 768 else None + pooled_prompt_embeds = embedding_providers[-1].get_pooled_embeddings(texts=[prompt_2], device=device) if \ + prompt_embeds[-1].shape[-1] > 768 else None + negative_pooled_prompt_embeds = embedding_providers[-1].get_pooled_embeddings(texts=[neg_prompt_2], + device=device) if \ + negative_prompt_embeds[-1].shape[-1] > 768 else None prompt_embeds = torch.cat(prompt_embeds, dim=-1) if len(prompt_embeds) > 1 else prompt_embeds[0] - negative_prompt_embeds = torch.cat(negative_prompt_embeds, dim=-1) if len(negative_prompt_embeds) > 1 else negative_prompt_embeds[0] + negative_prompt_embeds = torch.cat(negative_prompt_embeds, dim=-1) if len(negative_prompt_embeds) > 1 else \ + negative_prompt_embeds[0] debug(f'Prompt: shape={prompt_embeds.shape} negative={negative_prompt_embeds.shape}') if prompt_embeds.shape[1] != negative_prompt_embeds.shape[1]: [prompt_embeds, negative_prompt_embeds] = pad_to_same_length(pipe, [prompt_embeds, negative_prompt_embeds])