diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 97d98ed40..a7feb05d9 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -260,8 +260,10 @@ def prepare_embedding_providers(pipe, clip_skip) -> list[EmbeddingsProvider]: else: embedding_type = clip_skip if getattr(pipe, "prior_pipe", None) is not None and getattr(pipe.prior_pipe, "tokenizer", None) is not None and getattr(pipe.prior_pipe, "text_encoder", None) is not None: - provider = EmbeddingsProvider(padding_attention_mask_value=1 if "sote" in pipe.sd_checkpoint_info.name.lower() else 0, tokenizer=pipe.prior_pipe.tokenizer, text_encoder=pipe.prior_pipe.text_encoder, truncate=False, returned_embeddings_type=embedding_type, device=device) + provider = EmbeddingsProvider(padding_attention_mask_value=0, tokenizer=pipe.prior_pipe.tokenizer, text_encoder=pipe.prior_pipe.text_encoder, truncate=False, returned_embeddings_type=embedding_type, device=device) embeddings_providers.append(provider) + no_mask_provider = EmbeddingsProvider(padding_attention_mask_value=1 if "sote" in pipe.sd_checkpoint_info.name.lower() else 0, tokenizer=pipe.prior_pipe.tokenizer, text_encoder=pipe.prior_pipe.text_encoder, truncate=False, returned_embeddings_type=embedding_type, device=device) + embeddings_providers.append(no_mask_provider) elif getattr(pipe, "tokenizer", None) is not None and getattr(pipe, "text_encoder", None) is not None: provider = EmbeddingsProvider(tokenizer=pipe.tokenizer, text_encoder=pipe.text_encoder, truncate=False, returned_embeddings_type=embedding_type, device=device) embeddings_providers.append(provider) @@ -271,7 +273,7 @@ def prepare_embedding_providers(pipe, clip_skip) -> list[EmbeddingsProvider]: return embeddings_providers -def pad_to_same_length(pipe, embeds, embedding_providers=None): +def pad_to_same_length(pipe, embeds, empty_embedding_providers=None): if not hasattr(pipe, 'encode_prompt') and 'StableCascade' not in pipe.__class__.__name__: return embeds device = devices.device @@ -280,7 +282,7 @@ def pad_to_same_length(pipe, embeds, embedding_providers=None): else: try: if 'StableCascade' in pipe.__class__.__name__: - empty_embed = embedding_providers[0].get_embeddings_for_weighted_prompt_fragments(text_batch=[[""]], fragment_weights_batch=[[1]], should_return_tokens=False, device=device) + empty_embed = empty_embedding_providers[0].get_embeddings_for_weighted_prompt_fragments(text_batch=[[""]], fragment_weights_batch=[[1]], should_return_tokens=False, device=device) empty_embed = [empty_embed] else: empty_embed = pipe.encode_prompt("") @@ -344,6 +346,10 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c return prompt_embeds, pooled_prompt_embeds, None, None # no negative support embedding_providers = prepare_embedding_providers(pipe, clip_skip) + empty_embedding_providers = None + if 'StableCascade' in pipe.__class__.__name__: + empty_embedding_providers = [embedding_providers[1]] + embedding_providers = [embedding_providers[0]] prompt_embeds = [] negative_prompt_embeds = [] @@ -414,7 +420,7 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c negative_pooled_prompt_embeds = None debug(f'Prompt: positive={prompt_embeds.shape if prompt_embeds is not None else None} pooled={pooled_prompt_embeds.shape if pooled_prompt_embeds is not None else None} negative={negative_prompt_embeds.shape if negative_prompt_embeds is not None else None} pooled={negative_pooled_prompt_embeds.shape if negative_pooled_prompt_embeds is not None else None}') if prompt_embeds.shape[1] != negative_prompt_embeds.shape[1]: - [prompt_embeds, negative_prompt_embeds] = pad_to_same_length(pipe, [prompt_embeds, negative_prompt_embeds], embedding_providers=embedding_providers) + [prompt_embeds, negative_prompt_embeds] = pad_to_same_length(pipe, [prompt_embeds, negative_prompt_embeds], empty_embedding_providers=empty_embedding_providers) if SD3: device = devices.device t5_prompt_embed = pipe._get_t5_prompt_embeds( # pylint: disable=protected-access