diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 10c1bc3db..fd50ce14e 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -1,11 +1,16 @@ import torch import modules.shared as shared from compel import Compel, ReturnedEmbeddingsType +import diffusers +import typing -def compel_encode_prompt(pipeline, prompt, negative_prompt, prompt_2=None, negative_prompt_2=None, refiner=False): - if "XL" not in pipeline.__class__.__name__: - print(f"Compel parser is not configured for: {pipeline.__class__.__name__}") - return None, None, None, None +def compel_encode_prompt(pipeline: typing.Any, *args, **kwargs): + compel_encode_fn = COMPEL_ENCODE_FN_DICT.get(type(pipeline), None) + if compel_encode_fn is None: + raise TypeError(f"Compel encoding not yet supported for {type(pipeline).__name__}.") + return compel_encode_fn(pipeline, *args, **kwargs) + +def compel_encode_prompt_sdxl(pipeline: diffusers.StableDiffusionXLPipeline, prompt: str, negative_prompt: str, prompt_2: typing.Optional[str]=None, negative_prompt_2: typing.Optional[str]=None, refiner=False): compel_te1 = Compel( tokenizer=pipeline.tokenizer, text_encoder=pipeline.text_encoder, @@ -19,7 +24,7 @@ def compel_encode_prompt(pipeline, prompt, negative_prompt, prompt_2=None, negat returned_embeddings_type=ReturnedEmbeddingsType.PENULTIMATE_HIDDEN_STATES_NON_NORMALIZED, requires_pooled=True, ) - if not refiner: + if refiner is None: positive_te1 = compel_te1(prompt) positive_te2, pooled = compel_te2(prompt_2) positive = torch.cat((positive_te1, positive_te2), dim=-1) @@ -27,10 +32,12 @@ def compel_encode_prompt(pipeline, prompt, negative_prompt, prompt_2=None, negat negative_te1 = compel_te1(negative_prompt) negative_te2, negative_pooled = compel_te2(negative_prompt_2) negative = torch.cat((negative_te1, negative_te2), dim=-1) - if refiner: + else: positive, pooled = compel_te2(prompt) negative, negative_pooled = compel_te2(negative_prompt) [prompt_embed, negative_embed] = compel_te2.pad_conditioning_tensors_to_same_length([positive, negative]) - return prompt_embed, pooled, negative_embed, negative_pooled \ No newline at end of file + return prompt_embed, pooled, negative_embed, negative_pooled + +COMPEL_ENCODE_FN_DICT = {diffusers.StableDiffusionXLPipeline: compel_encode_prompt_sdxl} \ No newline at end of file