From 79d665e5dabe152a407c9e6c9c31dcaff98a01da Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 1 Apr 2025 22:16:32 -0400 Subject: [PATCH] prompt enhance lora extract/recombine Signed-off-by: Vladimir Mandic --- scripts/prompt_enhance.py | 27 ++++++++++++++++++--------- 1 file changed, 18 insertions(+), 9 deletions(-) diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index 292faae64..bbc52b0be 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -8,7 +8,7 @@ from modules import scripts, shared, devices, errors, processing, sd_models, sd_ debug_enabled = os.environ.get('SD_LLM_DEBUG', None) is not None -debug = shared.log.trace if debug_enabled else lambda *args, **kwargs: None +debug_log = shared.log.trace if debug_enabled else lambda *args, **kwargs: None @dataclass @@ -79,8 +79,7 @@ class Script(scripts.Script): gguf_args = {} if model_type is not None and model_file is not None and len(model_type) > 2 and len(model_file) > 2: - if debug: - shared.log.trace(f'Prompt enhance: gguf supported={self.options.supported}') + debug_log(f'Prompt enhance: gguf supported={self.options.supported}') if model_type not in self.options.supported: shared.log.error(f'Prompt enhance: name="{name}" repo="{model_repo}" fn="{model_file}" type={model_type} gguf not supported') shared.log.trace(f'Prompt enhance: gguf supported={self.options.supported}') @@ -112,7 +111,7 @@ class Script(scripts.Script): pretrained_model_name_or_path=model_repo, cache_dir=shared.opts.hfcache_dir, ) - if debug: + if debug_enabled: modules = sd_modules.get_model_stats(self.llm) + sd_modules.get_model_stats(self.tokenizer) for m in modules: shared.log.trace(f'Prompt enhance: {m}') @@ -158,7 +157,7 @@ class Script(scripts.Script): if 0 <= response.find('---') < self.options.max_delim_index: response, removed = response.split('---', maxsplit=1) if len(removed) > 0: - debug(f'Prompt enhance: max={self.options.max_delim_index} removed="{removed}"') + debug_log(f'Prompt enhance: max={self.options.max_delim_index} removed="{removed}"') # remove bullets and lists lines = [re.sub(r'^(\s*[-*]|\s*\d+)\s+', '', line).strip() for line in response.splitlines()] @@ -167,7 +166,7 @@ class Script(scripts.Script): response = response.strip() return response - def preprocess(self, response, prefix, suffix): + def post(self, response, prefix, suffix, networks): response = response.strip() prefix = prefix.strip() suffix = suffix.strip() @@ -175,8 +174,16 @@ class Script(scripts.Script): response = f'{prefix} {response}' if len(suffix) > 0: response = f'{response} {suffix}' + if len(networks) > 0: + response = f'{response} {" ".join(networks)}' return response + def extract(self, prompt): + pattern = r'(<.*?>)' + matches = re.findall(pattern, prompt) + filtered = re.sub(pattern, '', prompt) + return filtered, matches + def enhance(self, model: str=None, prompt:str=None, system:str=None, prefix:str=None, suffix:str=None, sample:bool=None, tokens:int=None, temperature:float=None, penalty:float=None): model = model or self.options.default prompt = prompt or self.prompt.value @@ -193,6 +200,8 @@ class Script(scripts.Script): if self.llm is None: shared.log.error('Prompt enhance: model not loaded') return prompt + prompt, networks = self.extract(prompt) + debug_log(f'Prompt enhance: networks={networks}') chat_template = [ { "role": "system", "content": system }, { "role": "user", "content": prompt }, @@ -226,7 +235,7 @@ class Script(scripts.Script): if shared.opts.diffusers_offload_mode != 'none': sd_models.move_model(self.llm, devices.cpu) devices.torch_gc() - if debug: + if debug_enabled: raw_response = self.tokenizer.batch_decode(outputs, skip_special_tokens=True, clean_up_tokenization_spaces=True) shared.log.trace(f'Prompt enhance: raw="{raw_response}"') outputs_cropped = outputs[:, input_len:] @@ -247,9 +256,9 @@ class Script(scripts.Script): is_censored = self.censored(response) if not is_censored: response = self.clean(response) - response = self.preprocess(response, prefix, suffix) + response = self.post(response, prefix, suffix, networks) shared.log.info(f'Prompt enhance: model="{model}" time={t1-t0:.2f} inputs={input_len} outputs={outputs.shape[-1]} prompt={len(prompt)} response={len(response)}') - if debug: + if debug_enabled: shared.log.trace(f'Prompt enhance: sample={sample} tokens={tokens} temperature={temperature} penalty={penalty}') shared.log.trace(f'Prompt enhance: prompt="{prompt}"') shared.log.trace(f'Prompt enhance: response="{response}"')