diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index f29a19b6d..7dcc04336 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -6,10 +6,8 @@ from compel import ReturnedEmbeddingsType from compel.embeddings_provider import BaseTextualInversionManager, EmbeddingsProvider from modules import shared, prompt_parser, devices - debug = shared.log.info if os.environ.get('SD_PROMPT_DEBUG', None) is not None else lambda *args, **kwargs: None - CLIP_SKIP_MAPPING = { None: ReturnedEmbeddingsType.LAST_HIDDEN_STATES_NORMALIZED, 1: ReturnedEmbeddingsType.LAST_HIDDEN_STATES_NORMALIZED, @@ -59,6 +57,7 @@ class DiffusersTextualInversionManager(BaseTextualInversionManager): debug(f'Prompt: expand={prompt}') return self.pipe.tokenizer.encode(prompt, add_special_tokens=False) + def get_prompt_schedule(p, prompt, steps): t0 = time.time() temp = [] @@ -69,14 +68,16 @@ def get_prompt_schedule(p, prompt, steps): for s in range(steps): if len(temp) < s + 1 <= chunk[0]: temp.append(chunk[1]) - debug(f'Prompt: schedule={temp} time={time.time()-t0}') + debug(f'Prompt: schedule={temp} time={time.time() - t0}') 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): + +def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, step: int = 1, 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) p.scheduled_prompt = scheduled or neg_scheduled @@ -87,23 +88,24 @@ def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, p.negative_pooleds = [] cache = {} - for i in range(len(positive_schedule)): - cached = cache.get(positive_schedule[i]+negative_schedule[i], None) + 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 + 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], - negative_schedule[i], - clip_skip) + 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) if prompt_embed is not None: - p.prompt_embeds.append(torch.cat([prompt_embed]*len(prompts), dim=0)) + 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)) + 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": - p.positive_pooleds.append(torch.cat([positive_pooled]*len(prompts), dim=0)) + p.positive_pooleds.append(torch.cat([positive_pooled] * len(prompts), dim=0)) if negative_pooled is not None and shared.sd_model_type == "sdxl": - p.negative_pooleds.append(torch.cat([negative_pooled]*len(negative_prompts), dim=0)) + p.negative_pooleds.append(torch.cat([negative_pooled] * len(negative_prompts), dim=0)) + debug(f"Prompt Parser: Elapsed Time {time.time() - t0}") return @@ -138,9 +140,9 @@ def prepare_embedding_providers(pipe, clip_skip): def pad_to_same_length(pipe, embeds): device = pipe.device if str(pipe.device) != 'meta' else devices.device - try: #SDXL + try: # SDXL empty_embed = pipe.encode_prompt("") - except Exception: #SD1.5 + except Exception: # 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]) @@ -174,7 +176,7 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c prompt_embeds = [] negative_prompt_embeds = [] pooled_prompt_embeds = None - negative_pooled_prompt_embeds = None + 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] @@ -188,8 +190,8 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c 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) provider_embed.append(embed) - text = text[pos+1:] - weights = weights[pos+1:] + text = text[pos + 1:] + weights = weights[pos + 1:] prompt_embeds.append(torch.cat(provider_embed, dim=1)) # 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) @@ -198,17 +200,17 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c if prompt_embeds[-1].shape[-1] > 768: if shared.opts.diffusers_pooled == "weighted": pooled_prompt_embeds = prompt_embeds[-1][ - torch.arange(prompt_embeds[-1].shape[0], device=device), - (ptokens.to(dtype=torch.int, device=device) == 49407) - .int() - .argmax(dim=-1), - ] + torch.arange(prompt_embeds[-1].shape[0], device=device), + (ptokens.to(dtype=torch.int, device=device) == 49407) + .int() + .argmax(dim=-1), + ] negative_pooled_prompt_embeds = negative_prompt_embeds[-1][ - torch.arange(negative_prompt_embeds[-1].shape[0], device=device), - (ntokens.to(dtype=torch.int, device=device) == 49407) - .int() - .argmax(dim=-1), - ] + torch.arange(negative_prompt_embeds[-1].shape[0], device=device), + (ntokens.to(dtype=torch.int, device=device) == 49407) + .int() + .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