diff --git a/modules/processing_args.py b/modules/processing_args.py index 8dd47bd87..34dd97a1c 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -117,7 +117,6 @@ 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) p.embedder = prompt_parser_diffusers.PromptEmbedder(prompts, negative_prompts, clip_skip, p) parser = shared.opts.prompt_attention except Exception as e: @@ -143,7 +142,7 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2 args['prompt'] = prompts if 'negative_prompt' in possible: 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') + args['negative_prompt_embeds'] = p.embedder('negative_prompt_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: diff --git a/modules/processing_callbacks.py b/modules/processing_callbacks.py index 2584ee796..5c24aead0 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.embedder is not None: + try: + if 'prompt_embeds' in kwargs: + kwargs["prompt_embeds"] = p.embedder("prompt_embeds", step + 1) + if 'negative_prompt_embeds' in kwargs: + kwargs["negative_prompt_embeds"] = p.embedder("negative_prompt_embeds", step + 1) + 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/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index ecf0cebd5..5ba0e8a74 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -2,8 +2,8 @@ import os import math import time import typing -import torch from collections import OrderedDict +import torch from compel.embeddings_provider import BaseTextualInversionManager, EmbeddingsProvider from transformers import PreTrainedTokenizer from modules import shared, prompt_parser, devices, sd_models @@ -43,16 +43,16 @@ def prepare_model(): 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.allsame = self.compare_prompts() # collapses batched prompts to single prompt if possible self.steps = p.steps self.clip_skip = clip_skip + # All embeds are nested lists, outer list batch length, inner schedule length self.prompt_embeds = [[]] * self.batchsize self.positive_pooleds = [[]] * self.batchsize - self.negative_embeds = [[]] * self.batchsize + self.negative_prompt_embeds = [[]] * self.batchsize self.negative_pooleds = [[]] * self.batchsize self.positive_schedule = None self.negative_schedule = None @@ -68,31 +68,31 @@ class PromptEmbedder: self.scheduled_encode(pipe, batchidx) else: self.encode(pipe, prompt, negative_prompt, batchidx) - if self.allsame: - self.fix_batch_embeds() - debug(f"Prompt encode: time={(time.time() - t0):.3f}") self.checkcache(p) + debug(f"Prompt encode: time={(time.time() - t0):.3f}") def checkcache(self, p): if shared.opts.sd_textencoder_cache_size == 0: return False + def flatten(xss): return [x for xs in xss for x in xs] # unpack EN data in case of TE LoRA en_data = p.extra_network_data en_data = [idx.items for item in en_data.values() for idx in item] - key = str([self.prompts, self.negative_prompts, self.batchsize, self.clip_skip, self.steps, en_data]) + effective_batch = 1 if self.allsame else self.batchsize + key = str([self.prompts, self.negative_prompts, effective_batch, self.clip_skip, self.steps, en_data]) item = cache.get(key) if not item: - if not any([flatten(emb) for emb in [self.prompt_embeds, - self.negative_embeds, - self.positive_pooleds, - self.negative_pooleds]]): + if not any(flatten(emb) for emb in [self.prompt_embeds, + self.negative_prompt_embeds, + self.positive_pooleds, + self.negative_pooleds]): return False else: cache[key] = {'prompt_embeds': self.prompt_embeds, - 'negative_embeds': self.negative_embeds, + 'negative_prompt_embeds': self.negative_prompt_embeds, 'positive_pooleds': self.positive_pooleds, 'negative_pooleds': self.negative_pooleds, } @@ -102,6 +102,11 @@ class PromptEmbedder: if item: self.__dict__.update(cache[key]) cache.move_to_end(key) + if self.allsame and len(self.prompt_embeds) < self.batchsize: # If current batch larger than cached + self.prompt_embeds = [self.prompt_embeds[0]] * self.batchsize + 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 debug(f"Prompt cache: Retrieving {key}") return True @@ -119,7 +124,7 @@ class PromptEmbedder: self.scheduled_prompt = scheduled or neg_scheduled def scheduled_encode(self, pipe, batchidx): - prompt_dict = {} + prompt_dict = {} # index cache 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)] @@ -131,20 +136,14 @@ class PromptEmbedder: self.encode(pipe, positive_prompt, negative_prompt, batchidx) prompt_dict[positive_prompt+negative_prompt] = i - def extend_embeds(self, batchidx, idx): + def extend_embeds(self, batchidx, idx): # Extends scheduled prompt via index self.prompt_embeds[batchidx].append(self.prompt_embeds[batchidx][idx]) - self.negative_embeds[batchidx].append(self.negative_embeds[batchidx][idx]) + self.negative_prompt_embeds[batchidx].append(self.negative_prompt_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 fix_batch_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( @@ -155,7 +154,7 @@ class PromptEmbedder: 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) + self.negative_prompt_embeds[batchidx].append(negative_embed) if positive_pooled is not None: self.positive_pooleds[batchidx].append(positive_pooled) if negative_pooled is not None: @@ -169,14 +168,14 @@ class PromptEmbedder: def __call__(self, key, step=0): batch = getattr(self, key) res = [] - for embed in batch: - if len(embed) == 0: + for i in range(self.batchsize): + if len(batch[i]) == 0: return None - if len(embed) == 1: - res.append(embed[0]) else: - res.append(embed[step]) - return torch.stack(res) + res.append(batch[i][step]) + if step != 0: # For Callback + res.append(batch[i][step]) # Diffusers internally doubles batch dimension + return torch.cat(res) def compel_hijack(self, token_ids: torch.Tensor, attention_mask: typing.Optional[torch.Tensor] = None) -> torch.Tensor: @@ -221,9 +220,9 @@ def insert_parser_highjack(pipename): debug("Load Standard Parser hijack") - insert_parser_highjack("Initialize") + # from https://github.com/damian0815/compel/blob/main/src/compel/diffusers_textual_inversion_manager.py class DiffusersTextualInversionManager(BaseTextualInversionManager): def __init__(self, pipe, tokenizer): @@ -270,12 +269,6 @@ class DiffusersTextualInversionManager(BaseTextualInversionManager): def get_prompt_schedule(prompt, steps): t0 = time.time() - if shared.native: - # TODO prompt scheduling - # prompt schedule returns array of prompts which would require that each prompt is fed to the model per-step - # prompt scheduling should instead interpolate between each prompt in schedule - # this temporarily disables prompt scheduling - return [prompt], False temp = [] schedule = prompt_parser.get_learned_conditioning_prompt_schedules([prompt], steps)[0] if all(x == schedule[0] for x in schedule): @@ -319,118 +312,6 @@ def get_tokens(msg, prompt): debug(f'Prompt tokenizer: type={msg} tokens={token_count} {tokens}') -def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, clip_skip: typing.Optional[int] = None): - params_match = prompts == cache.get('prompts', None) and negative_prompts == cache.get('negative_prompts', None) and clip_skip == cache.get('clip_skip', None) and steps == cache.get('steps', None) - if ( - 'StableDiffusion' not in pipe.__class__.__name__ and - 'DemoFusion' not in pipe.__class__.__name__ and - 'StableCascade' not in pipe.__class__.__name__ and - 'Flux' not in pipe.__class__.__name__ - ): - shared.log.warning(f"Prompt parser not supported: {pipe.__class__.__name__}") - return - elif shared.opts.sd_textencoder_cache and cache.get('model_type', None) == shared.sd_model_type and params_match: - p.prompt_embeds = cache.get('prompt_embeds', None) - p.positive_pooleds = cache.get('positive_pooleds', None) - p.negative_embeds = cache.get('negative_embeds', None) - p.negative_pooleds = cache.get('negative_pooleds', None) - p.scheduled_prompt = cache.get('scheduled_prompt', None) - debug("Prompt encode: cached") - return - else: - t0 = time.time() - 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() - - prompt_embeds, positive_pooleds, negative_embeds, negative_pooleds = [], [], [], [] - last_prompt, last_negative = None, None - for prompt, negative in zip(prompts, negative_prompts): - prompt_embed, positive_pooled, negative_embed, negative_pooled = None, None, None, None - if last_prompt == prompt and last_negative == negative: - prompt_embeds.append(prompt_embeds[-1]) - negative_embeds.append(negative_embeds[-1]) - if len(positive_pooleds) > 0: - positive_pooleds.append(positive_pooleds[-1]) - if len(negative_pooleds) > 0: - negative_pooleds.append(negative_pooleds[-1]) - continue - positive_schedule, scheduled = get_prompt_schedule(prompt, steps) - negative_schedule, neg_scheduled = get_prompt_schedule(negative, steps) - p.scheduled_prompt = scheduled or neg_scheduled - p.prompt_embeds = [] - p.positive_pooleds = [] - p.negative_embeds = [] - p.negative_pooleds = [] - - for i in range(max(len(positive_schedule), len(negative_schedule))): - positive_prompt = positive_schedule[i % len(positive_schedule)] - negative_prompt = negative_schedule[i % len(negative_schedule)] - 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, clip_skip) - else: - prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipe, positive_prompt, negative_prompt, clip_skip) - if prompt_embed is not None: - prompt_embeds.append(prompt_embed) - if negative_embed is not None: - negative_embeds.append(negative_embed) - if positive_pooled is not None: - positive_pooleds.append(positive_pooled) - if negative_pooled is not None: - negative_pooleds.append(negative_pooled) - last_prompt, last_negative = prompt, negative - # TODO prompt scheduling - # interpolation should happen here and then we can re-enable prompt scheduling - # ive tried simple torch.mean and its not good-enough - - def fix_length(embeds): - max_len = max([e.shape[1] for e in embeds if e is not None]) - for i, e in enumerate(embeds): - if e is not None and e.shape[1] < max_len: - expanded = torch.zeros((e.shape[0], max_len, e.shape[2]), device=e.device, dtype=e.dtype) - expanded[:, :e.shape[1], :] = e - embeds[i] = expanded - return torch.cat(embeds, dim=0).to(devices.device, dtype=devices.dtype) - - if len(prompt_embeds) > 0: - p.prompt_embeds.append(fix_length(prompt_embeds)) - if len(negative_embeds) > 0: - p.negative_embeds.append(fix_length(negative_embeds)) - if len(positive_pooleds) > 0: - p.positive_pooleds.append(fix_length(positive_pooleds)) - if len(negative_pooleds) > 0: - p.negative_pooleds.append(fix_length(negative_pooleds)) - - if shared.opts.sd_textencoder_cache and p.batch_size == 1: - cache.update({ - 'prompt_embeds': p.prompt_embeds, - 'negative_embeds': p.negative_embeds, - 'positive_pooleds': p.positive_pooleds, - 'negative_pooleds': p.negative_pooleds, - 'scheduled_prompt': p.scheduled_prompt, - 'prompts': prompts, - 'negative_prompts': negative_prompts, - 'clip_skip': clip_skip, - 'steps': steps, - 'model_type': shared.sd_model_type - }) - else: - cache.clear() - if debug_enabled: - get_tokens('positive', prompts[0]) - get_tokens('negative', negative_prompts[0]) - if shared.opts.diffusers_offload_mode == "balanced": - pipe = sd_models.apply_balanced_offload(pipe) - elif hasattr(pipe, "maybe_free_model_hooks"): - # text encoder will stay in the vram and cause oom, send everything back to cpu before continuing - pipe.maybe_free_model_hooks() - debug(f"Prompt encode: time={(time.time() - t0):.3f}") - devices.torch_gc() - return - - def normalize_prompt(pairs: list): num_words = 0 total_weight = 0 @@ -516,6 +397,7 @@ def pad_to_same_length(pipe, embeds, empty_embedding_providers=None): embeds[i] = embed return embeds + def split_prompts(prompt, SD3 = False): if prompt.find("TE2:") != -1: prompt, prompt2 = prompt.split("TE2:")