diff --git a/modules/processing_args.py b/modules/processing_args.py index fca0115de..7e9be06e9 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -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] diff --git a/modules/processing_class.py b/modules/processing_class.py index 6b10f07a7..fc697d997 100644 --- a/modules/processing_class.py +++ b/modules/processing_class.py @@ -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__}' diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 1e4f5246a..37d6fe9ae 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -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 diff --git a/modules/prompt_parser_xhinker.py b/modules/prompt_parser_xhinker.py index d2cc1971d..ed31fc5d4 100644 --- a/modules/prompt_parser_xhinker.py +++ b/modules/prompt_parser_xhinker.py @@ -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