From be264d14d6b4d080c992b607109879441603dd1e Mon Sep 17 00:00:00 2001 From: Disty0 Date: Sun, 26 May 2024 16:01:27 +0300 Subject: [PATCH] Stable Cascade prompt parser support --- modules/processing_args.py | 8 ++++++-- modules/prompt_parser_diffusers.py | 20 ++++++++++++++------ 2 files changed, 20 insertions(+), 8 deletions(-) diff --git a/modules/processing_args.py b/modules/processing_args.py index e3e99c9fc..0e46b49f0 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -103,7 +103,7 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2 shared.log.debug(f'Sampler: steps={len(timesteps)} timesteps={timesteps}') except Exception as e: shared.log.error(f'Sampler timesteps: {e}') - if shared.opts.prompt_attention != 'Fixed attention' and 'StableDiffusion' in model.__class__.__name__ and 'Onnx' not in model.__class__.__name__: + if shared.opts.prompt_attention != 'Fixed attention' and ('StableDiffusion' in model.__class__.__name__ or 'StableCascade' in model.__class__.__name__) and 'Onnx' not in model.__class__.__name__: try: prompt_parser_diffusers.encode_prompts(model, p, prompts, negative_prompts, steps=steps, clip_skip=clip_skip) parser = shared.opts.prompt_attention @@ -119,13 +119,17 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2 if 'prompt' in possible: if hasattr(model, 'text_encoder') and 'prompt_embeds' in possible and len(p.prompt_embeds) > 0 and p.prompt_embeds[0] is not None: args['prompt_embeds'] = p.prompt_embeds[0] - if 'XL' in model.__class__.__name__ and len(getattr(p, 'positive_pooleds', [])) > 0: + if 'StableCascade' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0: + args['prompt_embeds_pooled'] = p.positive_pooleds[0].unsqueeze(0) + elif 'XL' in model.__class__.__name__ and len(getattr(p, 'positive_pooleds', [])) > 0: args['pooled_prompt_embeds'] = p.positive_pooleds[0] else: args['prompt'] = prompts if 'negative_prompt' in possible: if hasattr(model, 'text_encoder') and 'negative_prompt_embeds' in possible and len(p.negative_embeds) > 0 and p.negative_embeds[0] is not None: args['negative_prompt_embeds'] = p.negative_embeds[0] + if 'StableCascade' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0: + args['negative_prompt_embeds_pooled'] = p.negative_pooleds[0].unsqueeze(0) if 'XL' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0: args['negative_pooled_prompt_embeds'] = p.negative_pooleds[0] else: diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 67a714426..fd9da0be2 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -130,7 +130,7 @@ def get_tokens(msg, prompt): def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, clip_skip: typing.Optional[int] = None): - if 'StableDiffusion' not in pipe.__class__.__name__ and 'DemoFusion' not in pipe.__class__.__name__: + if 'StableDiffusion' not in pipe.__class__.__name__ and 'DemoFusion' not in pipe.__class__.__name__ and 'StableCascade' not in pipe.__class__.__name__: shared.log.warning(f"Prompt parser not supported: {pipe.__class__.__name__}") return elif prompts == cache.get('prompts', None) and negative_prompts == cache.get('negative_prompts', None) and cache.get('model_type', None) == shared.sd_model_type: @@ -202,11 +202,16 @@ def get_prompts_with_weights(prompt: str): def prepare_embedding_providers(pipe, clip_skip) -> list[EmbeddingsProvider]: device = pipe.device if str(pipe.device) != 'meta' else devices.device embeddings_providers = [] - if 'XL' in pipe.__class__.__name__: + if 'StableCascade' in pipe.__class__.__name__: + embedding_type = -(clip_skip) + elif 'XL' in pipe.__class__.__name__: embedding_type = -(clip_skip + 1) else: embedding_type = clip_skip - if getattr(pipe, "tokenizer", None) is not None and getattr(pipe, "text_encoder", None) is not None: + 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(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) + 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) if getattr(pipe, "tokenizer_2", None) is not None and getattr(pipe, "text_encoder_2", None) is not None: @@ -216,11 +221,14 @@ def prepare_embedding_providers(pipe, clip_skip) -> list[EmbeddingsProvider]: def pad_to_same_length(pipe, embeds): - if not hasattr(pipe, 'encode_prompt'): + if not hasattr(pipe, 'encode_prompt') and not (hasattr(pipe, "prior_pipe") and hasattr(pipe.prior_pipe, "encode_prompt")): return embeds device = pipe.device if str(pipe.device) != 'meta' else devices.device - try: # SDXL - empty_embed = pipe.encode_prompt("") + try: + if getattr(pipe, "prior_pipe", None) and getattr(pipe.prior_pipe, "text_encoder", None) is not None: # Cascade + empty_embed = pipe.prior_pipe.encode_prompt(device, 1, 1, False, "") + else: # SDXL + empty_embed = pipe.encode_prompt("") except TypeError: # SD1.5 empty_embed = pipe.encode_prompt("", device, 1, False) max_token_count = max([embed.shape[1] for embed in embeds])