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:
@@ -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)
|
||||
@@ -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",
|
||||
|
||||
@@ -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}),
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
@@ -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(" ")
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user