mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
Rework prompt parsing/processing
- Return consistent structure
This commit is contained in:
+14
-20
@@ -151,33 +151,27 @@ def deactivate(p, extra_network_data=None, force=shared.opts.lora_force_reload):
|
||||
re_extra_net = re.compile(r"<(\w+):([^>]+)>")
|
||||
|
||||
|
||||
def parse_prompt(prompt):
|
||||
res = defaultdict(list)
|
||||
def parse_prompt(prompt: str | None) -> tuple[str, defaultdict[str, list[ExtraNetworkParams]]]:
|
||||
res: defaultdict[str, list[ExtraNetworkParams]] = defaultdict(list)
|
||||
if prompt is None:
|
||||
return prompt, res
|
||||
return "", res
|
||||
|
||||
def found(m):
|
||||
name = m.group(1)
|
||||
args = m.group(2)
|
||||
def found(m: re.Match[str]):
|
||||
name, args = m.group(1, 2)
|
||||
res[name].append(ExtraNetworkParams(items=args.split(":")))
|
||||
return ""
|
||||
if isinstance(prompt, list):
|
||||
prompt = [re.sub(re_extra_net, found, p) for p in prompt]
|
||||
else:
|
||||
prompt = re.sub(re_extra_net, found, prompt)
|
||||
return prompt, res
|
||||
|
||||
updated_prompt = re.sub(re_extra_net, found, prompt)
|
||||
return updated_prompt, res
|
||||
|
||||
|
||||
def parse_prompts(prompts):
|
||||
res = []
|
||||
extra_data = None
|
||||
if prompts is None:
|
||||
return prompts, extra_data
|
||||
|
||||
def parse_prompts(prompts: list[str]):
|
||||
updated_prompt_list: list[str] = []
|
||||
extra_data: defaultdict[str, list[ExtraNetworkParams]] = defaultdict(list)
|
||||
for prompt in prompts:
|
||||
updated_prompt, parsed_extra_data = parse_prompt(prompt)
|
||||
if extra_data is None:
|
||||
if not extra_data:
|
||||
extra_data = parsed_extra_data
|
||||
res.append(updated_prompt)
|
||||
updated_prompt_list.append(updated_prompt)
|
||||
|
||||
return res, extra_data
|
||||
return updated_prompt_list, extra_data
|
||||
|
||||
@@ -427,7 +427,10 @@ def update_token_counter(text):
|
||||
shared.log.debug('Tokenizer busy')
|
||||
return f"<span class='gr-box gr-text-input'>{token_count}/{max_length}</span>"
|
||||
from modules import extra_networks
|
||||
prompt, _ = extra_networks.parse_prompt(text)
|
||||
if isinstance(text, list):
|
||||
prompt, _ = extra_networks.parse_prompts(text)
|
||||
else:
|
||||
prompt, _ = extra_networks.parse_prompt(text)
|
||||
if shared.sd_loaded and hasattr(shared.sd_model, 'tokenizer') and shared.sd_model.tokenizer is not None:
|
||||
tokenizer = shared.sd_model.tokenizer
|
||||
# For multi-modal processors (e.g., PixtralProcessor), use the underlying text tokenizer
|
||||
|
||||
Reference in New Issue
Block a user