From 346ff8b7c4d9357822915820fedf00feb5cd92f3 Mon Sep 17 00:00:00 2001 From: Artificial Sweetener Date: Tue, 26 May 2026 17:45:21 -0400 Subject: [PATCH] feat(prompt-control): add schedule and encode prompt node Add Prompt-Control prompt parsing, lazy graph expansion, legacy and v3 node exports, and batch-aware conditioning support for KSampler Extras and tiled diffusion. --- simple_syrup/domain/prompt_control_prompt.py | 76 +++++ simple_syrup/nodes/__init__.py | 17 ++ simple_syrup/nodes/ksampler_extras.py | 114 +++++++- .../nodes/ksampler_tiled_diffusion.py | 4 +- simple_syrup/nodes/prompt_encode_style.py | 4 +- .../prompt_encode_style_and_normalization.py | 4 +- ..._and_encode_prompts_with_prompt_control.py | 118 ++++++++ simple_syrup/nodes_v3/__init__.py | 4 + ..._and_encode_prompts_with_prompt_control.py | 137 +++++++++ .../prompt_control_schedule_encode_graph.py | 195 +++++++++++++ .../tiled_diffusion_sampling_service.py | 124 ++++++++ tests/test_ksampler_extras_node.py | 107 +++++++ tests/test_ksampler_tiled_diffusion_node.py | 2 + tests/test_node_tooltips.py | 4 + tests/test_prompt_control_prompt.py | 83 ++++++ ...st_prompt_control_schedule_encode_graph.py | 224 ++++++++++++++ tests/test_prompt_encode_style_nodes.py | 20 +- tests/test_registration.py | 25 ++ ...encode_prompts_with_prompt_control_node.py | 273 ++++++++++++++++++ .../test_tiled_diffusion_sampling_service.py | 58 ++++ 20 files changed, 1561 insertions(+), 32 deletions(-) create mode 100644 simple_syrup/domain/prompt_control_prompt.py create mode 100644 simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py create mode 100644 simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py create mode 100644 simple_syrup/runtime/prompt_control_schedule_encode_graph.py create mode 100644 tests/test_prompt_control_prompt.py create mode 100644 tests/test_prompt_control_schedule_encode_graph.py create mode 100644 tests/test_schedule_and_encode_prompts_with_prompt_control_node.py diff --git a/simple_syrup/domain/prompt_control_prompt.py b/simple_syrup/domain/prompt_control_prompt.py new file mode 100644 index 0000000..a1bbe4a --- /dev/null +++ b/simple_syrup/domain/prompt_control_prompt.py @@ -0,0 +1,76 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Prepare Prompt-Control prompt text for scheduling and encoding.""" + +from __future__ import annotations + +import re +from dataclasses import dataclass + +from .conditioning_batch import split_prompt_batch + +PROMPT_TEXT_PATTERN = r"(?:^|>)([^<]+)(?=<|$)" +LORA_TAG_PATTERN = r"<[^>]*>" + + +@dataclass(frozen=True) +class PreparedPromptChunk: + """Store one prompt chunk's cleaned text and scheduling tags.""" + + text: str + lora_tags: str + + +@dataclass(frozen=True) +class PreparedPromptSide: + """Store ordered prompt chunks and all scheduling tags for one prompt side.""" + + chunks: tuple[PreparedPromptChunk, ...] + lora_tags: str + + +def extract_prompt_text(text: str) -> str: + """Return prompt text outside angle-bracket Prompt-Control tags.""" + + return _extract_all_matches(text, PROMPT_TEXT_PATTERN) + + +def extract_lora_tags(text: str) -> str: + """Return angle-bracket Prompt-Control tags joined with newlines.""" + + return _extract_all_matches(text, LORA_TAG_PATTERN) + + +def prepare_prompt_side(text: str, separator: str) -> PreparedPromptSide: + """Split a prompt side into cleaned chunks and aggregate LoRA tags.""" + + chunks = tuple( + PreparedPromptChunk( + text=extract_prompt_text(chunk), + lora_tags=extract_lora_tags(chunk), + ) + for chunk in split_prompt_batch(text, separator) + ) + lora_tags = "\n".join(chunk.lora_tags for chunk in chunks if chunk.lora_tags) + return PreparedPromptSide(chunks=chunks, lora_tags=lora_tags) + + +def apply_encode_style(encode_style: str, prompt_text: str) -> str: + """Prepend Prompt-Control encode style text exactly as provided.""" + + if not encode_style: + return prompt_text + return f"{encode_style}{prompt_text}" + + +def _extract_all_matches(text: str, pattern: str) -> str: + """Match Comfy's RegexExtract All Matches behavior.""" + + matches = re.findall(pattern, text, re.IGNORECASE) + if not matches: + return "" + if isinstance(matches[0], tuple): + return "\n".join(match[0] for match in matches) + return "\n".join(matches) diff --git a/simple_syrup/nodes/__init__.py b/simple_syrup/nodes/__init__.py index c06593d..7934259 100644 --- a/simple_syrup/nodes/__init__.py +++ b/simple_syrup/nodes/__init__.py @@ -6,6 +6,7 @@ from __future__ import annotations +from ..runtime.prompt_control_availability import prompt_control_is_available from .conditioning_batch_pack import ConditioningBatchAppend, ConditioningBatchStart from .detail_segs_as_regions import DetailSEGSAsRegions from .detail_segs_by_scale_factor import DetailSEGSByScaleFactor @@ -36,6 +37,8 @@ from .vae_options import VAEDecodeOptions, VAEEncodeOptions from .vitmatte_model_loader import ViTMatteModelLoader from .wd14_tagger_loader import WD14TaggerLoader +_PROMPT_CONTROL_EXPORTS: list[str] = [] + NODE_CLASS_MAPPINGS = { "SimpleSyrup.ConditioningBatchAppend": ConditioningBatchAppend, "SimpleSyrup.ConditioningBatchStart": ConditioningBatchStart, @@ -108,6 +111,19 @@ NODE_DISPLAY_NAME_MAPPINGS = { "SimpleSyrup.WD14TaggerLoader": "Load WD14 Tagger", } +if prompt_control_is_available(): + from .schedule_and_encode_prompts_with_prompt_control import ( + ScheduleAndEncodePromptsWithPromptControl, + ) + + NODE_CLASS_MAPPINGS["SimpleSyrup.ScheduleAndEncodePromptsWithPromptControl"] = ( + ScheduleAndEncodePromptsWithPromptControl + ) + NODE_DISPLAY_NAME_MAPPINGS[ + "SimpleSyrup.ScheduleAndEncodePromptsWithPromptControl" + ] = "Schedule & Encode Prompts" + _PROMPT_CONTROL_EXPORTS.append("ScheduleAndEncodePromptsWithPromptControl") + __all__ = [ "ConditioningBatchAppend", "ConditioningBatchStart", @@ -130,6 +146,7 @@ __all__ = [ "PromptSEGSWithSAM", "ResizeImageToTarget", "SAMModelLoader", + *_PROMPT_CONTROL_EXPORTS, "ScaleFactor", "Seed", "SimpleLoadAnima", diff --git a/simple_syrup/nodes/ksampler_extras.py b/simple_syrup/nodes/ksampler_extras.py index 8ca1699..b44b00e 100644 --- a/simple_syrup/nodes/ksampler_extras.py +++ b/simple_syrup/nodes/ksampler_extras.py @@ -9,6 +9,9 @@ from __future__ import annotations from importlib import import_module from typing import Any +import torch + +from ..domain.conditioning_batch import ConditioningBatch, select_conditioning from ..runtime import sampling_samplers, sampling_schedulers from . import tooltips @@ -74,11 +77,11 @@ class KSamplerExtras: {"tooltip": tooltips.SCHEDULER}, ), "positive": ( - "CONDITIONING", + "CONDITIONING,CONDITIONING_BATCH", {"tooltip": tooltips.POSITIVE_CONDITIONING}, ), "negative": ( - "CONDITIONING", + "CONDITIONING,CONDITIONING_BATCH", {"tooltip": tooltips.NEGATIVE_CONDITIONING}, ), "latent_image": ("LATENT", {"tooltip": tooltips.LATENT_IMAGE}), @@ -138,20 +141,37 @@ class KSamplerExtras: callback = latent_preview.prepare_callback(model, steps) disable_pbar = not comfy_utils.PROGRESS_BAR_ENABLED - samples = comfy_sample.sample_custom( - model, - noise, - cfg, - sampler, - sigmas, - positive, - negative, - latent_samples, - noise_mask=noise_mask, - callback=callback, - disable_pbar=disable_pbar, - seed=seed, - ) + if _uses_conditioning_batch(positive, negative): + samples = _sample_conditioning_batch( + comfy_sample=comfy_sample, + model=model, + noise=noise, + cfg=cfg, + sampler=sampler, + sigmas=sigmas, + positive=positive, + negative=negative, + latent_samples=latent_samples, + noise_mask=noise_mask, + callback=callback, + disable_pbar=disable_pbar, + seed=seed, + ) + else: + samples = comfy_sample.sample_custom( + model, + noise, + cfg, + sampler, + sigmas, + positive, + negative, + latent_samples, + noise_mask=noise_mask, + callback=callback, + disable_pbar=disable_pbar, + seed=seed, + ) output = latent_image.copy() output.pop("downscale_ratio_spacial", None) @@ -167,6 +187,68 @@ def _comfy_sample() -> Any: return comfy.sample +def _uses_conditioning_batch(positive: Any, negative: Any) -> bool: + """Return whether either conditioning input needs per-item selection.""" + + return isinstance(positive, ConditioningBatch) or isinstance( + negative, + ConditioningBatch, + ) + + +def _sample_conditioning_batch( + *, + comfy_sample: Any, + model: Any, + noise: torch.Tensor, + cfg: float, + sampler: Any, + sigmas: torch.Tensor, + positive: Any, + negative: Any, + latent_samples: torch.Tensor, + noise_mask: Any, + callback: Any, + disable_pbar: bool, + seed: int, +) -> torch.Tensor: + """Sample each latent batch item with its selected conditioning.""" + + sampled: list[torch.Tensor] = [] + for index in range(int(latent_samples.shape[0])): + sampled.append( + comfy_sample.sample_custom( + model, + noise[index : index + 1], + cfg, + sampler, + sigmas, + select_conditioning(positive, index), + select_conditioning(negative, index), + latent_samples[index : index + 1], + noise_mask=_slice_noise_mask(noise_mask, index, latent_samples), + callback=callback, + disable_pbar=disable_pbar, + seed=seed, + ) + ) + return torch.cat(sampled, dim=0) + + +def _slice_noise_mask( + noise_mask: Any, + index: int, + latent_samples: torch.Tensor, +) -> Any: + """Return the noise mask slice matching one latent batch item.""" + + if isinstance(noise_mask, torch.Tensor) and noise_mask.shape[0] == int( + latent_samples.shape[0], + ): + return noise_mask[index : index + 1] + return noise_mask + + def _comfy_utils() -> Any: """Import ComfyUI utility state lazily.""" diff --git a/simple_syrup/nodes/ksampler_tiled_diffusion.py b/simple_syrup/nodes/ksampler_tiled_diffusion.py index 0ba0cd7..ada4761 100644 --- a/simple_syrup/nodes/ksampler_tiled_diffusion.py +++ b/simple_syrup/nodes/ksampler_tiled_diffusion.py @@ -84,11 +84,11 @@ class KSamplerTiledDiffusion: {"tooltip": tooltips.SCHEDULER}, ), "positive": ( - "CONDITIONING", + "CONDITIONING,CONDITIONING_BATCH", {"tooltip": tooltips.POSITIVE_CONDITIONING}, ), "negative": ( - "CONDITIONING", + "CONDITIONING,CONDITIONING_BATCH", {"tooltip": tooltips.NEGATIVE_CONDITIONING}, ), "latent_image": ("LATENT", {"tooltip": tooltips.LATENT_IMAGE}), diff --git a/simple_syrup/nodes/prompt_encode_style.py b/simple_syrup/nodes/prompt_encode_style.py index cdf403e..34d65e4 100644 --- a/simple_syrup/nodes/prompt_encode_style.py +++ b/simple_syrup/nodes/prompt_encode_style.py @@ -15,8 +15,8 @@ class PromptEncodeStyle: """Build Prompt Control STYLE tags from encode-style selections.""" RETURN_TYPES = ("STRING",) - RETURN_NAMES = ("style_tag",) - OUTPUT_TOOLTIPS = ("Prompt Control STYLE tag text for prompt encoding workflows.",) + RETURN_NAMES = ("encode_style",) + OUTPUT_TOOLTIPS = ("Prompt Control encode style text for prompt workflows.",) FUNCTION = "build" CATEGORY = "SimpleSyrup/Prompt" DESCRIPTION = "Builds a Prompt Control STYLE tag from an encode-style selection." diff --git a/simple_syrup/nodes/prompt_encode_style_and_normalization.py b/simple_syrup/nodes/prompt_encode_style_and_normalization.py index 0664006..b2e95a9 100644 --- a/simple_syrup/nodes/prompt_encode_style_and_normalization.py +++ b/simple_syrup/nodes/prompt_encode_style_and_normalization.py @@ -19,9 +19,9 @@ class PromptEncodeStyleAndNormalization: """Build STYLE tags from encode-style and normalization selections.""" RETURN_TYPES = ("STRING",) - RETURN_NAMES = ("style_tag",) + RETURN_NAMES = ("encode_style",) OUTPUT_TOOLTIPS = ( - "Prompt Control STYLE tag text with the selected normalization behavior.", + "Prompt Control encode style text with the selected normalization behavior.", ) FUNCTION = "build" CATEGORY = "SimpleSyrup/Prompt" diff --git a/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py b/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py new file mode 100644 index 0000000..22cd362 --- /dev/null +++ b/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py @@ -0,0 +1,118 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Legacy ComfyUI node for Prompt-Control prompt scheduling and encoding.""" + +from __future__ import annotations + +from typing import Any + +from ..runtime.prompt_control_schedule_encode_graph import ( + PromptControlScheduleEncodeGraphBuilder, +) + + +class ScheduleAndEncodePromptsWithPromptControl: + """Schedule Prompt-Control LoRAs and encode prompts with optional batches.""" + + RETURN_TYPES = ( + "MODEL", + "CONDITIONING,CONDITIONING_BATCH", + "CONDITIONING,CONDITIONING_BATCH", + ) + RETURN_NAMES = ("model", "positive", "negative") + OUTPUT_TOOLTIPS = ( + "Model after LoRA tags from positive and negative prompts are scheduled.", + "Positive conditioning or SimpleSyrup conditioning batch.", + "Negative conditioning or SimpleSyrup conditioning batch.", + ) + FUNCTION = "execute" + CATEGORY = "SimpleSyrup/Conditioning" + DESCRIPTION = ( + "Schedules Prompt-Control LoRAs and encodes prompts, using [SEP] to " + "create SimpleSyrup conditioning batches." + ) + SEARCH_ALIASES = ["prompt control", "schedule prompts", "encode prompts"] + + @classmethod + def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]: + """Declare Prompt-Control schedule and encode inputs.""" + + return { + "required": { + "model": ( + "MODEL", + { + "rawLink": True, + "tooltip": ( + "Model that receives LoRA changes found in " + "Prompt-Control prompt tags." + ), + }, + ), + "clip": ( + "CLIP", + { + "rawLink": True, + "tooltip": ( + "CLIP connection used for scheduled hooks and " + "cleaned prompt encoding." + ), + }, + ), + "positive_prompt": ( + "STRING", + { + "default": "", + "multiline": False, + "tooltip": ( + "Positive Prompt-Control text. [SEP] creates a " + "conditioning batch for SimpleSyrup batch-aware nodes." + ), + }, + ), + "negative_prompt": ( + "STRING", + { + "default": "", + "multiline": False, + "tooltip": ( + "Negative Prompt-Control text. [SEP] creates a " + "conditioning batch for SimpleSyrup batch-aware nodes." + ), + }, + ), + }, + "optional": { + "encode_style": ( + "STRING", + { + "default": "", + "forceInput": True, + "tooltip": ( + "Encode style text from Prompt Encode Style or " + "Prompt Encode Style & Normalization." + ), + }, + ), + }, + } + + def execute( + self, + model: Any, + clip: Any, + positive_prompt: str, + negative_prompt: str, + encode_style: str = "", + ) -> Any: + """Build lazy Prompt-Control graph expansion for prompts.""" + + return PromptControlScheduleEncodeGraphBuilder().build( + model=model, + clip=clip, + positive_prompt=positive_prompt, + negative_prompt=negative_prompt, + encode_style=encode_style, + ) diff --git a/simple_syrup/nodes_v3/__init__.py b/simple_syrup/nodes_v3/__init__.py index fd6446e..efafa0f 100644 --- a/simple_syrup/nodes_v3/__init__.py +++ b/simple_syrup/nodes_v3/__init__.py @@ -32,6 +32,9 @@ def get_nodes() -> list[type[object]]: from .encode_prompt_batch_with_prompt_control import ( EncodePromptBatchWithPromptControl, ) + from .schedule_and_encode_prompts_with_prompt_control import ( + ScheduleAndEncodePromptsWithPromptControl, + ) return [ WD14TaggerLoaderV3, @@ -41,6 +44,7 @@ def get_nodes() -> list[type[object]]: VAEDecodeOptionsV3, VAEEncodeOptionsV3, EncodePromptBatchWithPromptControl, + ScheduleAndEncodePromptsWithPromptControl, ] diff --git a/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py b/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py new file mode 100644 index 0000000..5b33d56 --- /dev/null +++ b/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py @@ -0,0 +1,137 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Comfy v3 node for Prompt-Control prompt scheduling and encoding.""" + +from __future__ import annotations + +from importlib import import_module +from typing import TYPE_CHECKING, Any, ClassVar + +from ..runtime.prompt_control_schedule_encode_graph import ( + PromptControlScheduleEncodeGraphBuilder, +) + +if TYPE_CHECKING: + + class _ComfyNodeBase: + """Type-checking base for Comfy v3 nodes.""" + + RETURN_TYPES: ClassVar[list[str]] + RETURN_NAMES: ClassVar[list[str]] + + @classmethod + def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]: + """Return v1-compatible input metadata.""" + + raise NotImplementedError + +else: + _ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode + +_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io +MixedConditioningIO: Any = ( + None if TYPE_CHECKING else _comfy_io.Custom("CONDITIONING,CONDITIONING_BATCH") +) + + +class ScheduleAndEncodePromptsWithPromptControl(_ComfyNodeBase): + """Schedule Prompt-Control LoRAs and encode prompts with optional batches.""" + + @classmethod + def define_schema(cls) -> Any: + """Declare the Prompt-Control schedule and encode schema.""" + + return _comfy_io.Schema( + node_id="SimpleSyrup.ScheduleAndEncodePromptsWithPromptControl", + display_name="Schedule & Encode Prompts", + enable_expand=True, + category="SimpleSyrup/Conditioning", + description=( + "Schedules Prompt-Control LoRAs and encodes prompts, using [SEP] " + "to create SimpleSyrup conditioning batches." + ), + inputs=[ + _comfy_io.Model.Input( + "model", + raw_link=True, + tooltip=( + "Model that receives LoRA changes found in Prompt-Control " + "prompt tags." + ), + ), + _comfy_io.Clip.Input( + "clip", + raw_link=True, + tooltip=( + "CLIP connection used for scheduled hooks and cleaned " + "prompt encoding." + ), + ), + _comfy_io.String.Input( + "encode_style", + default="", + force_input=True, + optional=True, + tooltip=( + "Encode style text from Prompt Encode Style or " + "Prompt Encode Style & Normalization." + ), + ), + _comfy_io.String.Input( + "positive_prompt", + multiline=False, + default="", + tooltip=( + "Positive Prompt-Control text. [SEP] creates a " + "conditioning batch for SimpleSyrup batch-aware nodes." + ), + ), + _comfy_io.String.Input( + "negative_prompt", + multiline=False, + default="", + tooltip=( + "Negative Prompt-Control text. [SEP] creates a " + "conditioning batch for SimpleSyrup batch-aware nodes." + ), + ), + ], + outputs=[ + _comfy_io.Model.Output( + "model", + tooltip=( + "Model after LoRA tags from positive and negative prompts " + "are scheduled." + ), + ), + MixedConditioningIO.Output( + "positive", + tooltip="Positive conditioning or SimpleSyrup conditioning batch.", + ), + MixedConditioningIO.Output( + "negative", + tooltip="Negative conditioning or SimpleSyrup conditioning batch.", + ), + ], + ) + + @classmethod + def execute( + cls, + model: Any, + clip: Any, + positive_prompt: str, + negative_prompt: str, + encode_style: str = "", + ) -> Any: + """Build lazy Prompt-Control graph expansion for prompts.""" + + return PromptControlScheduleEncodeGraphBuilder().build( + model=model, + clip=clip, + positive_prompt=positive_prompt, + negative_prompt=negative_prompt, + encode_style=encode_style, + ) diff --git a/simple_syrup/runtime/prompt_control_schedule_encode_graph.py b/simple_syrup/runtime/prompt_control_schedule_encode_graph.py new file mode 100644 index 0000000..4a34052 --- /dev/null +++ b/simple_syrup/runtime/prompt_control_schedule_encode_graph.py @@ -0,0 +1,195 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Runtime graph expansion for Prompt-Control scheduling and prompt encoding.""" + +from __future__ import annotations + +import sys +from importlib import import_module +from typing import Any, cast + +from ..domain.prompt_control_prompt import ( + PreparedPromptSide, + apply_encode_style, + prepare_prompt_side, +) +from .prompt_control_availability import find_prompt_control_install + +PROMPT_CONTROL_MISSING_MESSAGE = ( + "Schedule & Encode Prompts requires comfyui-prompt-control. " + "Install Prompt Control or remove this node from the workflow." +) +PROMPT_BATCH_SEPARATOR = "[SEP]" + + +class PromptControlScheduleEncodeGraphBuilder: + """Build lazy Prompt-Control graphs for LoRA scheduling and prompt encoding.""" + + def build( + self, + model: Any, + clip: Any, + positive_prompt: str, + negative_prompt: str, + encode_style: str = "", + ) -> Any: + """Return an io.NodeOutput for scheduled model and encoded prompts.""" + + io, graph_utils, lazy_nodes = self._prompt_control_dependencies() + positive_side = prepare_prompt_side(positive_prompt, PROMPT_BATCH_SEPARATOR) + negative_side = prepare_prompt_side(negative_prompt, PROMPT_BATCH_SEPARATOR) + + expand: dict[str, dict[str, Any]] = {} + positive_lora = lazy_nodes.PCLazyLoraLoaderAdvanced.execute( + model=model, + clip=clip, + text=positive_side.lora_tags, + apply_hooks=True, + tags="", + start=0.0, + end=1.0, + num_steps=0, + ) + self._merge_expand(expand, positive_lora.expand, "positive LoRA scheduling") + + negative_lora = lazy_nodes.PCLazyLoraLoaderAdvanced.execute( + model=positive_lora.args[0], + clip=positive_lora.args[1], + text=negative_side.lora_tags, + apply_hooks=True, + tags="", + start=0.0, + end=1.0, + num_steps=0, + ) + self._merge_expand(expand, negative_lora.expand, "negative LoRA scheduling") + + scheduled_model = negative_lora.args[0] + scheduled_clip = negative_lora.args[1] + positive_conditioning = self._encode_side( + side=positive_side, + clip=scheduled_clip, + encode_style=encode_style, + graph_utils=graph_utils, + lazy_text_encoder=lazy_nodes.PCLazyTextEncodeAdvanced, + expand=expand, + label="positive prompt encoding", + ) + negative_conditioning = self._encode_side( + side=negative_side, + clip=scheduled_clip, + encode_style=encode_style, + graph_utils=graph_utils, + lazy_text_encoder=lazy_nodes.PCLazyTextEncodeAdvanced, + expand=expand, + label="negative prompt encoding", + ) + + return io.NodeOutput( + scheduled_model, + positive_conditioning, + negative_conditioning, + expand=expand, + ) + + def _encode_side( + self, + *, + side: PreparedPromptSide, + clip: Any, + encode_style: str, + graph_utils: Any, + lazy_text_encoder: Any, + expand: dict[str, dict[str, Any]], + label: str, + ) -> Any: + """Encode one prompt side and return conditioning or conditioning batch.""" + + conditioning_outputs: list[Any] = [] + for index, chunk in enumerate(side.chunks): + text = apply_encode_style(encode_style, chunk.text) + node_output = lazy_text_encoder.execute( + clip=clip, + text=text, + tags="", + start=0.0, + end=1.0, + num_steps=0, + ) + self._merge_expand( + expand, + node_output.expand, + f"{label} chunk {index}", + ) + conditioning_outputs.append(node_output.args[0]) + + if len(conditioning_outputs) == 1: + return conditioning_outputs[0] + + pack_graph = graph_utils.GraphBuilder() + current = pack_graph.node( + "SimpleSyrup.ConditioningBatchStart", + conditioning=conditioning_outputs[0], + ) + for conditioning in conditioning_outputs[1:]: + current = pack_graph.node( + "SimpleSyrup.ConditioningBatchAppend", + batch=current.out(0), + conditioning=conditioning, + ) + self._merge_expand( + expand, + cast(dict[str, dict[str, Any]], pack_graph.finalize()), + f"{label} batch packing", + ) + return current.out(0) + + def _prompt_control_dependencies(self) -> tuple[Any, Any, Any]: + """Import Prompt-Control and Comfy graph helpers on demand.""" + + try: + io = import_module("comfy_api.latest.io") + except ModuleNotFoundError: + comfy_api = import_module("comfy_api.latest") + io = comfy_api.io + try: + graph_utils = import_module("comfy_execution.graph_utils") + lazy_nodes = self._import_prompt_control_lazy_nodes() + except ModuleNotFoundError as exc: + raise RuntimeError(PROMPT_CONTROL_MISSING_MESSAGE) from exc + return io, graph_utils, lazy_nodes + + def _import_prompt_control_lazy_nodes(self) -> Any: + """Import Prompt-Control lazy nodes from installed or sibling paths.""" + + try: + return import_module("prompt_control.nodes_lazy") + except ModuleNotFoundError: + availability = find_prompt_control_install() + if availability.root_path is not None: + root_path = str(availability.root_path) + if root_path not in sys.path: + sys.path.insert(0, root_path) + return import_module("prompt_control.nodes_lazy") + + def _merge_expand( + self, + target: dict[str, dict[str, Any]], + source: object, + operation: str, + ) -> None: + """Merge a lazy expand graph and reject duplicate generated node ids.""" + + if not source: + return + expand = cast(dict[str, dict[str, Any]], source) + overlap = set(target).intersection(expand) + if overlap: + overlapping_ids = ", ".join(sorted(overlap)) + raise ValueError( + f"Prompt-Control graph expansion generated duplicate node ids " + f"during {operation}: {overlapping_ids}." + ) + target.update(expand) diff --git a/simple_syrup/services/tiled_diffusion_sampling_service.py b/simple_syrup/services/tiled_diffusion_sampling_service.py index 2688779..33d9738 100644 --- a/simple_syrup/services/tiled_diffusion_sampling_service.py +++ b/simple_syrup/services/tiled_diffusion_sampling_service.py @@ -8,6 +8,9 @@ from __future__ import annotations from typing import Any +import torch + +from ..domain.conditioning_batch import ConditioningBatch, select_conditioning from ..domain.tiled_diffusion import validate_tiled_diffusion_mode from ..runtime import mixture_of_diffusers_sampling, multidiffusion_sampling from ..runtime.detail_previews import DetailPreviewContext @@ -42,6 +45,26 @@ class TiledDiffusionSamplingService: """Sample a latent with the selected tiled diffusion method.""" validate_tiled_diffusion_mode(diffusion_mode) + if self._uses_conditioning_batch(positive, negative): + return self._sample_conditioning_batch( + diffusion_mode=diffusion_mode, + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=positive, + negative=negative, + latent_image=latent_image, + denoise=denoise, + latent_tile_width=latent_tile_width, + latent_tile_height=latent_tile_height, + latent_tile_overlap=latent_tile_overlap, + latent_tile_batch_size=latent_tile_batch_size, + preview_context=preview_context, + differential_diffusion=differential_diffusion, + ) if diffusion_mode == "multidiffusion": return multidiffusion_sampling.sample_multidiffusion( model=model, @@ -79,3 +102,104 @@ class TiledDiffusionSamplingService: preview_context=preview_context, differential_diffusion=differential_diffusion, ) + + def _sample_conditioning_batch( + self, + *, + diffusion_mode: str, + model: Any, + seed: int, + steps: int, + cfg: float, + sampler_name: str, + scheduler: str, + positive: Any, + negative: Any, + latent_image: Latent, + denoise: float, + latent_tile_width: int, + latent_tile_height: int, + latent_tile_overlap: int, + latent_tile_batch_size: int, + preview_context: DetailPreviewContext | None, + differential_diffusion: bool, + ) -> Latent: + """Sample latent batch items one at a time with selected conditioning.""" + + latent_samples = latent_image["samples"] + if not isinstance(latent_samples, torch.Tensor): + raise TypeError("Tiled diffusion latent samples must be a torch.Tensor.") + + outputs: list[torch.Tensor] = [] + for index in range(int(latent_samples.shape[0])): + item_latent = self._single_item_latent(latent_image, index) + output = self.sample( + diffusion_mode=diffusion_mode, + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=select_conditioning(positive, index), + negative=select_conditioning(negative, index), + latent_image=item_latent, + denoise=denoise, + latent_tile_width=latent_tile_width, + latent_tile_height=latent_tile_height, + latent_tile_overlap=latent_tile_overlap, + latent_tile_batch_size=latent_tile_batch_size, + preview_context=preview_context, + differential_diffusion=differential_diffusion, + ) + output_samples = output["samples"] + if not isinstance(output_samples, torch.Tensor): + raise TypeError( + "Tiled diffusion output samples must be a torch.Tensor." + ) + outputs.append(output_samples) + + result = latent_image.copy() + result.pop("downscale_ratio_spacial", None) + result["samples"] = torch.cat(outputs, dim=0) + return result + + def _single_item_latent(self, latent_image: Latent, index: int) -> Latent: + """Return a latent dictionary for one batch item.""" + + latent_samples = latent_image["samples"] + if not isinstance(latent_samples, torch.Tensor): + raise TypeError("Tiled diffusion latent samples must be a torch.Tensor.") + item = latent_image.copy() + item["samples"] = latent_samples[index : index + 1] + if "batch_index" in item: + item["batch_index"] = [item["batch_index"][index]] + if "noise_mask" in item: + item["noise_mask"] = self._slice_noise_mask( + item["noise_mask"], + index, + latent_samples, + ) + return item + + def _slice_noise_mask( + self, + noise_mask: Any, + index: int, + latent_samples: torch.Tensor, + ) -> Any: + """Return the noise mask slice matching one latent batch item.""" + + if isinstance(noise_mask, torch.Tensor) and noise_mask.shape[0] == int( + latent_samples.shape[0], + ): + return noise_mask[index : index + 1] + return noise_mask + + def _uses_conditioning_batch(self, positive: Any, negative: Any) -> bool: + """Return whether tiled sampling needs per-item conditioning selection.""" + + return isinstance(positive, ConditioningBatch) or isinstance( + negative, + ConditioningBatch, + ) diff --git a/tests/test_ksampler_extras_node.py b/tests/test_ksampler_extras_node.py index 79a7776..dd65dee 100644 --- a/tests/test_ksampler_extras_node.py +++ b/tests/test_ksampler_extras_node.py @@ -12,6 +12,7 @@ from typing import Any import torch +from simple_syrup.domain.conditioning_batch import ConditioningBatch from simple_syrup.nodes.ksampler_extras import KSamplerExtras from simple_syrup.runtime import sampling_samplers, sampling_schedulers @@ -63,6 +64,8 @@ def test_input_types_match_simple_ksampler_contract() -> None: "latent_image", "denoise", ) + assert required["positive"][0] == "CONDITIONING,CONDITIONING_BATCH" + assert required["negative"][0] == "CONDITIONING,CONDITIONING_BATCH" def test_node_metadata_matches_contract() -> None: @@ -289,3 +292,107 @@ def test_sample_delegates_to_runtime_helpers( assert calls["sample_custom"]["sampler"] is sampler assert calls["sample_custom"]["sigmas"] is fixed_sigmas assert calls["sample_custom"]["disable_pbar"] is True + + +def test_sample_selects_conditioning_batch_per_latent_item( + monkeypatch: Any, +) -> None: + """Conditioning batches are selected before calling Comfy sampling.""" + + calls: list[dict[str, Any]] = [] + model = FakeModel() + sampler = FakeSampler() + latent_samples = torch.arange(2 * 4 * 2 * 2, dtype=torch.float32).reshape( + (2, 4, 2, 2) + ) + fixed_noise = torch.full_like(latent_samples, 2.0) + fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) + noise_mask = torch.ones((2, 1, 2, 2), dtype=torch.float32) + latent_image: dict[str, Any] = { + "samples": latent_samples, + "batch_index": [4, 9], + "noise_mask": noise_mask, + } + + monkeypatch.setattr( + sampling_samplers, + "resolve_sampler", + lambda sampler_name: sampler, + ) + monkeypatch.setattr( + sampling_schedulers, + "calculate_sigmas", + lambda **kwargs: fixed_sigmas, + ) + monkeypatch.setattr( + comfy_sample, + "fix_empty_latent_channels", + lambda model, samples, downscale_ratio_spacial: samples, + ) + monkeypatch.setattr( + comfy_sample, + "prepare_noise", + lambda samples, seed, batch_inds: fixed_noise, + ) + monkeypatch.setattr( + latent_preview, + "prepare_callback", + lambda received_model, steps: "callback", + ) + + def fake_sample_custom( + received_model: FakeModel, + noise: torch.Tensor, + cfg: float, + received_sampler: FakeSampler, + sigmas: torch.Tensor, + positive: object, + negative: object, + latent_image: torch.Tensor, + noise_mask: torch.Tensor | None, + callback: str, + disable_pbar: bool, + seed: int, + ) -> torch.Tensor: + """Record one per-item sample call and return a marked tensor.""" + + del received_model, cfg, received_sampler, sigmas, callback, disable_pbar, seed + calls.append( + { + "noise": noise, + "positive": positive, + "negative": negative, + "latent_image": latent_image, + "noise_mask": noise_mask, + } + ) + return torch.full_like(latent_image, float(len(calls))) + + monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) + monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", False) + + (output,) = KSamplerExtras().sample( + model=model, + seed=123, + steps=2, + cfg=7.5, + sampler_name="lcm", + scheduler="GITS", + positive=ConditioningBatch(("positive-0", "positive-1")), + negative=ConditioningBatch(("negative-last",)), + latent_image=latent_image, + denoise=0.8, + ) + + assert len(calls) == 2 + assert calls[0]["positive"] == "positive-0" + assert calls[1]["positive"] == "positive-1" + assert calls[0]["negative"] == "negative-last" + assert calls[1]["negative"] == "negative-last" + assert torch.equal(calls[0]["noise"], fixed_noise[0:1]) + assert torch.equal(calls[1]["noise"], fixed_noise[1:2]) + assert torch.equal(calls[0]["noise_mask"], noise_mask[0:1]) + assert torch.equal(calls[1]["noise_mask"], noise_mask[1:2]) + assert output["samples"].shape == latent_samples.shape + assert torch.equal(output["samples"][0], torch.full((4, 2, 2), 1.0)) + assert torch.equal(output["samples"][1], torch.full((4, 2, 2), 2.0)) diff --git a/tests/test_ksampler_tiled_diffusion_node.py b/tests/test_ksampler_tiled_diffusion_node.py index 5d59e2b..bd771b5 100644 --- a/tests/test_ksampler_tiled_diffusion_node.py +++ b/tests/test_ksampler_tiled_diffusion_node.py @@ -50,6 +50,8 @@ def test_input_types_match_tiled_diffusion_contract( "mixture_of_diffusers", ] assert required["diffusion_mode"][1]["default"] == "multidiffusion" + assert required["positive"][0] == "CONDITIONING,CONDITIONING_BATCH" + assert required["negative"][0] == "CONDITIONING,CONDITIONING_BATCH" assert required["latent_tile_width"][1]["default"] == 128 assert required["latent_tile_width"][1]["max"] == 512 assert required["latent_tile_height"][1]["default"] == 128 diff --git a/tests/test_node_tooltips.py b/tests/test_node_tooltips.py index 3988dfb..20713d8 100644 --- a/tests/test_node_tooltips.py +++ b/tests/test_node_tooltips.py @@ -18,6 +18,9 @@ from simple_syrup.nodes_v3.encode_prompt_batch_with_prompt_control import ( EncodePromptBatchWithPromptControl, ) from simple_syrup.nodes_v3.scale_factor import ScaleFactorV3 +from simple_syrup.nodes_v3.schedule_and_encode_prompts_with_prompt_control import ( + ScheduleAndEncodePromptsWithPromptControl, +) from simple_syrup.nodes_v3.simple_load_checkpoint import SimpleLoadCheckpointV3 from simple_syrup.nodes_v3.tile_and_tag_segs import TileAndTagSEGSV3 from simple_syrup.nodes_v3.vae_decode_options import VAEDecodeOptionsV3 @@ -122,6 +125,7 @@ def test_legacy_named_outputs_provide_tooltips() -> None: VAEEncodeOptionsV3, WD14TaggerLoaderV3, EncodePromptBatchWithPromptControl, + ScheduleAndEncodePromptsWithPromptControl, ], ) def test_v3_nodes_provide_tooltip_metadata( diff --git a/tests/test_prompt_control_prompt.py b/tests/test_prompt_control_prompt.py new file mode 100644 index 0000000..af1ccaf --- /dev/null +++ b/tests/test_prompt_control_prompt.py @@ -0,0 +1,83 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for Prompt-Control prompt preparation helpers.""" + +from __future__ import annotations + +import pytest + +from simple_syrup.domain.prompt_control_prompt import ( + apply_encode_style, + extract_lora_tags, + extract_prompt_text, + prepare_prompt_side, +) + + +def test_extract_prompt_text_matches_sugarcubes_regex_behavior() -> None: + """Visible text outside angle tags is retained and joined with newlines.""" + + text = "portrait cinematic " + + assert extract_prompt_text(text) == "portrait \n cinematic " + + +def test_extract_prompt_text_keeps_text_after_adjacent_tag() -> None: + """Text after a closing angle tag is captured as a later prompt segment.""" + + assert extract_prompt_text("face") == "face" + + +def test_extract_lora_tags_returns_angle_tags_joined_with_newlines() -> None: + """All Prompt-Control angle tags are retained for LoRA scheduling.""" + + text = "portrait cinematic " + + assert extract_lora_tags(text) == "\n" + + +def test_prepare_prompt_side_splits_ordered_chunks_and_aggregates_loras() -> None: + """Separator-delimited chunks preserve order and collect all tags.""" + + side = prepare_prompt_side( + "face [SEP] hair ", + "[SEP]", + ) + + assert [chunk.text for chunk in side.chunks] == ["face ", "hair "] + assert [chunk.lora_tags for chunk in side.chunks] == [ + "", + "", + ] + assert side.lora_tags == "\n" + + +def test_prepare_prompt_side_preserves_empty_chunks() -> None: + """Batch splitting keeps empty entries so batch positions stay explicit.""" + + side = prepare_prompt_side("face [SEP] ", "[SEP]") + + assert [chunk.text for chunk in side.chunks] == ["face", ""] + assert [chunk.lora_tags for chunk in side.chunks] == ["", ""] + assert side.lora_tags == "" + + +def test_prepare_prompt_side_rejects_empty_separator() -> None: + """Empty separators are rejected by the shared batch splitter.""" + + with pytest.raises(ValueError, match="separator must not be empty"): + prepare_prompt_side("face", "") + + +def test_apply_encode_style_prepends_style_without_extra_formatting() -> None: + """Encode style text is used exactly as produced by style nodes.""" + + assert apply_encode_style("STYLE(A1111) ", "face") == "STYLE(A1111) face" + + +def test_apply_encode_style_keeps_prompt_when_style_is_blank() -> None: + """Blank encode style leaves cleaned prompt text unchanged.""" + + assert apply_encode_style("", "face") == "face" diff --git a/tests/test_prompt_control_schedule_encode_graph.py b/tests/test_prompt_control_schedule_encode_graph.py new file mode 100644 index 0000000..91f5f8c --- /dev/null +++ b/tests/test_prompt_control_schedule_encode_graph.py @@ -0,0 +1,224 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for Prompt-Control schedule and encode lazy graph expansion.""" + +from __future__ import annotations + +import sys +from importlib import import_module +from types import ModuleType +from typing import Any, cast + +import pytest + +from simple_syrup.runtime.prompt_control_schedule_encode_graph import ( + PROMPT_CONTROL_MISSING_MESSAGE, + PromptControlScheduleEncodeGraphBuilder, +) + + +def test_schedule_encode_graph_builds_single_conditioning_outputs( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Single prompts return direct Prompt-Control conditioning links.""" + + calls = _install_fake_prompt_control(monkeypatch) + + output = PromptControlScheduleEncodeGraphBuilder().build( + model=["model", 0], + clip=["clip", 0], + positive_prompt="face ", + negative_prompt="blur ", + encode_style="STYLE(A1111) ", + ) + + assert output.args[0] == ["lora_negative", 0] + assert output.args[1] == ["encode_0", 0] + assert output.args[2] == ["encode_1", 0] + assert output.expand is not None + assert not any( + node["class_type"].startswith("SimpleSyrup.ConditioningBatch") + for node in output.expand.values() + ) + assert calls["lora"] == [ + { + "model": ["model", 0], + "clip": ["clip", 0], + "text": "", + }, + { + "model": ["lora_positive", 0], + "clip": ["lora_positive", 1], + "text": "", + }, + ] + assert calls["encode"] == [ + {"clip": ["lora_negative", 1], "text": "STYLE(A1111) face "}, + {"clip": ["lora_negative", 1], "text": "STYLE(A1111) blur "}, + ] + + +def test_schedule_encode_graph_packs_only_multichunk_sides( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A side with multiple chunks becomes a conditioning batch.""" + + _install_fake_prompt_control(monkeypatch) + + output = PromptControlScheduleEncodeGraphBuilder().build( + model=["model", 0], + clip=["clip", 0], + positive_prompt="face [SEP] hair", + negative_prompt="blur", + ) + + assert output.args[1] != ["encode_0", 0] + assert output.args[2] == ["encode_2", 0] + assert output.expand is not None + pack_nodes = [ + node + for node in output.expand.values() + if node["class_type"].startswith("SimpleSyrup.ConditioningBatch") + ] + assert [node["class_type"] for node in pack_nodes] == [ + "SimpleSyrup.ConditioningBatchStart", + "SimpleSyrup.ConditioningBatchAppend", + ] + + +def test_schedule_encode_graph_collects_loras_from_all_chunks( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """LoRA scheduling sees all tags from every separator-delimited chunk.""" + + calls = _install_fake_prompt_control(monkeypatch) + + PromptControlScheduleEncodeGraphBuilder().build( + model=["model", 0], + clip=["clip", 0], + positive_prompt="face [SEP] hair ", + negative_prompt="blur [SEP] noise ", + ) + + assert calls["lora"][0]["text"] == "\n" + assert calls["lora"][1]["text"] == "\n" + + +def test_schedule_encode_graph_reports_duplicate_expand_ids( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Duplicate generated node ids are rejected with context.""" + + _install_fake_prompt_control(monkeypatch, duplicate_lora_ids=True) + + with pytest.raises(ValueError, match="negative LoRA scheduling"): + PromptControlScheduleEncodeGraphBuilder().build( + model=["model", 0], + clip=["clip", 0], + positive_prompt="face", + negative_prompt="blur", + ) + + +def test_schedule_encode_graph_reports_missing_prompt_control( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Missing Prompt-Control dependency raises an actionable error.""" + + def fake_import_module(name: str) -> Any: + if name == "prompt_control.nodes_lazy": + raise ModuleNotFoundError(name) + return import_module(name) + + monkeypatch.setattr( + "simple_syrup.runtime.prompt_control_schedule_encode_graph.import_module", + fake_import_module, + ) + + with pytest.raises(RuntimeError, match="requires comfyui-prompt-control"): + PromptControlScheduleEncodeGraphBuilder().build( + model=["model", 0], + clip=["clip", 0], + positive_prompt="face", + negative_prompt="blur", + ) + assert PROMPT_CONTROL_MISSING_MESSAGE.startswith("Schedule & Encode Prompts") + + +def _install_fake_prompt_control( + monkeypatch: pytest.MonkeyPatch, + *, + duplicate_lora_ids: bool = False, +) -> dict[str, list[dict[str, Any]]]: + """Install graph-expanding Prompt-Control stand-ins for builder tests.""" + + prompt_control = ModuleType("prompt_control") + nodes_lazy = ModuleType("prompt_control.nodes_lazy") + io = import_module("comfy_api.latest").io + calls: dict[str, list[dict[str, Any]]] = {"lora": [], "encode": []} + + class FakePCLazyLoraLoaderAdvanced: + """Prompt-Control LoRA scheduler test double.""" + + @staticmethod + def execute( + model: Any, + clip: Any, + text: str, + apply_hooks: bool, + tags: str, + start: float, + end: float, + num_steps: int, + ) -> Any: + """Return deterministic model and clip links.""" + + assert apply_hooks is True + assert tags == "" + assert start == 0.0 + assert end == 1.0 + assert num_steps == 0 + side = "positive" if not calls["lora"] else "negative" + calls["lora"].append({"model": model, "clip": clip, "text": text}) + node_id = "duplicate_lora" if duplicate_lora_ids else f"lora_{side}" + return io.NodeOutput( + [f"lora_{side}", 0], + [f"lora_{side}", 1], + None, + expand={node_id: {"class_type": "PromptControl.FakeLora"}}, + ) + + class FakePCLazyTextEncodeAdvanced: + """Prompt-Control text encoder test double.""" + + @staticmethod + def execute( + clip: Any, + text: str, + tags: str, + start: float, + end: float, + num_steps: int, + ) -> Any: + """Return deterministic conditioning links.""" + + assert tags == "" + assert start == 0.0 + assert end == 1.0 + assert num_steps == 0 + index = len(calls["encode"]) + calls["encode"].append({"clip": clip, "text": text}) + node_id = f"encode_{index}" + return io.NodeOutput( + [node_id, 0], + expand={node_id: {"class_type": "PromptControl.FakeTextEncode"}}, + ) + + cast(Any, nodes_lazy).PCLazyLoraLoaderAdvanced = FakePCLazyLoraLoaderAdvanced + cast(Any, nodes_lazy).PCLazyTextEncodeAdvanced = FakePCLazyTextEncodeAdvanced + cast(Any, prompt_control).nodes_lazy = nodes_lazy + monkeypatch.setitem(sys.modules, "prompt_control", prompt_control) + monkeypatch.setitem(sys.modules, "prompt_control.nodes_lazy", nodes_lazy) + return calls diff --git a/tests/test_prompt_encode_style_nodes.py b/tests/test_prompt_encode_style_nodes.py index 7d6da2f..13271ba 100644 --- a/tests/test_prompt_encode_style_nodes.py +++ b/tests/test_prompt_encode_style_nodes.py @@ -29,7 +29,7 @@ def test_prompt_encode_style_node_contract_constants() -> None: """Style-only node constants match the public ComfyUI contract.""" assert PromptEncodeStyle.RETURN_TYPES == ("STRING",) - assert PromptEncodeStyle.RETURN_NAMES == ("style_tag",) + assert PromptEncodeStyle.RETURN_NAMES == ("encode_style",) assert PromptEncodeStyle.FUNCTION == "build" assert PromptEncodeStyle.CATEGORY == "SimpleSyrup/Prompt" @@ -60,18 +60,18 @@ def test_prompt_encode_style_node_builds_style_tag( ) -> None: """Style-only node formats Prompt Control STYLE tags from combo values.""" - (style_tag,) = PromptEncodeStyle().build(encode_style) + (encode_style_text,) = PromptEncodeStyle().build(encode_style) - assert style_tag == expected - assert style_tag.endswith(" ") - assert not style_tag.endswith(" ") + assert encode_style_text == expected + assert encode_style_text.endswith(" ") + assert not encode_style_text.endswith(" ") def test_prompt_encode_style_and_normalization_node_contract_constants() -> None: """Normalization node constants match the public ComfyUI contract.""" assert PromptEncodeStyleAndNormalization.RETURN_TYPES == ("STRING",) - assert PromptEncodeStyleAndNormalization.RETURN_NAMES == ("style_tag",) + assert PromptEncodeStyleAndNormalization.RETURN_NAMES == ("encode_style",) assert PromptEncodeStyleAndNormalization.FUNCTION == "build" assert PromptEncodeStyleAndNormalization.CATEGORY == "SimpleSyrup/Prompt" @@ -110,10 +110,10 @@ def test_prompt_encode_style_and_normalization_node_builds_style_tag( ) -> None: """Normalization node formats Prompt Control STYLE tags from combo values.""" - (style_tag,) = PromptEncodeStyleAndNormalization().build( + (encode_style_text,) = PromptEncodeStyleAndNormalization().build( encode_style, normalization ) - assert style_tag == expected - assert style_tag.endswith(" ") - assert not style_tag.endswith(" ") + assert encode_style_text == expected + assert encode_style_text.endswith(" ") + assert not encode_style_text.endswith(" ") diff --git a/tests/test_registration.py b/tests/test_registration.py index d550fb5..350554f 100644 --- a/tests/test_registration.py +++ b/tests/test_registration.py @@ -15,6 +15,8 @@ from types import ModuleType import pytest +from simple_syrup.runtime.prompt_control_availability import prompt_control_is_available + def test_package_exports_node_mappings() -> None: """Root package import exposes ComfyUI mapping dictionaries.""" @@ -189,6 +191,28 @@ def test_prompt_control_encode_style_clean_break_id_is_removed() -> None: ) +def test_prompt_control_schedule_node_is_legacy_registered_when_available() -> None: + """Prompt-Control schedule node is exported through legacy mappings.""" + + if not prompt_control_is_available(): + pytest.skip("Prompt-Control is not installed in this test environment.") + + package = importlib.import_module("SimpleSyrup") + nodes_module = importlib.import_module("SimpleSyrup.simple_syrup.nodes") + + registered = package.NODE_CLASS_MAPPINGS[ + "SimpleSyrup.ScheduleAndEncodePromptsWithPromptControl" + ] + assert registered.__name__ == "ScheduleAndEncodePromptsWithPromptControl" + assert ( + package.NODE_DISPLAY_NAME_MAPPINGS[ + "SimpleSyrup.ScheduleAndEncodePromptsWithPromptControl" + ] + == "Schedule & Encode Prompts" + ) + assert "ScheduleAndEncodePromptsWithPromptControl" in nodes_module.__all__ + + def test_prompt_segs_with_sam_node_is_registered() -> None: """Prompt SEGS w/ SAM node id maps to its class and display name.""" @@ -530,6 +554,7 @@ def test_v3_entrypoint_registers_tile_and_prompt_control_batch_nodes( "VAEDecodeOptionsV3", "VAEEncodeOptionsV3", "EncodePromptBatchWithPromptControl", + "ScheduleAndEncodePromptsWithPromptControl", ] assert "prompt_control.nodes_lazy" not in sys.modules diff --git a/tests/test_schedule_and_encode_prompts_with_prompt_control_node.py b/tests/test_schedule_and_encode_prompts_with_prompt_control_node.py new file mode 100644 index 0000000..0c54e30 --- /dev/null +++ b/tests/test_schedule_and_encode_prompts_with_prompt_control_node.py @@ -0,0 +1,273 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for the Prompt-Control schedule and encode node schema.""" + +from __future__ import annotations + +from typing import Any + +from simple_syrup.nodes.schedule_and_encode_prompts_with_prompt_control import ( + ScheduleAndEncodePromptsWithPromptControl as LegacyScheduleAndEncode, +) +from simple_syrup.nodes_v3.schedule_and_encode_prompts_with_prompt_control import ( + ScheduleAndEncodePromptsWithPromptControl, +) + + +def test_legacy_schedule_and_encode_prompt_control_node_contract() -> None: + """The legacy node exposes the same workflow-facing contract.""" + + inputs = LegacyScheduleAndEncode.INPUT_TYPES() + + assert LegacyScheduleAndEncode.RETURN_TYPES == ( + "MODEL", + "CONDITIONING,CONDITIONING_BATCH", + "CONDITIONING,CONDITIONING_BATCH", + ) + assert LegacyScheduleAndEncode.RETURN_NAMES == ( + "model", + "positive", + "negative", + ) + assert LegacyScheduleAndEncode.FUNCTION == "execute" + assert LegacyScheduleAndEncode.CATEGORY == "SimpleSyrup/Conditioning" + assert list(inputs["required"]) == [ + "model", + "clip", + "positive_prompt", + "negative_prompt", + ] + assert list(inputs["optional"]) == ["encode_style"] + assert inputs["required"]["model"][0] == "MODEL" + assert inputs["required"]["clip"][0] == "CLIP" + assert inputs["optional"]["encode_style"][0] == "STRING" + assert inputs["optional"]["encode_style"][1]["forceInput"] is True + assert inputs["required"]["positive_prompt"][1]["multiline"] is False + assert inputs["required"]["negative_prompt"][1]["multiline"] is False + + +def test_legacy_schedule_and_encode_prompt_control_execute_delegates( + monkeypatch: Any, +) -> None: + """Legacy node execution delegates to the shared runtime builder.""" + + calls: list[dict[str, Any]] = [] + + class FakeBuilder: + """Runtime builder double.""" + + def build(self, **kwargs: Any) -> str: + """Record builder arguments and return a fixed output.""" + + calls.append(kwargs) + return "legacy-node-output" + + monkeypatch.setattr( + "simple_syrup.nodes.schedule_and_encode_prompts_with_prompt_control." + "PromptControlScheduleEncodeGraphBuilder", + FakeBuilder, + ) + + output = LegacyScheduleAndEncode().execute( + model="model", + clip="clip", + encode_style="STYLE(A1111) ", + positive_prompt="positive", + negative_prompt="negative", + ) + + assert output == "legacy-node-output" + assert calls == [ + { + "model": "model", + "clip": "clip", + "encode_style": "STYLE(A1111) ", + "positive_prompt": "positive", + "negative_prompt": "negative", + } + ] + + +def test_legacy_schedule_and_encode_prompt_control_omits_encode_style( + monkeypatch: Any, +) -> None: + """Legacy node execution no-ops style behavior when the optional socket is empty.""" + + calls: list[dict[str, Any]] = [] + + class FakeBuilder: + """Runtime builder double.""" + + def build(self, **kwargs: Any) -> str: + """Record builder arguments and return a fixed output.""" + + calls.append(kwargs) + return "legacy-node-output" + + monkeypatch.setattr( + "simple_syrup.nodes.schedule_and_encode_prompts_with_prompt_control." + "PromptControlScheduleEncodeGraphBuilder", + FakeBuilder, + ) + + output = LegacyScheduleAndEncode().execute( + model="model", + clip="clip", + positive_prompt="positive", + negative_prompt="negative", + ) + + assert output == "legacy-node-output" + assert calls == [ + { + "model": "model", + "clip": "clip", + "encode_style": "", + "positive_prompt": "positive", + "negative_prompt": "negative", + } + ] + + +def test_schedule_and_encode_prompt_control_node_schema() -> None: + """The v3 node exposes the planned schedule and encode contract.""" + + schema = ScheduleAndEncodePromptsWithPromptControl.define_schema() + + assert schema.node_id == "SimpleSyrup.ScheduleAndEncodePromptsWithPromptControl" + assert schema.display_name == "Schedule & Encode Prompts" + assert schema.enable_expand is True + assert schema.category == "SimpleSyrup/Conditioning" + assert [output.io_type for output in schema.outputs] == [ + "MODEL", + "CONDITIONING,CONDITIONING_BATCH", + "CONDITIONING,CONDITIONING_BATCH", + ] + assert [output.id for output in schema.outputs] == [ + "model", + "positive", + "negative", + ] + assert [input_item.id for input_item in schema.inputs] == [ + "model", + "clip", + "encode_style", + "positive_prompt", + "negative_prompt", + ] + assert schema.inputs[2].optional is True + + +def test_schedule_and_encode_prompt_control_input_types() -> None: + """The finalized v3 schema exposes Comfy-compatible sockets.""" + + inputs = ScheduleAndEncodePromptsWithPromptControl.INPUT_TYPES() + + assert ScheduleAndEncodePromptsWithPromptControl.RETURN_TYPES == [ + "MODEL", + "CONDITIONING,CONDITIONING_BATCH", + "CONDITIONING,CONDITIONING_BATCH", + ] + assert ScheduleAndEncodePromptsWithPromptControl.RETURN_NAMES == [ + "model", + "positive", + "negative", + ] + assert list(inputs["required"]) == [ + "model", + "clip", + "positive_prompt", + "negative_prompt", + ] + assert list(inputs["optional"]) == ["encode_style"] + assert inputs["required"]["model"][0] == "MODEL" + assert inputs["required"]["clip"][0] == "CLIP" + assert inputs["optional"]["encode_style"][0] == "STRING" + assert inputs["optional"]["encode_style"][1]["forceInput"] is True + assert inputs["required"]["positive_prompt"][1]["multiline"] is False + assert inputs["required"]["negative_prompt"][1]["multiline"] is False + + +def test_schedule_and_encode_prompt_control_execute_delegates( + monkeypatch: Any, +) -> None: + """Node execution delegates behavior to the runtime graph builder.""" + + calls: list[dict[str, Any]] = [] + + class FakeBuilder: + """Runtime builder double.""" + + def build(self, **kwargs: Any) -> str: + """Record builder arguments and return a fixed output.""" + + calls.append(kwargs) + return "node-output" + + monkeypatch.setattr( + "simple_syrup.nodes_v3.schedule_and_encode_prompts_with_prompt_control." + "PromptControlScheduleEncodeGraphBuilder", + FakeBuilder, + ) + + output = ScheduleAndEncodePromptsWithPromptControl.execute( + model="model", + clip="clip", + encode_style="STYLE(A1111) ", + positive_prompt="positive", + negative_prompt="negative", + ) + + assert output == "node-output" + assert calls == [ + { + "model": "model", + "clip": "clip", + "encode_style": "STYLE(A1111) ", + "positive_prompt": "positive", + "negative_prompt": "negative", + } + ] + + +def test_schedule_and_encode_prompt_control_omits_encode_style( + monkeypatch: Any, +) -> None: + """Node execution no-ops style behavior when the optional socket is empty.""" + + calls: list[dict[str, Any]] = [] + + class FakeBuilder: + """Runtime builder double.""" + + def build(self, **kwargs: Any) -> str: + """Record builder arguments and return a fixed output.""" + + calls.append(kwargs) + return "node-output" + + monkeypatch.setattr( + "simple_syrup.nodes_v3.schedule_and_encode_prompts_with_prompt_control." + "PromptControlScheduleEncodeGraphBuilder", + FakeBuilder, + ) + + output = ScheduleAndEncodePromptsWithPromptControl.execute( + model="model", + clip="clip", + positive_prompt="positive", + negative_prompt="negative", + ) + + assert output == "node-output" + assert calls == [ + { + "model": "model", + "clip": "clip", + "encode_style": "", + "positive_prompt": "positive", + "negative_prompt": "negative", + } + ] diff --git a/tests/test_tiled_diffusion_sampling_service.py b/tests/test_tiled_diffusion_sampling_service.py index fa616d1..66f8886 100644 --- a/tests/test_tiled_diffusion_sampling_service.py +++ b/tests/test_tiled_diffusion_sampling_service.py @@ -11,6 +11,7 @@ from typing import Any import pytest import torch +from simple_syrup.domain.conditioning_batch import ConditioningBatch from simple_syrup.services.tiled_diffusion_sampling_service import ( TiledDiffusionSamplingService, ) @@ -158,6 +159,63 @@ def test_service_forwards_differential_diffusion_request( assert calls["differential_diffusion"] is True +def test_service_selects_conditioning_batch_per_latent_item( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Batch conditioning is selected before tiled runtime dispatch.""" + + calls: list[dict[str, Any]] = [] + + def fake_multidiffusion(**kwargs: Any) -> dict[str, Any]: + """Record per-item runtime arguments and return marked samples.""" + + calls.append(kwargs) + return { + "samples": torch.full_like( + kwargs["latent_image"]["samples"], + float(len(calls)), + ) + } + + monkeypatch.setattr( + "simple_syrup.services.tiled_diffusion_sampling_service." + "multidiffusion_sampling.sample_multidiffusion", + fake_multidiffusion, + ) + + latent_samples = torch.zeros((2, 4, 4, 4)) + noise_mask = torch.ones((2, 1, 4, 4)) + result = TiledDiffusionSamplingService().sample( + **( + _sample_kwargs(diffusion_mode="multidiffusion") + | { + "positive": ConditioningBatch(("positive-0", "positive-1")), + "negative": ConditioningBatch(("negative-last",)), + "latent_image": { + "samples": latent_samples, + "batch_index": [7, 11], + "noise_mask": noise_mask, + "downscale_ratio_spacial": 2, + }, + } + ) + ) + + assert len(calls) == 2 + assert calls[0]["positive"] == "positive-0" + assert calls[1]["positive"] == "positive-1" + assert calls[0]["negative"] == "negative-last" + assert calls[1]["negative"] == "negative-last" + assert calls[0]["latent_image"]["batch_index"] == [7] + assert calls[1]["latent_image"]["batch_index"] == [11] + assert torch.equal(calls[0]["latent_image"]["noise_mask"], noise_mask[0:1]) + assert torch.equal(calls[1]["latent_image"]["noise_mask"], noise_mask[1:2]) + assert "downscale_ratio_spacial" not in result + assert result["samples"].shape == latent_samples.shape + assert torch.equal(result["samples"][0], torch.full((4, 4, 4), 1.0)) + assert torch.equal(result["samples"][1], torch.full((4, 4, 4), 2.0)) + + def test_invalid_mode_fails_before_runtime_call( monkeypatch: pytest.MonkeyPatch, ) -> None: