diff --git a/__init__.py b/__init__.py index b3c39c1..d47bbe5 100644 --- a/__init__.py +++ b/__init__.py @@ -1,11 +1,38 @@ from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS +from .nodes import PromptRelayEncode, PromptRelayEncodeTimeline from .smart_nodes import PromptRelaySmartEncode, PromptRelaySmartEncodeTest +from comfy_api.latest import ComfyExtension, io +from typing_extensions import override -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" +class PromptRelay(ComfyExtension): + @override + async def get_node_list(self) -> list[type[io.ComfyNode]]: + return [ + PromptRelayEncode, + PromptRelayEncodeTimeline, + PromptRelaySmartEncode, + PromptRelaySmartEncodeTest, + ] + + +async def comfy_entrypoint() -> PromptRelay: + return PromptRelay() + +NODE_CLASS_MAPPINGS = { + "PromptRelayEncode": PromptRelayEncode, + "PromptRelayEncodeTimeline": PromptRelayEncodeTimeline, + "PromptRelaySmartEncode": PromptRelaySmartEncode, + "PromptRelaySmartEncodeTest": PromptRelaySmartEncodeTest +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "PromptRelayEncode": "Prompt Relay Encode", + "PromptRelayEncodeTimeline": "Prompt Relay Encode (Timeline)", + "PromptRelaySmartEncode": "Prompt Relay Encode (Smart)", + "PromptRelaySmartEncodeTest": "Prompt Relay Smart Encode Test" +} + WEB_DIRECTORY = "./web" diff --git a/parser.py b/parser.py index c961546..581fc8a 100644 --- a/parser.py +++ b/parser.py @@ -13,15 +13,19 @@ def parse_smart_prompt(text): if current_text.strip(): parsed_segments.append(_extract_inline(current_text, None)) + # parts looks like: [text_before, val1, val2_opt, text_after, val1, val2_opt, text_after...] for i in range(1, len(parts), 3): val1 = float(parts[i]) val2_str = parts[i+1] text_part = parts[i+2] + # If it's a range like "Second 1-3:", weight is 3-1 = 2 if val2_str is not None: weight = float(val2_str) - val1 else: - weight = val1 + # If it's just "Second 1:", "Second 2:", we assume each step is an equal chunk + # so base weight is 1.0, not the index itself! + weight = 1.0 parsed_segments.append(_extract_inline(text_part, weight)) @@ -48,3 +52,4 @@ def _extract_inline(text, default_weight): text = re.sub(inline_pattern, "", text) return {"text": text.strip(), "weight": weight} + diff --git a/smart_nodes.py b/smart_nodes.py index 98cac9d..3707c88 100644 --- a/smart_nodes.py +++ b/smart_nodes.py @@ -20,12 +20,15 @@ class PromptRelaySmartEncode(io.ComfyNode): io.Model.Input("model"), io.Clip.Input("clip"), io.Latent.Input("latent"), - io.String.Input("global_prompt", multiline=True, default=""), + io.String.Input( + "global_prompt", multiline=True, default="", + tooltip="Conditions entire video. Leave empty to auto-use the first parsed segment from smart_prompt as the global anchor." + ), io.String.Input( "smart_prompt", multiline=True, default="", tooltip="Enter prompt using Smart Syntax:\\n1. Inline: 'text one [0-50] | text two [50-100]'\\n2. Block: 'Second 1:\\ntext one\\nSecond 2:\\ntext two'\\nSyntax is auto-stripped and normalized evenly or proportionally." ), - io.Bool.Input("normalize_by_tokens", default=False, tooltip="If true, scales the calculated length of each segment by its token count."), + io.Boolean.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=[ @@ -69,8 +72,12 @@ class PromptRelaySmartEncode(io.ComfyNode): scale_factor = 100000.0 segment_lengths_str = ", ".join(str(int(w * scale_factor)) for w in weights_list) + global_prompt_str = global_prompt.strip() + if not global_prompt_str and valid_segments: + global_prompt_str = valid_segments[0]["text"] + patched, conditioning = _encode_relay( - model, clip, latent, global_prompt, local_prompts_str, segment_lengths_str, epsilon + model, clip, latent, global_prompt_str, local_prompts_str, segment_lengths_str, epsilon ) return io.NodeOutput(patched, conditioning) @@ -91,7 +98,7 @@ class PromptRelaySmartEncodeTest(io.ComfyNode): "smart_prompt", multiline=True, default="", tooltip="Enter prompt using Smart Syntax:\\n1. Inline: 'text one [0-50] | text two [50-100]'\\n2. Block: 'Second 1:\\ntext one\\nSecond 2:\\ntext two'\\nSyntax is auto-stripped and normalized evenly or proportionally." ), - io.Bool.Input("normalize_by_tokens", default=False), + io.Boolean.Input("normalize_by_tokens", default=False), io.Clip.Input("clip", optional=True), ], outputs=[