From f3442abc929c966581a48e74cb4e782bfda2a39d Mon Sep 17 00:00:00 2001 From: AI-Casanova <54461896+AI-Casanova@users.noreply.github.com> Date: Sat, 26 Oct 2024 18:47:11 -0500 Subject: [PATCH] Prompt LRU Cache --- modules/extra_networks.py | 2 +- modules/prompt_parser_diffusers.py | 44 +++++++++++++++++++++++++++--- modules/shared.py | 3 +- 3 files changed, 43 insertions(+), 6 deletions(-) diff --git a/modules/extra_networks.py b/modules/extra_networks.py index a574e8469..673549b6b 100644 --- a/modules/extra_networks.py +++ b/modules/extra_networks.py @@ -102,8 +102,8 @@ def activate(p, extra_network_data, step=0): except Exception as e: errors.display(e, f"Activating network: type={extra_network_name}") + p.extra_network_data = extra_network_data if stepwise: - p.extra_network_data = extra_network_data shared.opts.data['lora_functional'] = functional diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 42a6bcbbd..ecf0cebd5 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -3,11 +3,11 @@ import math import time import typing import torch +from collections import OrderedDict from compel.embeddings_provider import BaseTextualInversionManager, EmbeddingsProvider 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 @@ -15,7 +15,7 @@ debug('Trace: PROMPT') orig_encode_token_ids_to_embeddings = EmbeddingsProvider._encode_token_ids_to_embeddings # pylint: disable=protected-access token_dict = None # used by helper get_tokens token_type = None # used by helper get_tokens -cache = {} +cache = OrderedDict() def prompt_compatible(): @@ -57,6 +57,9 @@ class PromptEmbedder: self.positive_schedule = None self.negative_schedule = None self.scheduled_prompt = False + earlyout = self.checkcache(p) + if earlyout: + return pipe = prepare_model() # per prompt in batch for batchidx, (prompt, negative_prompt) in enumerate(zip(self.prompts, self.negative_prompts)): @@ -66,8 +69,41 @@ class PromptEmbedder: else: self.encode(pipe, prompt, negative_prompt, batchidx) if self.allsame: - self.duplicate_embeds() + self.fix_batch_embeds() debug(f"Prompt encode: time={(time.time() - t0):.3f}") + self.checkcache(p) + + 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]) + 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]]): + return False + else: + cache[key] = {'prompt_embeds': self.prompt_embeds, + 'negative_embeds': self.negative_embeds, + 'positive_pooleds': self.positive_pooleds, + 'negative_pooleds': self.negative_pooleds, + } + debug(f"Prompt cache: Adding {key}") + while len(cache) > int(shared.opts.sd_textencoder_cache_size): + cache.popitem(last=False) + if item: + self.__dict__.update(cache[key]) + cache.move_to_end(key) + debug(f"Prompt cache: Retrieving {key}") + return True def compare_prompts(self): same = (self.prompts == [self.prompts[0]] * len(self.prompts) and @@ -103,7 +139,7 @@ class PromptEmbedder: if len(self.negative_pooleds[batchidx]) > 0: self.negative_pooleds[batchidx].append(self.negative_pooleds[batchidx][idx]) - def duplicate_embeds(self): + 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 diff --git a/modules/shared.py b/modules/shared.py index f7be44390..4622f1d2c 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -435,7 +435,8 @@ options_templates.update(options_section(('sd', "Execution & Models"), { "sd_model_dict": OptionInfo('None', "Use separate base dict", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_tiles()}, refresh=refresh_checkpoints), "sd_checkpoint_autoload": OptionInfo(True, "Model autoload on start"), "sd_checkpoint_autodownload": OptionInfo(True, "Model auto-download on demand"), - "sd_textencoder_cache": OptionInfo(True, "Cache text encoder results"), + "sd_textencoder_cache": OptionInfo(True, "Cache text encoder results", gr.Checkbox, {"visible": False}), + "sd_textencoder_cache_size": OptionInfo(4, "Text encoder results LRU cache size", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1}), "stream_load": OptionInfo(False, "Load models using stream loading method", gr.Checkbox, {"visible": not native }), "model_reuse_dict": OptionInfo(False, "Reuse loaded model dictionary", gr.Checkbox, {"visible": False}), "prompt_mean_norm": OptionInfo(False, "Prompt attention normalization", gr.Checkbox),