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.
This commit is contained in:
Artificial Sweetener
2026-05-26 17:45:21 -04:00
parent 92c53b3493
commit 346ff8b7c4
20 changed files with 1561 additions and 32 deletions
@@ -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)
+17
View File
@@ -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",
+98 -16
View File
@@ -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."""
@@ -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}),
+2 -2
View File
@@ -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."
@@ -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"
@@ -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,
)
+4
View File
@@ -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,
]
@@ -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,
)
@@ -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)
@@ -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,
)
+107
View File
@@ -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))
@@ -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
+4
View File
@@ -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(
+83
View File
@@ -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 <lora:face:1.0> cinematic <foo>"
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("<lora:a:1.0>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 <lora:face:1.0> cinematic <lora:light:0.5>"
assert extract_lora_tags(text) == "<lora:face:1.0>\n<lora:light:0.5>"
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 <lora:a:1.0> [SEP] hair <lora:b:0.5>",
"[SEP]",
)
assert [chunk.text for chunk in side.chunks] == ["face ", "hair "]
assert [chunk.lora_tags for chunk in side.chunks] == [
"<lora:a:1.0>",
"<lora:b:0.5>",
]
assert side.lora_tags == "<lora:a:1.0>\n<lora:b:0.5>"
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"
@@ -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 <lora:positive:1.0>",
negative_prompt="blur <lora:negative:0.5>",
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": "<lora:positive:1.0>",
},
{
"model": ["lora_positive", 0],
"clip": ["lora_positive", 1],
"text": "<lora:negative:0.5>",
},
]
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 <lora:a:1> [SEP] hair <lora:b:1>",
negative_prompt="blur <lora:c:1> [SEP] noise <lora:d:1>",
)
assert calls["lora"][0]["text"] == "<lora:a:1>\n<lora:b:1>"
assert calls["lora"][1]["text"] == "<lora:c:1>\n<lora:d:1>"
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
+10 -10
View File
@@ -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(" ")
+25
View File
@@ -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
@@ -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",
}
]
@@ -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: