add attention masks

This commit is contained in:
Enes Sadık Özbek
2025-06-23 01:27:17 +00:00
parent 5e9c6bfd4e
commit 954138bfc1
4 changed files with 124 additions and 48 deletions
+4
View File
@@ -183,6 +183,8 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t
args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds')
elif 'Flux' in model.__class__.__name__:
args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds')
elif 'Chroma' in model.__class__.__name__:
args['prompt_attention_mask'] = prompt_parser_diffusers.embedder('prompt_attention_masks')
else:
args['prompt'] = prompts
if 'negative_prompt' in possible:
@@ -199,6 +201,8 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t
args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds')
elif 'StableDiffusion3' in model.__class__.__name__:
args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds')
elif 'Chroma' in model.__class__.__name__:
args['negative_prompt_attention_mask'] = prompt_parser_diffusers.embedder('negative_prompt_attention_masks')
else:
if 'PixArtSigmaPipeline' in model.__class__.__name__: # pixart-sigma pipeline throws list-of-list for negative prompt
args['negative_prompt'] = negative_prompts[0]
+4
View File
@@ -152,6 +152,8 @@ class StableDiffusionProcessing:
self.positive_pooleds = []
self.negative_embeds = []
self.negative_pooleds = []
self.prompt_attention_masks = []
self.negative_prompt_attention_masks = []
self.disable_extra_networks = False
self.iteration = 0
self.network_data = {}
@@ -321,6 +323,8 @@ class StableDiffusionProcessing:
self.positive_pooleds = []
self.negative_embeds = []
self.negative_pooleds = []
self.prompt_attention_masks = []
self.negative_prompt_attention_mask = []
def __str__(self):
return f'{self.__class__.__name__}: {self.__dict__}'
+41 -12
View File
@@ -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
+75 -36
View File
@@ -88,8 +88,8 @@ def get_prompts_tokens_with_weights(
def get_prompts_tokens_with_weights_t5(
t5_tokenizer: T5Tokenizer
, prompt: str
t5_tokenizer: T5Tokenizer,
prompt: str
):
"""
Get prompt token ids and weights, this function works for both prompt and negative prompt
@@ -98,18 +98,21 @@ def get_prompts_tokens_with_weights_t5(
prompt = "empty"
texts_and_weights = parse_prompt_attention(prompt)
text_tokens, text_weights = [], []
text_tokens, text_weights, text_masks = [], [], []
for word, weight in texts_and_weights:
# tokenize and discard the starting and the ending token
token = t5_tokenizer(
word
, truncation=False # so that tokenize whatever length prompt
, add_special_tokens=True
).input_ids
# the returned token is a 1d list: [320, 1125, 539, 320]
inputs = t5_tokenizer(
word,
truncation=False, # so that tokenize whatever length prompt
add_special_tokens=True,
)
token = inputs.input_ids
mask = inputs.attention_mask
# merge the new tokens to the all tokens holder: text_tokens
text_tokens = [*text_tokens, *token]
text_masks = [*text_masks, *mask]
# each token chunk will come with one weight, like ['red cat', 2.0]
# need to expand weight for each token.
@@ -117,7 +120,7 @@ def get_prompts_tokens_with_weights_t5(
# append the weight back to the weight holder: text_weights
text_weights = [*text_weights, *chunk_weights]
return text_tokens, text_weights
return text_tokens, text_weights, text_masks
def group_tokens_and_weights(
@@ -1070,11 +1073,11 @@ def get_weighted_text_embeddings_sd3(
)
# tokenizer 3
prompt_tokens_3, prompt_weights_3 = get_prompts_tokens_with_weights_t5(
prompt_tokens_3, prompt_weights_3, _ = get_prompts_tokens_with_weights_t5(
pipe.tokenizer_3, prompt
)
neg_prompt_tokens_3, neg_prompt_weights_3 = get_prompts_tokens_with_weights_t5(
neg_prompt_tokens_3, neg_prompt_weights_3, _ = get_prompts_tokens_with_weights_t5(
pipe.tokenizer_3, neg_prompt
)
@@ -1366,7 +1369,7 @@ def get_weighted_text_embeddings_flux1(
)
# tokenizer 2 - google/t5-v1_1-xxl
prompt_tokens_2, prompt_weights_2 = get_prompts_tokens_with_weights_t5(
prompt_tokens_2, prompt_weights_2, _ = get_prompts_tokens_with_weights_t5(
pipe.tokenizer_2, prompt2
)
@@ -1443,23 +1446,26 @@ def get_weighted_text_embeddings_chroma(
neg_prompt (str)
device (torch.device, optional): Device to run the embeddings on.
Returns:
prompt_embeds (T5 prompt embeds as torch.Tensor)
neg_prompt_embeds (T5 prompt embeds as torch.Tensor)
prompt_embeds (torch.Tensor)
prompt_attention_mask (torch.Tensor)
neg_prompt_embeds (torch.Tensor)
neg_prompt_attention_mask (torch.Tensor)
"""
if device is None:
device = pipe.text_encoder.device
# prompt
prompt_tokens, prompt_weights = get_prompts_tokens_with_weights_t5(
# positive prompt
prompt_tokens, prompt_weights, prompt_masks = get_prompts_tokens_with_weights_t5(
pipe.tokenizer, prompt
)
prompt_tokens = torch.tensor([prompt_tokens], dtype=torch.long)
prompt_tokens = torch.tensor([prompt_tokens], dtype=torch.long).to(device)
prompt_masks = torch.tensor([prompt_masks], dtype=torch.long).to(device)
t5_prompt_embeds = pipe.text_encoder(prompt_tokens.to(device))[0].squeeze(0)
t5_prompt_embeds = pipe.text_encoder(prompt_tokens, output_hidden_states=False, attention_mask=prompt_masks)[0].squeeze(0)
t5_prompt_embeds = t5_prompt_embeds.to(device=device)
# add weight to t5 prompt
# add weight to t5 positive embeddings
for z in range(len(prompt_weights)):
if prompt_weights[z] != 1.0:
t5_prompt_embeds[z] = t5_prompt_embeds[z] * prompt_weights[z]
@@ -1467,34 +1473,67 @@ def get_weighted_text_embeddings_chroma(
t5_prompt_embeds = t5_prompt_embeds.to(dtype=pipe.text_encoder.dtype, device=device)
# negative prompt
neg_prompt_tokens, neg_prompt_weights = get_prompts_tokens_with_weights_t5(
neg_prompt_tokens, neg_prompt_weights, neg_prompt_masks = get_prompts_tokens_with_weights_t5(
pipe.tokenizer, neg_prompt
)
neg_prompt_tokens = torch.tensor([neg_prompt_tokens], dtype=torch.long)
neg_prompt_tokens = torch.tensor([neg_prompt_tokens], dtype=torch.long).to(device)
neg_prompt_masks = torch.tensor([neg_prompt_masks], dtype=torch.long).to(device)
t5_neg_prompt_embeds = pipe.text_encoder(neg_prompt_tokens.to(device))[0].squeeze(0)
t5_neg_prompt_embeds = pipe.text_encoder(neg_prompt_tokens, output_hidden_states=False, attention_mask=neg_prompt_masks)[0].squeeze(0)
t5_neg_prompt_embeds = t5_neg_prompt_embeds.to(device=device)
# add weight to neg t5 embeddings
# add weight to negative t5 embeddings
for z in range(len(neg_prompt_weights)):
if neg_prompt_weights[z] != 1.0:
t5_neg_prompt_embeds[z] = t5_neg_prompt_embeds[z] * neg_prompt_weights[z]
t5_neg_prompt_embeds = t5_neg_prompt_embeds.unsqueeze(0)
t5_neg_prompt_embeds = t5_neg_prompt_embeds.to(dtype=pipe.text_encoder.dtype, device=device)
def pad_prompt_embeds_to_same_size(prompt_embeds_a, prompt_embeds_b):
size_a = prompt_embeds_a.size(1)
size_b = prompt_embeds_b.size(1)
embeds, masks = pad_prompt_embeds_to_same_size_chroma(
pipe,
[t5_prompt_embeds, t5_neg_prompt_embeds],
[prompt_masks, neg_prompt_masks]
)
if size_a < size_b:
pad_size = size_b - size_a
prompt_embeds_a = F.pad(prompt_embeds_a, (0, 0, 0, pad_size)) # Pad dim=1
elif size_b < size_a:
pad_size = size_a - size_b
prompt_embeds_b = F.pad(prompt_embeds_b, (0, 0, 0, pad_size)) # Pad dim=1
return embeds[0], masks[0], embeds[1], masks[1]
return prompt_embeds_a, prompt_embeds_b
# chroma needs positive and negative prompt embeddings to have the same length (for now)
return pad_prompt_embeds_to_same_size(t5_prompt_embeds, t5_neg_prompt_embeds)
def pad_prompt_embeds_to_same_size_chroma(pipe, embeds, masks):
"""
Implementation of Chroma's padding for prompt embeddings.
Pads the embeddings to the maximum length found in the batch, while ensuring
that the padding tokens are masked correctly and keeping one padding token unmasked.
https://huggingface.co/lodestones/Chroma#tldr-masking-t5-padding-tokens-enhanced-fidelity-and-increased-stability-during-training
"""
pad_token_id = pipe.tokenizer.pad_token_id
max_token_count = max([embed.shape[1] for embed in embeds])
padded_embeds = []
padded_masks = []
for embed, mask in zip(embeds, masks):
current_length = embed.shape[1]
if current_length < max_token_count:
pad_length = max_token_count - current_length
embed_pad = torch.full(
(1, pad_length, embed.shape[-1]),
fill_value=pad_token_id,
dtype=embed.dtype,
device=embed.device
)
padded_embed = torch.cat([embed, embed_pad], dim=1)
mask_pad = torch.ones(1, pad_length, device=mask.device)
mask_pad[0, 0] = 0 # keep one padding token unmasked, see linked Chroma docs
padded_mask = torch.cat([mask, mask_pad], dim=1)
padded_embeds.append(padded_embed)
padded_masks.append(padded_mask)
else:
padded_embeds.append(embed)
padded_masks.append(mask)
return padded_embeds, padded_masks