diffusers prompt handle BREAK

This commit is contained in:
Vladimir Mandic
2023-11-13 17:49:56 -05:00
parent 39e5d614ff
commit 69c5cff112
3 changed files with 21 additions and 4 deletions
+1
View File
@@ -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
+2
View File
@@ -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:
+18 -4
View File
@@ -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)