diff --git a/modules/processing_args.py b/modules/processing_args.py index 1ea91fb08..8dd47bd87 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -117,7 +117,8 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2 'Flux' in model.__class__.__name__ ): try: - prompt_parser_diffusers.encode_prompts(model, p, prompts, negative_prompts, steps=steps, clip_skip=clip_skip) + # prompt_parser_diffusers.encode_prompts(model, p, prompts, negative_prompts, steps=steps, clip_skip=clip_skip) + p.embedder = prompt_parser_diffusers.PromptEmbedder(prompts, negative_prompts, clip_skip, p) parser = shared.opts.prompt_attention except Exception as e: shared.log.error(f'Prompt parser encode: {e}') @@ -128,27 +129,27 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2 if 'prompt' in possible: if 'OmniGen' in model.__class__.__name__: prompts = [p.replace('|image|', '<|image_1|>') for p in prompts] - 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 hasattr(model, 'text_encoder') and 'prompt_embeds' in possible and p.embedder is not None: + args['prompt_embeds'] = p.embedder('prompt_embeds') 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] - elif 'StableDiffusion3' in model.__class__.__name__ and len(getattr(p, 'positive_pooleds', [])) > 0: - args['pooled_prompt_embeds'] = p.positive_pooleds[0] - elif 'Flux' in model.__class__.__name__ and len(getattr(p, 'positive_pooleds', [])) > 0: - args['pooled_prompt_embeds'] = p.positive_pooleds[0] + args['prompt_embeds_pooled'] = p.embedder('positive_pooleds').unsqueeze(0) + elif 'XL' in model.__class__.__name__ and p.embedder is not None: + args['pooled_prompt_embeds'] = p.embedder('positive_pooleds') + elif 'StableDiffusion3' in model.__class__.__name__ and p.embedder is not None: + args['pooled_prompt_embeds'] = p.embedder('positive_pooleds') + elif 'Flux' in model.__class__.__name__ and p.embedder is not None: + args['pooled_prompt_embeds'] = p.embedder('positive_pooleds') 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] - if 'StableDiffusion3' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0: - args['negative_pooled_prompt_embeds'] = p.negative_pooleds[0] + if hasattr(model, 'text_encoder') and 'negative_prompt_embeds' in possible and p.embedder is not None: + args['negative_prompt_embeds'] = p.embedder('negative_embeds') + if 'StableCascade' in model.__class__.__name__ and p.embedder is not None: + args['negative_prompt_embeds_pooled'] = p.embedder('negative_pooleds').unsqueeze(0) + if 'XL' in model.__class__.__name__ and p.embedder is not None: + args['negative_pooled_prompt_embeds'] = p.embedder('negative_pooleds') + if 'StableDiffusion3' in model.__class__.__name__ and p.embedder is not None: + args['negative_pooled_prompt_embeds'] = p.embedder('negative_pooleds') 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_callbacks.py b/modules/processing_callbacks.py index 47c8e8827..2584ee796 100644 --- a/modules/processing_callbacks.py +++ b/modules/processing_callbacks.py @@ -67,14 +67,14 @@ def diffusers_callback(pipe, step: int, timestep: int, kwargs: dict): pipe.set_ip_adapter_scale(ip_adapter_scales) if step != getattr(pipe, 'num_timesteps', 0): kwargs = processing_correction.correction_callback(p, timestep, kwargs) - if p.scheduled_prompt and 'prompt_embeds' in kwargs and 'negative_prompt_embeds' in kwargs: - try: - i = (step + 1) % len(p.prompt_embeds) - kwargs["prompt_embeds"] = p.prompt_embeds[i][0:1].expand(kwargs["prompt_embeds"].shape) - j = (step + 1) % len(p.negative_embeds) - kwargs["negative_prompt_embeds"] = p.negative_embeds[j][0:1].expand(kwargs["negative_prompt_embeds"].shape) - except Exception as e: - shared.log.debug(f"Callback: {e}") + # if p.scheduled_prompt and 'prompt_embeds' in kwargs and 'negative_prompt_embeds' in kwargs: + # try: + # i = (step + 1) % len(p.prompt_embeds) + # kwargs["prompt_embeds"] = p.prompt_embeds[i][0:1].expand(kwargs["prompt_embeds"].shape) + # j = (step + 1) % len(p.negative_embeds) + # kwargs["negative_prompt_embeds"] = p.negative_embeds[j][0:1].expand(kwargs["negative_prompt_embeds"].shape) + # except Exception as e: + # shared.log.debug(f"Callback: {e}") if step == int(getattr(pipe, 'num_timesteps', 100) * p.cfg_end) and 'prompt_embeds' in kwargs and 'negative_prompt_embeds' in kwargs: if "PAG" in shared.sd_model.__class__.__name__: pipe._guidance_scale = 1.001 if pipe._guidance_scale > 1 else pipe._guidance_scale # pylint: disable=protected-access diff --git a/modules/processing_class.py b/modules/processing_class.py index 9265ea3cf..d38aae790 100644 --- a/modules/processing_class.py +++ b/modules/processing_class.py @@ -204,11 +204,12 @@ class StableDiffusionProcessing: self.hdr_color_picker=hdr_color_picker self.hdr_tint_ratio=hdr_tint_ratio # globals - self.scheduled_prompt: bool = False - self.prompt_embeds = [] - self.positive_pooleds = [] - self.negative_embeds = [] - self.negative_pooleds = [] + self.embedder = None + # self.scheduled_prompt: bool = False + # self.prompt_embeds = [] + # self.positive_pooleds = [] + # self.negative_embeds = [] + # self.negative_pooleds = [] @property def sd_model(self): diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index cc814f379..42a6bcbbd 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -7,6 +7,7 @@ from compel.embeddings_provider import BaseTextualInversionManager, EmbeddingsPr from transformers import PreTrainedTokenizer from modules import shared, prompt_parser, devices, sd_models from modules.prompt_parser_xhinker import get_weighted_text_embeddings_sd15, get_weighted_text_embeddings_sdxl_2p, get_weighted_text_embeddings_sd3, get_weighted_text_embeddings_flux1 +from modules.processing_helpers import fix_prompts debug_enabled = os.environ.get('SD_PROMPT_DEBUG', None) debug = shared.log.trace if os.environ.get('SD_PROMPT_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -17,6 +18,131 @@ token_type = None # used by helper get_tokens cache = {} +def prompt_compatible(): + if ( + 'StableDiffusion' not in shared.sd_model.__class__.__name__ and + 'DemoFusion' not in shared.sd_model.__class__.__name__ and + 'StableCascade' not in shared.sd_model.__class__.__name__ and + 'Flux' not in shared.sd_model.__class__.__name__ + ): + shared.log.warning(f"Prompt parser not supported: {shared.sd_model.__class__.__name__}") + return False + return True + + +def prepare_model(): + pipe = shared.sd_model + if shared.opts.diffusers_offload_mode == "balanced": + pipe = sd_models.apply_balanced_offload(pipe) + elif hasattr(pipe, "maybe_free_model_hooks"): + pipe.maybe_free_model_hooks() + devices.torch_gc() + return pipe + + +class PromptEmbedder: + def __init__(self, prompts, negative_prompts, clip_skip, p): + t0 = time.time() + # self.prompts, self.negative_prompts, _, _ = fix_prompts(prompts, negative_prompts, None, None) + self.prompts = prompts + self.negative_prompts = negative_prompts + self.batchsize = len(self.prompts) + self.allsame = self.compare_prompts() # collapses batched prompts to single prompt if same + self.steps = p.steps + self.clip_skip = clip_skip + self.prompt_embeds = [[]] * self.batchsize + self.positive_pooleds = [[]] * self.batchsize + self.negative_embeds = [[]] * self.batchsize + self.negative_pooleds = [[]] * self.batchsize + self.positive_schedule = None + self.negative_schedule = None + self.scheduled_prompt = False + pipe = prepare_model() + # per prompt in batch + for batchidx, (prompt, negative_prompt) in enumerate(zip(self.prompts, self.negative_prompts)): + self.prepare_schedule(prompt, negative_prompt) + if self.scheduled_prompt: + self.scheduled_encode(pipe, batchidx) + else: + self.encode(pipe, prompt, negative_prompt, batchidx) + if self.allsame: + self.duplicate_embeds() + debug(f"Prompt encode: time={(time.time() - t0):.3f}") + + def compare_prompts(self): + same = (self.prompts == [self.prompts[0]] * len(self.prompts) and + self.negative_prompts == [self.negative_prompts[0]] * len(self.negative_prompts)) + if same: + self.prompts = [self.prompts[0]] + self.negative_prompts = [self.negative_prompts[0]] + return same + + def prepare_schedule(self, prompt, negative_prompt): + self.positive_schedule, scheduled = get_prompt_schedule(prompt, self.steps) + self.negative_schedule, neg_scheduled = get_prompt_schedule(negative_prompt, self.steps) + self.scheduled_prompt = scheduled or neg_scheduled + + def scheduled_encode(self, pipe, batchidx): + prompt_dict = {} + for i in range(max(len(self.positive_schedule), len(self.negative_schedule))): + positive_prompt = self.positive_schedule[i % len(self.positive_schedule)] + negative_prompt = self.negative_schedule[i % len(self.negative_schedule)] + # skip repeated scheduled subprompts + idx = prompt_dict.get(positive_prompt+negative_prompt) + if idx is not None: + self.extend_embeds(batchidx, idx) + continue + self.encode(pipe, positive_prompt, negative_prompt, batchidx) + prompt_dict[positive_prompt+negative_prompt] = i + + def extend_embeds(self, batchidx, idx): + self.prompt_embeds[batchidx].append(self.prompt_embeds[batchidx][idx]) + self.negative_embeds[batchidx].append(self.negative_embeds[batchidx][idx]) + if len(self.positive_pooleds[batchidx]) > 0: + 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]) + + def duplicate_embeds(self): + self.prompt_embeds = self.prompt_embeds[0] * self.batchsize + self.positive_pooleds = self.positive_pooleds[0] * self.batchsize + self.negative_embeds = self.negative_embeds[0] * self.batchsize + self.negative_pooleds = self.negative_pooleds[0] * self.batchsize + + def encode(self, pipe, positive_prompt, negative_prompt, batchidx): + if shared.opts.prompt_attention == "xhinker parser" or 'Flux' in pipe.__class__.__name__: + prompt_embed, positive_pooled, negative_embed, negative_pooled = 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) + if prompt_embed is not None: + self.prompt_embeds[batchidx].append(prompt_embed) + if negative_embed is not None: + self.negative_embeds[batchidx].append(negative_embed) + if positive_pooled is not None: + self.positive_pooleds[batchidx].append(positive_pooled) + if negative_pooled is not None: + self.negative_pooleds[batchidx].append(negative_pooled) + + if debug_enabled: + get_tokens('positive', positive_prompt) + get_tokens('negative', negative_prompt) + pipe = prepare_model() + + def __call__(self, key, step=0): + batch = getattr(self, key) + res = [] + for embed in batch: + if len(embed) == 0: + return None + if len(embed) == 1: + res.append(embed[0]) + else: + res.append(embed[step]) + return torch.stack(res) + + def compel_hijack(self, token_ids: torch.Tensor, attention_mask: typing.Optional[torch.Tensor] = None) -> torch.Tensor: if not devices.same_device(self.text_encoder.device, devices.device): sd_models.move_model(self.text_encoder, devices.device)