From 4fa421ef90b5573bc0341133dc4d657ac477b160 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 21 May 2024 16:34:38 -0400 Subject: [PATCH] add prompt caching and tokenizer info --- CHANGELOG.md | 3 ++ modules/processing.py | 11 +------ modules/prompt_parser_diffusers.py | 48 +++++++++++++++++++++++++++++- modules/sd_hijack.py | 2 +- modules/ui_common.py | 23 +++++++------- 5 files changed, 63 insertions(+), 24 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index ecfd4a0ca..9728fa4c9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -123,6 +123,9 @@ Can improve memory utilization on compatible GPUs (RTX and newer) - Torch dynamic profiling You can enable/disable full torch profiling in settings top menu on-the-fly + - Prompt caching - if you use the same prompt multiple times, no need to re-parse and encode it + Useful for batches as prompt processing is ~0.1sec on each pass + - Enhance `SD_PROMPT_DEBUG` to show actual tokens used - Support controlnet manually downloads models in both standalone and diffusers format For standalone, simply copy safetensors file to `models/control/controlnet` folder For diffusers format, create folder with model name in `models/control/controlnet/` diff --git a/modules/processing.py b/modules/processing.py index ff38de4fe..ae81ce833 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -114,6 +114,7 @@ class Processed: return create_infotext(p, self.all_prompts, self.all_seeds, self.all_subseeds, comments=[], position_in_batch=index % self.batch_size, iteration=index // self.batch_size) + def process_images(p: StableDiffusionProcessing) -> Processed: debug(f'Process images: {vars(p)}') if not hasattr(p.sd_model, 'sd_checkpoint_info'): @@ -212,16 +213,6 @@ def process_images(p: StableDiffusionProcessing) -> Processed: def process_init(p: StableDiffusionProcessing): seed = get_fixed_seed(p.seed) subseed = get_fixed_seed(p.subseed) - """ - if type(p.prompt) == list: - p.all_prompts = [shared.prompt_styles.apply_styles_to_prompt(x, p.styles) for x in p.prompt] - else: - p.all_prompts = p.batch_size * p.n_iter * [shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles)] - if type(p.negative_prompt) == list: - p.all_negative_prompts = [shared.prompt_styles.apply_negative_styles_to_prompt(x, p.styles) for x in p.negative_prompt] - else: - p.all_negative_prompts = p.batch_size * p.n_iter * [shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles)] - """ reset_prompts = False if p.all_prompts is None: p.all_prompts = p.prompt if isinstance(p.prompt, list) else p.batch_size * p.n_iter * [p.prompt] diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 6d2a7d2f6..72602d4fd 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -8,9 +8,14 @@ from transformers import PreTrainedTokenizer from modules import shared, prompt_parser, devices +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 debug('Trace: PROMPT') orig_encode_token_ids_to_embeddings = EmbeddingsProvider._encode_token_ids_to_embeddings # pylint: disable=protected-access +token_dict = None +token_type = None +cache = {} +cache_type = None def compel_hijack(self, token_ids: torch.Tensor, @@ -97,10 +102,41 @@ def get_prompt_schedule(prompt, steps): return temp, len(schedule) > 1 +def get_tokens(msg, prompt): + global token_dict, token_type # pylint: disable=global-statement + if shared.backend != shared.Backend.DIFFUSERS: + return + if shared.sd_loaded and hasattr(shared.sd_model, 'tokenizer') and shared.sd_model.tokenizer is not None: + if token_dict is None or token_type != shared.sd_model_type: + token_type = shared.sd_model_type + fn = os.path.join(shared.sd_model.tokenizer.name_or_path, 'tokenizer', 'vocab.json') + token_dict = shared.readfile(fn, silent=True) + for k, v in shared.sd_model.tokenizer.added_tokens_decoder.items(): + token_dict[str(v)] = k + shared.log.debug(f'Tokenizer: words={len(token_dict)} file="{fn}"') + has_bos_token = shared.sd_model.tokenizer.bos_token_id is not None + has_eos_token = shared.sd_model.tokenizer.eos_token_id is not None + ids = shared.sd_model.tokenizer(prompt) + ids = getattr(ids, 'input_ids', []) + tokens = [] + for i in ids: + tokens.append(list(token_dict.keys())[list(token_dict.values()).index(i)]) + token_count = len(ids) - int(has_bos_token) - int(has_eos_token) + shared.log.trace(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): if 'StableDiffusion' not in pipe.__class__.__name__ and 'DemoFusion' 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: + 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() positive_schedule, scheduled = get_prompt_schedule(prompts[0], steps) @@ -111,7 +147,6 @@ def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, c p.negative_embeds = [] p.negative_pooleds = [] - cache = {} 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)] @@ -124,12 +159,23 @@ def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, c prompt_embed, positive_pooled, negative_embed, negative_pooled = results if prompt_embed is not None: p.prompt_embeds.append(torch.cat([prompt_embed] * len(prompts), dim=0)) + cache['prompt_embeds'] = p.prompt_embeds if negative_embed is not None: p.negative_embeds.append(torch.cat([negative_embed] * len(negative_prompts), dim=0)) + cache['negative_embeds'] = p.negative_embeds if positive_pooled is not None: p.positive_pooleds.append(torch.cat([positive_pooled] * len(prompts), dim=0)) + cache['positive_pooleds'] = p.positive_pooleds if negative_pooled is not None: p.negative_pooleds.append(torch.cat([negative_pooled] * len(negative_prompts), dim=0)) + cache['negative_pooleds'] = p.negative_pooleds + + cache['prompts'] = prompts + cache['negative_prompts'] = negative_prompts + cache['model_type'] = shared.sd_model_type + if debug_enabled: + get_tokens('positive', prompts[0]) + get_tokens('negative', negative_prompts[0]) debug(f"Prompt encode: time={(time.time() - t0):.3f}") return diff --git a/modules/sd_hijack.py b/modules/sd_hijack.py index 6360982b8..53e4272ef 100644 --- a/modules/sd_hijack.py +++ b/modules/sd_hijack.py @@ -252,7 +252,7 @@ class StableDiffusionModelHijack: def get_prompt_lengths(self, text): if self.clip is None: return 0, 0 - _, token_count = self.clip.process_texts([text]) + chunks, token_count = self.clip.process_texts([text]) return token_count, self.clip.get_target_prompt_token_count(token_count) diff --git a/modules/ui_common.py b/modules/ui_common.py index 77ce9e6a9..1b8aea378 100644 --- a/modules/ui_common.py +++ b/modules/ui_common.py @@ -382,24 +382,23 @@ def update_token_counter(text, steps): shared.log.info('Tokenizer busy') return f"{token_count}/{max_length}" from modules import extra_networks - try: - text, _ = extra_networks.parse_prompt(text) - _, prompt_flat_list, _ = prompt_parser.get_multicond_prompt_list([text]) - prompt_schedules = prompt_parser.get_learned_conditioning_prompt_schedules(prompt_flat_list, steps) - except Exception: - prompt_schedules = [[[steps, text]]] - flat_prompts = reduce(lambda list1, list2: list1+list2, prompt_schedules) - prompts = [prompt_text for step, prompt_text in flat_prompts] + prompt, _ = extra_networks.parse_prompt(text) if shared.backend == shared.Backend.ORIGINAL: from modules import sd_hijack + try: + _, prompt_flat_list, _ = prompt_parser.get_multicond_prompt_list([text]) + prompt_schedules = prompt_parser.get_learned_conditioning_prompt_schedules(prompt_flat_list, steps) + except Exception: + prompt_schedules = [[[steps, text]]] + flat_prompts = reduce(lambda list1, list2: list1+list2, prompt_schedules) + prompts = [prompt_text for _step, prompt_text in flat_prompts] token_count, max_length = max([sd_hijack.model_hijack.get_prompt_lengths(prompt) for prompt in prompts], key=lambda args: args[0]) elif shared.backend == shared.Backend.DIFFUSERS: if shared.sd_loaded and hasattr(shared.sd_model, 'tokenizer') and shared.sd_model.tokenizer is not None: has_bos_token = shared.sd_model.tokenizer.bos_token_id is not None has_eos_token = shared.sd_model.tokenizer.eos_token_id is not None - ids = [shared.sd_model.tokenizer(prompt) for prompt in prompts] - if len(ids) > 0 and hasattr(ids[0], 'input_ids'): - ids = [x.input_ids for x in ids] - token_count = max([len(x) for x in ids]) - int(has_bos_token) - int(has_eos_token) + ids = shared.sd_model.tokenizer(prompt) + ids = getattr(ids, 'input_ids', []) + token_count = len(ids) - int(has_bos_token) - int(has_eos_token) max_length = shared.sd_model.tokenizer.model_max_length - int(has_bos_token) - int(has_eos_token) return f"{token_count}/{max_length}"