mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
add attention masks
This commit is contained in:
@@ -62,6 +62,8 @@ class PromptEmbedder:
|
||||
self.positive_pooleds = [[]] * self.batchsize
|
||||
self.negative_prompt_embeds = [[]] * self.batchsize
|
||||
self.negative_pooleds = [[]] * self.batchsize
|
||||
self.prompt_attention_masks = [[]] * self.batchsize
|
||||
self.negative_prompt_attention_masks = [[]] * self.batchsize
|
||||
self.positive_schedule = None
|
||||
self.negative_schedule = None
|
||||
self.scheduled_prompt = False
|
||||
@@ -109,13 +111,17 @@ class PromptEmbedder:
|
||||
if not any(flatten(emb) for emb in [self.prompt_embeds,
|
||||
self.negative_prompt_embeds,
|
||||
self.positive_pooleds,
|
||||
self.negative_pooleds]):
|
||||
self.negative_pooleds,
|
||||
self.prompt_attention_masks,
|
||||
self.negative_prompt_attention_masks]):
|
||||
return False
|
||||
else:
|
||||
cache[key] = {'prompt_embeds': self.prompt_embeds,
|
||||
'negative_prompt_embeds': self.negative_prompt_embeds,
|
||||
'positive_pooleds': self.positive_pooleds,
|
||||
'negative_pooleds': self.negative_pooleds,
|
||||
'prompt_attention_masks': self.prompt_attention_masks,
|
||||
'negative_prompt_attention_masks': self.negative_prompt_attention_masks,
|
||||
}
|
||||
debug(f"Prompt cache: add={key}")
|
||||
while len(cache) > int(shared.opts.sd_textencoder_cache_size):
|
||||
@@ -129,6 +135,8 @@ class PromptEmbedder:
|
||||
self.positive_pooleds = [self.positive_pooleds[0]] * self.batchsize
|
||||
self.negative_prompt_embeds = [self.negative_prompt_embeds[0]] * self.batchsize
|
||||
self.negative_pooleds = [self.negative_pooleds[0]] * self.batchsize
|
||||
self.prompt_attention_masks = [self.prompt_attention_masks[0]] * self.batchsize
|
||||
self.negative_prompt_attention_masks = [self.negative_prompt_attention_masks[0]] * self.batchsize
|
||||
debug(f"Prompt cache: get={key}")
|
||||
return True
|
||||
|
||||
@@ -167,15 +175,33 @@ class PromptEmbedder:
|
||||
self.positive_pooleds[batchidx].append(self.positive_pooleds[batchidx][idx])
|
||||
if len(self.negative_pooleds[batchidx]) > 0:
|
||||
self.negative_pooleds[batchidx].append(self.negative_pooleds[batchidx][idx])
|
||||
if len(self.prompt_attention_masks[batchidx]) > 0:
|
||||
self.prompt_attention_masks[batchidx].append(self.prompt_attention_masks[batchidx][idx])
|
||||
if len(self.negative_prompt_attention_masks[batchidx]) > 0:
|
||||
self.negative_prompt_attention_masks[batchidx].append(self.negative_prompt_attention_masks[batchidx][idx])
|
||||
|
||||
def encode(self, pipe, positive_prompt, negative_prompt, batchidx):
|
||||
global last_attention # pylint: disable=global-statement
|
||||
self.attention = shared.opts.prompt_attention
|
||||
last_attention = self.attention
|
||||
if self.attention == "xhinker":
|
||||
prompt_embed, positive_pooled, negative_embed, negative_pooled = get_xhinker_text_embeddings(pipe, positive_prompt, negative_prompt, self.clip_skip)
|
||||
(
|
||||
prompt_embed,
|
||||
positive_pooled,
|
||||
prompt_attention_mask,
|
||||
negative_embed,
|
||||
negative_pooled,
|
||||
negative_prompt_attention_mask
|
||||
) = get_xhinker_text_embeddings(pipe, positive_prompt, negative_prompt, self.clip_skip)
|
||||
else:
|
||||
prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipe, positive_prompt, negative_prompt, self.clip_skip)
|
||||
(
|
||||
prompt_embed,
|
||||
positive_pooled,
|
||||
prompt_attention_mask,
|
||||
negative_embed,
|
||||
negative_pooled,
|
||||
negative_prompt_attention_mask
|
||||
) = get_weighted_text_embeddings(pipe, positive_prompt, negative_prompt, self.clip_skip)
|
||||
if prompt_embed is not None:
|
||||
self.prompt_embeds[batchidx].append(prompt_embed)
|
||||
if negative_embed is not None:
|
||||
@@ -184,7 +210,10 @@ class PromptEmbedder:
|
||||
self.positive_pooleds[batchidx].append(positive_pooled)
|
||||
if negative_pooled is not None:
|
||||
self.negative_pooleds[batchidx].append(negative_pooled)
|
||||
|
||||
if prompt_attention_mask is not None:
|
||||
self.prompt_attention_masks[batchidx].append(prompt_attention_mask)
|
||||
if negative_prompt_attention_mask is not None:
|
||||
self.negative_prompt_attention_masks[batchidx].append(negative_prompt_attention_mask)
|
||||
if debug_enabled:
|
||||
get_tokens(pipe, 'positive', positive_prompt)
|
||||
get_tokens(pipe, 'negative', negative_prompt)
|
||||
@@ -509,11 +538,11 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c
|
||||
|
||||
if "Flux" in pipe.__class__.__name__: # clip is only used for the pooled embeds
|
||||
prompt_embeds, pooled_prompt_embeds, _ = pipe.encode_prompt(prompt=prompt, prompt_2=prompt_2, device=device, num_images_per_prompt=1)
|
||||
return prompt_embeds, pooled_prompt_embeds, None, None # no negative support
|
||||
return prompt_embeds, pooled_prompt_embeds, None, None, None, None # no negative support
|
||||
|
||||
if "Chroma" in pipe.__class__.__name__: # does not use clip and has no pooled embeds
|
||||
prompt_embeds, _, _, negative_prompt_embeds, _, _ = pipe.encode_prompt(prompt=prompt, negative_prompt=neg_prompt, device=device, num_images_per_prompt=1)
|
||||
return prompt_embeds, None, negative_prompt_embeds, None
|
||||
prompt_embeds, _, prompt_attention_mask, negative_prompt_embeds, _, negative_prompt_attention_mask = pipe.encode_prompt(prompt=prompt, negative_prompt=neg_prompt, device=device, num_images_per_prompt=1)
|
||||
return prompt_embeds, None, prompt_attention_mask, negative_prompt_embeds, None, negative_prompt_attention_mask
|
||||
|
||||
if "HiDreamImage" in pipe.__class__.__name__: # clip is only used for the pooled embeds
|
||||
prompt_embeds_t5, negative_prompt_embeds_t5, prompt_embeds_llama3, negative_prompt_embeds_llama3, pooled_prompt_embeds, negative_pooled_prompt_embeds = pipe.encode_prompt(
|
||||
@@ -523,7 +552,7 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c
|
||||
)
|
||||
prompt_embeds = [prompt_embeds_t5, prompt_embeds_llama3]
|
||||
negative_prompt_embeds = [negative_prompt_embeds_t5, negative_prompt_embeds_llama3]
|
||||
return prompt_embeds, pooled_prompt_embeds, negative_prompt_embeds, negative_pooled_prompt_embeds
|
||||
return prompt_embeds, pooled_prompt_embeds, None, negative_prompt_embeds, negative_pooled_prompt_embeds, None
|
||||
|
||||
if prompt != prompt_2:
|
||||
ps = [get_prompts_with_weights(pipe, p) for p in [prompt, prompt_2]]
|
||||
@@ -636,7 +665,7 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c
|
||||
negative_prompt_embeds, (0, t5_negative_prompt_embed.shape[-1] - negative_prompt_embeds.shape[-1])
|
||||
).to(device)
|
||||
negative_prompt_embeds = torch.cat([negative_prompt_embeds, t5_negative_prompt_embed], dim=-2)
|
||||
return prompt_embeds, pooled_prompt_embeds, negative_prompt_embeds, negative_pooled_prompt_embeds
|
||||
return prompt_embeds, pooled_prompt_embeds, None, negative_prompt_embeds, negative_pooled_prompt_embeds, None
|
||||
|
||||
|
||||
def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", clip_skip: int = None):
|
||||
@@ -650,7 +679,7 @@ def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", cl
|
||||
neg_prompt_2 = pipe.maybe_convert_prompt(neg_prompt_2, pipe.tokenizer_2)
|
||||
except Exception:
|
||||
pass
|
||||
prompt_embed = positive_pooled = negative_embed = negative_pooled = None
|
||||
prompt_embed = positive_pooled = negative_embed = negative_pooled = prompt_attention_mask = negative_prompt_attention_mask = None
|
||||
|
||||
te1_device, te2_device, te3_device = None, None, None
|
||||
if hasattr(pipe, "text_encoder") and pipe.text_encoder.device != devices.device:
|
||||
@@ -668,7 +697,7 @@ def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", cl
|
||||
elif 'Flux' in pipe.__class__.__name__:
|
||||
prompt_embed, positive_pooled = get_weighted_text_embeddings_flux1(pipe=pipe, prompt=prompt, prompt2=prompt_2, device=devices.device)
|
||||
elif 'Chroma' in pipe.__class__.__name__:
|
||||
prompt_embed, negative_embed = get_weighted_text_embeddings_chroma(pipe=pipe, prompt=prompt, neg_prompt=neg_prompt, device=devices.device)
|
||||
prompt_embed, prompt_attention_mask, negative_embed, negative_prompt_attention_mask = get_weighted_text_embeddings_chroma(pipe=pipe, prompt=prompt, neg_prompt=neg_prompt, device=devices.device)
|
||||
elif 'XL' in pipe.__class__.__name__:
|
||||
prompt_embed, negative_embed, positive_pooled, negative_pooled = get_weighted_text_embeddings_sdxl_2p(pipe=pipe, prompt=prompt, prompt_2=prompt_2, neg_prompt=neg_prompt, neg_prompt_2=neg_prompt_2)
|
||||
else:
|
||||
@@ -681,4 +710,4 @@ def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", cl
|
||||
if te3_device is not None:
|
||||
sd_models.move_model(pipe.text_encoder_3, te1_device, force=True)
|
||||
|
||||
return prompt_embed, positive_pooled, negative_embed, negative_pooled
|
||||
return prompt_embed, positive_pooled, prompt_attention_mask, negative_embed, negative_pooled, negative_prompt_attention_mask
|
||||
|
||||
Reference in New Issue
Block a user