Add automatic parsing for segmentation of prompt based on syntax
This commit is contained in:
@@ -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"
|
||||
|
||||
|
||||
@@ -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}
|
||||
+132
@@ -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))
|
||||
Reference in New Issue
Block a user