From ddac624ad198b09e527ca058e504b384bcbbd68e Mon Sep 17 00:00:00 2001 From: asagi4 <130366179+asagi4@users.noreply.github.com> Date: Mon, 2 Jun 2025 17:01:44 +0300 Subject: [PATCH] Add NBREAK (Warning: unstable. Name is likely to change) NBREAK should have the same behaviour as ComfyUI's ConditioningConcat See #111 --- prompt_control/prompts.py | 58 +++++++++++++++++++++++++-------------- 1 file changed, 38 insertions(+), 20 deletions(-) diff --git a/prompt_control/prompts.py b/prompt_control/prompts.py index 67fd9ae..689e121 100644 --- a/prompt_control/prompts.py +++ b/prompt_control/prompts.py @@ -147,23 +147,7 @@ def tokenize_chunks(clip, text, need_word_ids): return tokens -def encode_prompt_segment( - clip, - text, - settings, - default_style="comfy", - default_normalization="none", - clip_weights=None, -) -> list[tuple[torch.Tensor, dict[str]]]: - style, normalization, text = get_style(text, default_style, default_normalization) - clip_weights, text = get_clipweights(text, clip_weights) - text, cuts = parse_cuts(text) - extra = {} - if clip_weights: - extra["clip_weights"] = clip_weights - if cuts: - extra["cuts"] = cuts - +def tokenize(clip, text): # defaults=None means there is no argument parsing at all text, l_prompts = get_function(text, "CLIP_L", defaults=None) text, te_prompts = get_function(text, "TE", defaults=None) @@ -207,18 +191,52 @@ def encode_prompt_segment( empty = clip.tokenize("", return_word_ids=need_word_ids) tokens[k] += empty[k] - tokens = fix_word_ids(tokens) + return fix_word_ids(tokens) + +def encode_prompt_segment( + clip, + text, + settings, + default_style="comfy", + default_normalization="none", + clip_weights=None, +) -> list[tuple[torch.Tensor, dict[str]]]: + style, normalization, text = get_style(text, default_style, default_normalization) + clip_weights, text = get_clipweights(text, clip_weights) + text, cuts = parse_cuts(text) + extra = {} + if clip_weights: + extra["clip_weights"] = clip_weights + if cuts: + extra["cuts"] = cuts + + empty = clip.tokenize("") tes = [] - for k in tokens: + for k in empty: if k in ["g", "l"]: tes.append(f"clip_{k}") else: tes.append(k) clip = hook_te(clip, tes, style, normalization, extra) + chunks = re.split(r"\bNBREAK\b", text) + conds = [] + for c in chunks: + tokens = tokenize(clip, c) + conds.append(clip.encode_from_tokens_scheduled(tokens, add_dict=settings)) - return clip.encode_from_tokens_scheduled(tokens, add_dict=settings) + count = len(conds[0]) + base = conds[0] + for cond in conds[1:]: + assert len(cond) == count, "Conditioning length mismatch" + # Pooled gets ignored + for i in range(count): + c1 = base[i][0] + c2 = cond[i][0] + base[i][0] = torch.cat((c1, c2), 1) + + return base def apply_weights(output, te_name, spec):