From e0e243509dc2a32761d4927f55e9a0e7e8831ba8 Mon Sep 17 00:00:00 2001 From: silveroxides Date: Tue, 28 Apr 2026 10:17:42 +0200 Subject: [PATCH] Add automatic parsing for segmentation of prompt based on syntax --- __init__.py | 7 +++ parser.py | 50 +++++++++++++++++++ smart_nodes.py | 132 +++++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 189 insertions(+) create mode 100644 parser.py create mode 100644 smart_nodes.py diff --git a/__init__.py b/__init__.py index ba51731..b3c39c1 100644 --- a/__init__.py +++ b/__init__.py @@ -1,4 +1,11 @@ from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS +from .smart_nodes import PromptRelaySmartEncode, PromptRelaySmartEncodeTest + +NODE_CLASS_MAPPINGS["PromptRelaySmartEncode"] = PromptRelaySmartEncode +NODE_CLASS_MAPPINGS["PromptRelaySmartEncodeTest"] = PromptRelaySmartEncodeTest + +NODE_DISPLAY_NAME_MAPPINGS["PromptRelaySmartEncode"] = "Prompt Relay Encode (Smart)" +NODE_DISPLAY_NAME_MAPPINGS["PromptRelaySmartEncodeTest"] = "Prompt Relay Smart Encode Test" WEB_DIRECTORY = "./web" diff --git a/parser.py b/parser.py new file mode 100644 index 0000000..c961546 --- /dev/null +++ b/parser.py @@ -0,0 +1,50 @@ +import re + +def parse_smart_prompt(text): + block_pattern = r"(?im)^[a-z]+\s+([\d\.]+)(?:[:\-]([\d\.]+))?\s*:\s*$" + super_segments = text.split('|') + + parsed_segments = [] + + for super_seg in super_segments: + parts = re.split(block_pattern, super_seg) + current_text = parts[0] + + if current_text.strip(): + parsed_segments.append(_extract_inline(current_text, None)) + + for i in range(1, len(parts), 3): + val1 = float(parts[i]) + val2_str = parts[i+1] + text_part = parts[i+2] + + if val2_str is not None: + weight = float(val2_str) - val1 + else: + weight = val1 + + parsed_segments.append(_extract_inline(text_part, weight)) + + for seg in parsed_segments: + if not seg["text"]: + seg["text"] = " " + + return parsed_segments + +def _extract_inline(text, default_weight): + inline_pattern = r"\[([\d\.]+)(?:[:\-]([\d\.]+))?\]" + match = re.search(inline_pattern, text) + weight = 1.0 + if default_weight is not None: + weight = default_weight + + if match: + val1 = float(match.group(1)) + val2_str = match.group(2) + if val2_str is not None: + weight = float(val2_str) - val1 + else: + weight = val1 + text = re.sub(inline_pattern, "", text) + + return {"text": text.strip(), "weight": weight} diff --git a/smart_nodes.py b/smart_nodes.py new file mode 100644 index 0000000..c3e5629 --- /dev/null +++ b/smart_nodes.py @@ -0,0 +1,132 @@ +import logging +from comfy_api.latest import io +from .nodes import _encode_relay +from .prompt_relay import get_raw_tokenizer +from .parser import parse_smart_prompt + +log = logging.getLogger(__name__) + +class PromptRelaySmartEncode(io.ComfyNode): + """Parses advanced syntax into Prompt Relay segments and lengths.""" + + @classmethod + def define_schema(cls): + return io.Schema( + node_id="PromptRelaySmartEncode", + display_name="Prompt Relay Encode (Smart)", + category="conditioning/prompt_relay", + description="Parses syntax like [0-50] or block headers (Second 1:) to automatically calculate segment lengths.", + inputs=[ + io.Model.Input("model"), + io.Clip.Input("clip"), + io.Latent.Input("latent"), + io.String.Input("global_prompt", multiline=True, default=""), + io.String.Input("smart_prompt", multiline=True, default=""), + io.Bool.Input("normalize_by_tokens", default=False, tooltip="If true, scales the calculated length of each segment by its token count."), + io.Float.Input("epsilon", default=1e-3, min=1e-6, max=0.99, step=1e-4), + ], + outputs=[ + io.Model.Output(display_name="model"), + io.Conditioning.Output(display_name="positive"), + ], + ) + + @classmethod + def execute(cls, model, clip, latent, global_prompt, smart_prompt, normalize_by_tokens, epsilon) -> io.NodeOutput: + parsed = parse_smart_prompt(smart_prompt) + + valid_segments = [s for s in parsed if s["text"].strip()] + if not valid_segments: + valid_segments = [{"text": " ", "weight": 1.0}] + + raw_tokenizer = get_raw_tokenizer(clip) if normalize_by_tokens else None + + local_prompts_list = [] + weights_list = [] + + for seg in valid_segments: + text = seg["text"] + weight = seg["weight"] + + if normalize_by_tokens and raw_tokenizer: + try: + tokens = raw_tokenizer(text)["input_ids"] + has_eos = getattr(raw_tokenizer, "add_eos", False) + token_count = len(tokens) - (1 if has_eos else 0) + token_count = max(1, token_count) + weight *= token_count + except Exception as e: + log.warning(f"Token counting failed for segment '{text}': {e}") + + local_prompts_list.append(text) + weights_list.append(weight) + + local_prompts_str = " | ".join(local_prompts_list) + + scale_factor = 100000.0 + segment_lengths_str = ", ".join(str(int(w * scale_factor)) for w in weights_list) + + patched, conditioning = _encode_relay( + model, clip, latent, global_prompt, local_prompts_str, segment_lengths_str, epsilon + ) + + return io.NodeOutput(patched, conditioning) + + +class PromptRelaySmartEncodeTest(io.ComfyNode): + """Test node for Prompt Relay Smart Encode syntax parsing.""" + + @classmethod + def define_schema(cls): + return io.Schema( + node_id="PromptRelaySmartEncodeTest", + display_name="Prompt Relay Smart Encode Test", + category="conditioning/prompt_relay", + description="Outputs the parsed syntax for testing purposes.", + inputs=[ + io.String.Input("smart_prompt", multiline=True, default=""), + io.Bool.Input("normalize_by_tokens", default=False), + io.Clip.Input("clip", optional=True), + ], + outputs=[ + io.String.Output(display_name="parsed_output"), + ], + ) + + @classmethod + def execute(cls, smart_prompt, normalize_by_tokens, clip=None) -> io.NodeOutput: + parsed = parse_smart_prompt(smart_prompt) + + valid_segments = [s for s in parsed if s["text"].strip()] + if not valid_segments: + valid_segments = [{"text": " ", "weight": 1.0}] + + raw_tokenizer = None + if normalize_by_tokens and clip is not None: + from .prompt_relay import get_raw_tokenizer + raw_tokenizer = get_raw_tokenizer(clip) + + output_lines = [] + for i, seg in enumerate(valid_segments): + text = seg["text"] + weight = seg["weight"] + + base_weight = weight + token_count = None + + if normalize_by_tokens and raw_tokenizer: + try: + tokens = raw_tokenizer(text)["input_ids"] + has_eos = getattr(raw_tokenizer, "add_eos", False) + token_count = len(tokens) - (1 if has_eos else 0) + token_count = max(1, token_count) + weight *= token_count + except Exception: + pass + + line = f"Segment {i+1}: text='{text}', base_weight={base_weight}" + if token_count is not None: + line += f", tokens={token_count}, final_weight={weight}" + output_lines.append(line) + + return io.NodeOutput("\n".join(output_lines))