add prompt caching and tokenizer info

This commit is contained in:
Vladimir Mandic
2024-05-21 16:34:38 -04:00
parent b404b0354b
commit 4fa421ef90
5 changed files with 63 additions and 24 deletions
+1 -10
View File
@@ -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]
+47 -1
View File
@@ -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
+1 -1
View File
@@ -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)
+11 -12
View File
@@ -382,24 +382,23 @@ def update_token_counter(text, steps):
shared.log.info('Tokenizer busy')
return f"<span class='gr-box gr-text-input'>{token_count}/{max_length}</span>"
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"<span class='gr-box gr-text-input'>{token_count}/{max_length}</span>"