mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
add attention masks
This commit is contained in:
@@ -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]
|
||||
|
||||
@@ -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__}'
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user