diff --git a/modules/extra_networks.py b/modules/extra_networks.py index 054bc5c2b..a3cd8936a 100644 --- a/modules/extra_networks.py +++ b/modules/extra_networks.py @@ -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 diff --git a/modules/ui_common.py b/modules/ui_common.py index b96e0daff..3b43ea566 100644 --- a/modules/ui_common.py +++ b/modules/ui_common.py @@ -427,7 +427,10 @@ def update_token_counter(text): shared.log.debug('Tokenizer busy') return f"{token_count}/{max_length}" 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