From 69c5cff112e8b868feced025b398596f35ea248b Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 13 Nov 2023 17:49:56 -0500 Subject: [PATCH] diffusers prompt handle BREAK --- CHANGELOG.md | 1 + modules/processing_diffusers.py | 2 ++ modules/prompt_parser_diffusers.py | 22 ++++++++++++++++++---- 3 files changed, 21 insertions(+), 4 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index b8111bfc7..aedf129b8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -23,6 +23,7 @@ - Fix **DPM SDE** scheduler - Better support for SD 1.5 **inpainting** models - Add support for **OpenAI Consistency decoder VAE** + - Enhance prompt parsing with long prompts and support for *BREAK* keyword - Update to `diffusers==0.23.0` - **Extra networks** - Use multi-threading for 5x load speedup diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 9a56fc028..bf3b6d430 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -183,6 +183,8 @@ def process_diffusers(p: StableDiffusionProcessing, seeds, prompts, negative_pro negative_prompts = [negative_prompts] while len(negative_prompts) < len(prompts): negative_prompts.append(negative_prompts[-1]) + while len(prompts) < len(negative_prompts): + prompts.append(prompts[-1]) if type(prompts_2) is str: prompts_2 = [prompts_2] if type(prompts_2) is list: diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 600c66d4a..b9fbf8908 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -69,7 +69,7 @@ def encode_prompts(pipeline, prompts: list, negative_prompts: list, clip_skip: t negative_embeds = [] negative_pooleds = [] for i in range(len(prompts)): - prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings_sdxl(pipeline,prompts[i], negative_prompts[i], clip_skip) + prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipeline,prompts[i], negative_prompts[i], clip_skip) prompt_embeds.append(prompt_embed) positive_pooleds.append(positive_pooled) negative_embeds.append(negative_embed) @@ -113,6 +113,7 @@ def prepare_embedding_providers(pipe, clip_skip): embeddings_providers.append(embedding) return embeddings_providers + def pad_to_same_length(embeds): try: #SDXL empty_embed = shared.sd_model.encode_prompt("") @@ -127,7 +128,8 @@ def pad_to_same_length(embeds): embeds[i] = embed return embeds -def get_weighted_text_embeddings_sdxl(pipe, prompt: str = "", neg_prompt: str = "", clip_skip: int = None): + +def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", clip_skip: int = None): prompt_2 = prompt.split("TE2:")[-1] neg_prompt_2 = neg_prompt.split("TE2:")[-1] prompt = prompt.split("TE2:")[0] @@ -152,8 +154,20 @@ def get_weighted_text_embeddings_sdxl(pipe, prompt: str = "", neg_prompt: str = negative_pooled_prompt_embeds = None for i in range(len(embedding_providers)): - embed, ptokens = embedding_providers[i].get_embeddings_for_weighted_prompt_fragments(text_batch=[positives[i]], fragment_weights_batch=[positive_weights[i]], device=pipe.device, should_return_tokens=True) - prompt_embeds.append(embed) + # add BREAK keyword that splits the prompt into multiple fragments + text = positives[i] + weights = positive_weights[i] + text.append('BREAK') + weights.append(-1) + provider_embed = [] + while 'BREAK' in text: + pos = text.index('BREAK') + embed, ptokens = embedding_providers[i].get_embeddings_for_weighted_prompt_fragments(text_batch=[text[:pos]], fragment_weights_batch=[weights[:pos]], device=pipe.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)) + # 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=pipe.device, should_return_tokens=True) negative_prompt_embeds.append(embed)