fix(regional): align prompt batches and LoRA hooks

This commit is contained in:
Artificial Sweetener
2026-07-31 22:05:57 -04:00
parent 250ee42ea1
commit 546ee5db01
30 changed files with 1323 additions and 150 deletions
@@ -0,0 +1,94 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Plan authored and global-fallback positions across two prompt sides."""
from __future__ import annotations
from dataclasses import dataclass
from typing import TypeVar
SegmentValue = TypeVar("SegmentValue")
@dataclass(frozen=True)
class PromptSegmentSource:
"""Identify one effective segment's authored source position."""
source_index: int
authored: bool
def resolve(self, values: tuple[SegmentValue, ...]) -> SegmentValue:
"""Return the authored value selected for this effective position."""
return values[self.source_index]
@dataclass(frozen=True)
class PromptSideAlignment:
"""Describe effective positions for one authored prompt side."""
authored_count: int
sources: tuple[PromptSegmentSource, ...]
def materialize(
self,
values: tuple[SegmentValue, ...],
) -> tuple[SegmentValue, ...]:
"""Resolve effective values while rejecting a mismatched authored side."""
if len(values) != self.authored_count:
raise ValueError(
"prompt alignment expected "
f"{self.authored_count} authored segments but received {len(values)}."
)
return tuple(source.resolve(values) for source in self.sources)
@dataclass(frozen=True)
class PromptSegmentAlignment:
"""Store matched effective positions for positive and negative prompts."""
positive: PromptSideAlignment
negative: PromptSideAlignment
@property
def segment_count(self) -> int:
"""Return the shared number of effective prompt positions."""
return len(self.positive.sources)
def build_prompt_segment_alignment(
*,
positive_count: int,
negative_count: int,
) -> PromptSegmentAlignment:
"""Align prompt sides by filling missing positions from each global entry."""
if positive_count < 1 or negative_count < 1:
raise ValueError("each prompt side must contain at least one authored segment.")
segment_count = max(positive_count, negative_count)
return PromptSegmentAlignment(
positive=_build_side_alignment(positive_count, segment_count),
negative=_build_side_alignment(negative_count, segment_count),
)
def _build_side_alignment(
authored_count: int,
segment_count: int,
) -> PromptSideAlignment:
"""Return authored positions followed by global-entry fallback positions."""
return PromptSideAlignment(
authored_count=authored_count,
sources=tuple(
PromptSegmentSource(
source_index=index if index < authored_count else 0,
authored=index < authored_count,
)
for index in range(segment_count)
),
)
@@ -6,8 +6,6 @@
from __future__ import annotations
from collections.abc import Sequence
import torch
from .detailer_masks import gaussian_feather_mask
@@ -43,19 +41,3 @@ def regional_mask(mask_batch: torch.Tensor, index: int) -> torch.Tensor:
if index < 0 or index >= int(mask_batch.shape[0]):
raise IndexError(f"regional mask index {index} is out of range.")
return mask_batch[index : index + 1]
def complementary_global_prompt_mask(
mask_batch: torch.Tensor,
mask_indices: Sequence[int],
regional_prompt_weight: float,
) -> torch.Tensor:
"""Return global influence that recedes across accumulated region coverage."""
if not mask_indices:
raise ValueError("global prompt masking requires at least one regional mask.")
coverage = torch.zeros_like(mask_batch[0:1])
for index in mask_indices:
coverage.add_(regional_mask(mask_batch, index))
coverage.clamp_(0.0, 1.0)
return 1.0 - coverage * regional_prompt_weight
+17 -10
View File
@@ -8,8 +8,8 @@ from __future__ import annotations
from typing import Any, ClassVar
from ..domain.conditioning_batch import split_prompt_batch
from ..runtime.conditioning_encoding import ComfyConditioningEncoder
from ..services.prompt_batch_encoding_service import PromptBatchEncodingService
class EncodePromptBatch:
@@ -23,10 +23,16 @@ class EncodePromptBatch:
)
FUNCTION = "encode"
CATEGORY = "SimpleSyrup/Conditioning"
DESCRIPTION = "Encodes [SEP]-separated prompts into ordered conditioning batches."
DESCRIPTION = (
"Encodes [SEP]-separated prompts into matched conditioning batches, "
"reusing each side's global prompt when a regional entry is missing."
)
SEARCH_ALIASES = ["conditioning batch", "prompt batch", "segs prompts"]
encoder_class: ClassVar[type[ComfyConditioningEncoder]] = ComfyConditioningEncoder
service_class: ClassVar[type[PromptBatchEncodingService]] = (
PromptBatchEncodingService
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
@@ -49,7 +55,8 @@ class EncodePromptBatch:
"default": "",
"multiline": True,
"tooltip": (
"Ordered positive prompt entries separated by [SEP]."
"Ordered positive prompt entries separated by [SEP]; "
"the global entry fills missing positive regions."
),
},
),
@@ -59,7 +66,8 @@ class EncodePromptBatch:
"default": "",
"multiline": True,
"tooltip": (
"Ordered negative prompt entries separated by [SEP]."
"Ordered negative prompt entries separated by [SEP]; "
"the global entry fills missing negative regions."
),
},
),
@@ -82,10 +90,9 @@ class EncodePromptBatch:
) -> tuple[object, object]:
"""Encode positive and negative prompt batches."""
encoder = self.encoder_class()
positive_chunks = split_prompt_batch(positive_prompt, separator)
negative_chunks = split_prompt_batch(negative_prompt, separator)
return (
encoder.encode_batch(clip, positive_chunks),
encoder.encode_batch(clip, negative_chunks),
return self.service_class(self.encoder_class()).encode(
clip=clip,
positive_prompt=positive_prompt,
negative_prompt=negative_prompt,
separator=separator,
)
@@ -31,7 +31,7 @@ class ScheduleAndEncodePromptsWithPromptControl:
CATEGORY = "SimpleSyrup/Conditioning"
DESCRIPTION = (
"Schedules Prompt-Control LoRAs and encodes prompts. [SEP] creates "
"conditioning batches with segment-local LoRA hooks."
"matched conditioning batches using global text for missing regions."
)
SEARCH_ALIASES = ["prompt control", "schedule prompts", "encode prompts"]
@@ -68,7 +68,7 @@ class ScheduleAndEncodePromptsWithPromptControl:
"multiline": False,
"tooltip": (
"Positive Prompt-Control text; [SEP] creates ordered "
"entries with segment-local LoRA hooks."
"entries, and global text fills missing positive regions."
),
},
),
@@ -79,7 +79,7 @@ class ScheduleAndEncodePromptsWithPromptControl:
"multiline": False,
"tooltip": (
"Negative Prompt-Control text; [SEP] creates ordered "
"entries sharing each index's LoRA hooks."
"entries, and global text fills missing negative regions."
),
},
),
+6
View File
@@ -101,16 +101,22 @@ def get_nodes() -> list[type[object]]:
if not prompt_control_is_available():
return nodes
from .attach_regional_global_conditioning import (
AttachRegionalGlobalConditioningV3,
)
from .encode_prompt_batch_with_prompt_control import (
EncodePromptBatchWithPromptControl,
)
from .prepare_regional_lora_hooks import PrepareRegionalLoraHooksV3
from .schedule_and_encode_prompts_with_prompt_control import (
ScheduleAndEncodePromptsWithPromptControl,
)
return [
*nodes,
AttachRegionalGlobalConditioningV3,
EncodePromptBatchWithPromptControl,
PrepareRegionalLoraHooksV3,
ScheduleAndEncodePromptsWithPromptControl,
]
@@ -0,0 +1,71 @@
# 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 internal node for regional global-prompt companions."""
from __future__ import annotations
from importlib import import_module
from typing import TYPE_CHECKING, Any, ClassVar
from ..runtime.regional_conditioning_companion import attach_global_companion
if TYPE_CHECKING:
class _ComfyNodeBase:
"""Type-checking base for Comfy v3 nodes."""
RETURN_TYPES: ClassVar[list[str]]
RETURN_NAMES: ClassVar[list[str]]
else:
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
class AttachRegionalGlobalConditioningV3(_ComfyNodeBase):
"""Attach the hooked global share required by one regional LoRA."""
@classmethod
def define_schema(cls) -> Any:
"""Declare the internal companion-conditioning schema."""
return _comfy_io.Schema(
node_id="SimpleSyrup.AttachRegionalGlobalConditioning",
display_name="Attach Regional Global Conditioning",
category="SimpleSyrup/Internal",
description=(
"Keeps a region's global and local prompt shares under the same "
"regional LoRA."
),
inputs=[
_comfy_io.Conditioning.Input(
"conditioning",
tooltip="Regional prompt conditioning to preserve.",
),
_comfy_io.Conditioning.Input(
"global_conditioning",
tooltip=(
"Global prompt encoded with the same regional LoRA hooks."
),
),
],
outputs=[
_comfy_io.Conditioning.Output(
"conditioning",
tooltip="Regional conditioning carrying its hooked global share.",
)
],
)
@classmethod
def execute(
cls,
conditioning: object,
global_conditioning: object,
) -> tuple[object]:
"""Attach the global companion to one regional conditioning."""
return (attach_global_companion(conditioning, global_conditioning),)
@@ -47,8 +47,8 @@ class EncodePromptBatchWithPromptControl(_ComfyNodeBase):
enable_expand=True,
category="SimpleSyrup/Conditioning",
description=(
"Encodes [SEP]-separated prompts into per-segment Prompt Control "
"conditioning batches with segment-local LoRA hooks."
"Encodes [SEP]-separated prompts into matched Prompt Control "
"batches, reusing each side's global text for missing regions."
),
inputs=[
_comfy_io.Clip.Input(
@@ -65,7 +65,8 @@ class EncodePromptBatchWithPromptControl(_ComfyNodeBase):
default="",
tooltip=(
"Positive Prompt Control prompts in positional order; each "
"segment keeps its aligned LoRA hooks."
"segment keeps its aligned LoRA hooks, and global text "
"fills missing positive regions."
),
),
_comfy_io.String.Input(
@@ -74,7 +75,8 @@ class EncodePromptBatchWithPromptControl(_ComfyNodeBase):
default="",
tooltip=(
"Negative Prompt Control prompts in positional order; each "
"segment shares hooks with the matching positive index."
"segment shares hooks with the matching positive index, "
"and global text fills missing negative regions."
),
),
_comfy_io.String.Input(
@@ -0,0 +1,72 @@
# 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 internal node for regional LoRA CLIP preparation."""
from __future__ import annotations
from importlib import import_module
from typing import TYPE_CHECKING, Any, ClassVar
from ..runtime.regional_lora_hooks import prepare_regional_lora_clip
if TYPE_CHECKING:
class _ComfyNodeBase:
"""Type-checking base for Comfy v3 nodes."""
RETURN_TYPES: ClassVar[list[str]]
RETURN_NAMES: ClassVar[list[str]]
else:
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
class PrepareRegionalLoraHooksV3(_ComfyNodeBase):
"""Prepare only the CLIP work required by regional LoRA hooks."""
@classmethod
def define_schema(cls) -> Any:
"""Declare the internal hook conversion schema."""
return _comfy_io.Schema(
node_id="SimpleSyrup.PrepareRegionalLoraHooks",
display_name="Prepare Regional LoRA Hooks",
category="SimpleSyrup/Internal",
description=(
"Keeps model-only regional LoRAs away from the text encoder while "
"preparing text-encoder patches when present."
),
inputs=[
_comfy_io.Clip.Input(
"clip",
tooltip="CLIP model used to encode this regional prompt.",
),
_comfy_io.Hooks.Input(
"hooks",
tooltip=("Prompt-Control LoRA hooks to inspect and prepare."),
),
],
outputs=[
_comfy_io.Clip.Output(
"clip",
tooltip=(
"Original CLIP for model-only LoRAs, or a hook-prepared "
"CLIP when text-encoder weights are present."
),
),
_comfy_io.Hooks.Output(
"hooks",
tooltip=("Regional LoRA hooks to attach after prompt encoding."),
),
],
)
@classmethod
def execute(cls, clip: Any, hooks: object) -> tuple[Any, object]:
"""Return the appropriate encoding CLIP and unchanged model hooks."""
return prepare_regional_lora_clip(clip, hooks)
@@ -50,7 +50,8 @@ class ScheduleAndEncodePromptsWithPromptControl(_ComfyNodeBase):
category="SimpleSyrup/Conditioning",
description=(
"Schedules Prompt-Control LoRAs and encodes prompts. With [SEP], "
"each conditioning entry keeps only its segment's LoRA hooks."
"both sides are matched using global text for missing regions and "
"share each segment's LoRA hooks."
),
inputs=[
_comfy_io.Model.Input(
@@ -85,7 +86,8 @@ class ScheduleAndEncodePromptsWithPromptControl(_ComfyNodeBase):
default="",
tooltip=(
"Positive Prompt-Control text; [SEP] creates ordered "
"conditioning entries with segment-local LoRA hooks."
"conditioning entries, and global text fills missing "
"positive regions."
),
),
_comfy_io.String.Input(
@@ -94,7 +96,8 @@ class ScheduleAndEncodePromptsWithPromptControl(_ComfyNodeBase):
default="",
tooltip=(
"Negative Prompt-Control text; [SEP] creates ordered "
"conditioning entries sharing each index's LoRA hooks."
"conditioning entries, and global text fills missing "
"negative regions."
),
),
],
@@ -12,7 +12,10 @@ from ..domain.prompt_control_prompt import PreparedPromptSide
from ..services.prompt_control_segment_planning_service import (
PromptControlSegmentPlanningService,
)
from .prompt_control_graph_adapter import PromptControlGraphAdapter
from .prompt_control_graph_adapter import (
PromptControlGraphAdapter,
RegionalSegmentEncoding,
)
PROMPT_CONTROL_MISSING_MESSAGE = (
"Encode Prompt Batch w/ Prompt Control requires comfyui-prompt-control. "
@@ -42,7 +45,7 @@ class PromptControlBatchGraphBuilder:
)
adapter = self.graph_adapter_class.load(PROMPT_CONTROL_MISSING_MESSAGE)
expand: dict[str, dict[str, Any]] = {}
segment_clips = tuple(
segment_encodings = tuple(
adapter.clip_with_hooks(
clip=clip,
lora_tags=hook.lora_tags,
@@ -53,14 +56,14 @@ class PromptControlBatchGraphBuilder:
)
positive = self._encode_side(
plan.positive,
segment_clips=segment_clips,
segment_encodings=segment_encodings,
adapter=adapter,
expand=expand,
label="positive",
)
negative = self._encode_side(
plan.negative,
segment_clips=segment_clips,
segment_encodings=segment_encodings,
adapter=adapter,
expand=expand,
label="negative",
@@ -71,7 +74,7 @@ class PromptControlBatchGraphBuilder:
self,
side: PreparedPromptSide,
*,
segment_clips: tuple[Any, ...],
segment_encodings: tuple[RegionalSegmentEncoding, ...],
adapter: PromptControlGraphAdapter,
expand: dict[str, dict[str, Any]],
label: str,
@@ -80,13 +83,29 @@ class PromptControlBatchGraphBuilder:
outputs = [
adapter.encode_segment(
clip=segment_clips[index],
segment=segment_encodings[index],
text=chunk.text,
expand=expand,
label=f"{label} segment {index}",
)
for index, chunk in enumerate(side.chunks)
]
for index in range(1, len(outputs)):
segment = segment_encodings[index]
if segment.hooks is None:
continue
global_companion = adapter.encode_segment(
segment=segment,
text=side.chunks[0].text,
expand=expand,
label=f"{label} segment {index} global companion",
)
outputs[index] = adapter.attach_global_companion(
conditioning=outputs[index],
global_conditioning=global_companion,
expand=expand,
label=f"{label} segment {index}",
)
return adapter.pack_conditionings(
outputs,
expand=expand,
@@ -7,12 +7,21 @@
from __future__ import annotations
import sys
from dataclasses import dataclass
from importlib import import_module
from typing import Any, cast
from .prompt_control_availability import find_prompt_control_install
@dataclass(frozen=True)
class RegionalSegmentEncoding:
"""Carry one segment's encoding CLIP and post-encode model hooks."""
clip: Any
hooks: Any | None = None
class PromptControlGraphAdapter:
"""Build hook-aware lazy graph fragments behind one runtime boundary."""
@@ -76,7 +85,7 @@ class PromptControlGraphAdapter:
def encode_segment(
self,
*,
clip: Any,
segment: RegionalSegmentEncoding,
text: str,
expand: dict[str, dict[str, Any]],
label: str,
@@ -84,7 +93,7 @@ class PromptControlGraphAdapter:
"""Encode one segment with a previously prepared CLIP link."""
output = self._lazy_nodes.PCLazyTextEncodeAdvanced.execute(
clip=clip,
clip=segment.clip,
text=text,
tags="",
start=0.0,
@@ -92,7 +101,24 @@ class PromptControlGraphAdapter:
num_steps=0,
)
self.merge_expand(expand, output.expand, f"{label} text encoding")
return output.args[0]
conditioning = output.args[0]
if segment.hooks is None:
return conditioning
graph = self._graph_utils.GraphBuilder()
conditioned = graph.node(
"ConditioningSetProperties",
cond_NEW=conditioning,
hooks=segment.hooks,
strength=1.0,
set_cond_area="default",
)
self.merge_expand(
expand,
cast(dict[str, dict[str, Any]], graph.finalize()),
f"{label} model hook attachment",
)
return conditioned.out(0)
def clip_with_hooks(
self,
@@ -101,26 +127,30 @@ class PromptControlGraphAdapter:
lora_tags: str,
expand: dict[str, dict[str, Any]],
label: str,
) -> Any:
"""Return a CLIP link sharing one segment's hooks across prompt sides."""
) -> RegionalSegmentEncoding:
"""Return separate encoding-CLIP and post-encode model-hook links."""
if not lora_tags:
return clip
return RegionalSegmentEncoding(clip=clip)
graph = self._graph_utils.GraphBuilder()
hooks = graph.node("PCLoraHooksFromText", text=lora_tags)
hooked_clip = graph.node(
"SetClipHooks",
parsed_hooks = graph.node(
"PCLoraHooksFromText",
text=lora_tags,
)
regional_hooks = graph.node(
"SimpleSyrup.PrepareRegionalLoraHooks",
clip=clip,
hooks=hooks.out(0),
apply_to_conds=True,
schedule_clip=True,
hooks=parsed_hooks.out(0),
)
self.merge_expand(
expand,
cast(dict[str, dict[str, Any]], graph.finalize()),
f"{label} LoRA hooks",
)
return hooked_clip.out(0)
return RegionalSegmentEncoding(
clip=regional_hooks.out(0),
hooks=regional_hooks.out(1),
)
def pack_conditionings(
self,
@@ -152,6 +182,29 @@ class PromptControlGraphAdapter:
)
return current.out(0)
def attach_global_companion(
self,
*,
conditioning: Any,
global_conditioning: Any,
expand: dict[str, dict[str, Any]],
label: str,
) -> Any:
"""Attach one hooked global prompt share to a regional conditioning."""
graph = self._graph_utils.GraphBuilder()
attached = graph.node(
"SimpleSyrup.AttachRegionalGlobalConditioning",
conditioning=conditioning,
global_conditioning=global_conditioning,
)
self.merge_expand(
expand,
cast(dict[str, dict[str, Any]], graph.finalize()),
f"{label} global companion attachment",
)
return attached.out(0)
def merge_expand(
self,
target: dict[str, dict[str, Any]],
@@ -13,7 +13,10 @@ from ..services.prompt_control_segment_planning_service import (
PromptControlSegmentPlan,
PromptControlSegmentPlanningService,
)
from .prompt_control_graph_adapter import PromptControlGraphAdapter
from .prompt_control_graph_adapter import (
PromptControlGraphAdapter,
RegionalSegmentEncoding,
)
PROMPT_CONTROL_MISSING_MESSAGE = (
"Schedule & Encode Prompts requires comfyui-prompt-control. "
@@ -52,7 +55,7 @@ class PromptControlScheduleEncodeGraphBuilder:
adapter=adapter,
expand=expand,
)
segment_clips = self._segment_clips(
segment_encodings = self._segment_encodings(
plan=plan,
clip=encoding_clip,
adapter=adapter,
@@ -60,7 +63,7 @@ class PromptControlScheduleEncodeGraphBuilder:
)
positive = self._encode_side(
plan.positive,
segment_clips=segment_clips,
segment_encodings=segment_encodings,
encode_style=encode_style,
adapter=adapter,
expand=expand,
@@ -68,7 +71,7 @@ class PromptControlScheduleEncodeGraphBuilder:
)
negative = self._encode_side(
plan.negative,
segment_clips=segment_clips,
segment_encodings=segment_encodings,
encode_style=encode_style,
adapter=adapter,
expand=expand,
@@ -102,18 +105,18 @@ class PromptControlScheduleEncodeGraphBuilder:
expand=expand,
)
def _segment_clips(
def _segment_encodings(
self,
*,
plan: PromptControlSegmentPlan,
clip: Any,
adapter: PromptControlGraphAdapter,
expand: dict[str, dict[str, Any]],
) -> tuple[Any, ...]:
"""Create one shared hooked CLIP link per batched segment index."""
) -> tuple[RegionalSegmentEncoding, ...]:
"""Create one shared encoding context per batched segment index."""
if not plan.is_batched:
return (clip,)
return (RegionalSegmentEncoding(clip=clip),)
return tuple(
adapter.clip_with_hooks(
clip=clip,
@@ -128,7 +131,7 @@ class PromptControlScheduleEncodeGraphBuilder:
self,
side: PreparedPromptSide,
*,
segment_clips: tuple[Any, ...],
segment_encodings: tuple[RegionalSegmentEncoding, ...],
encode_style: str,
adapter: PromptControlGraphAdapter,
expand: dict[str, dict[str, Any]],
@@ -138,13 +141,29 @@ class PromptControlScheduleEncodeGraphBuilder:
outputs = [
adapter.encode_segment(
clip=segment_clips[index],
segment=segment_encodings[index],
text=apply_encode_style(encode_style, chunk.text),
expand=expand,
label=f"{label} segment {index}",
)
for index, chunk in enumerate(side.chunks)
]
for index in range(1, len(outputs)):
segment = segment_encodings[index]
if segment.hooks is None:
continue
global_companion = adapter.encode_segment(
segment=segment,
text=apply_encode_style(encode_style, side.chunks[0].text),
expand=expand,
label=f"{label} segment {index} global companion",
)
outputs[index] = adapter.attach_global_companion(
conditioning=outputs[index],
global_conditioning=global_companion,
expand=expand,
label=f"{label} segment {index}",
)
return adapter.pack_conditionings(
outputs,
expand=expand,
@@ -0,0 +1,64 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Carry hooked global conditioning alongside one regional conditioning entry."""
from __future__ import annotations
from typing import Any, TypeAlias
Conditioning: TypeAlias = list[list[Any]]
_GLOBAL_COMPANION_KEY = "simple_syrup.regional_global_companion"
def attach_global_companion(
conditioning: object,
global_conditioning: object,
) -> Conditioning:
"""Attach one hooked global conditioning without changing Comfy's public type."""
local = _validate_conditioning(conditioning, "regional conditioning")
companion = _validate_conditioning(
global_conditioning,
"regional global companion",
)
attached = _copy_conditioning(local)
attached[0][1][_GLOBAL_COMPANION_KEY] = _copy_conditioning(companion)
return attached
def detach_global_companion(
conditioning: Conditioning,
) -> tuple[Conditioning, Conditioning | None]:
"""Remove and return an attached global companion before Comfy sampling."""
detached = _copy_conditioning(conditioning)
companion = detached[0][1].pop(_GLOBAL_COMPANION_KEY, None)
if companion is None:
return detached, None
return detached, _validate_conditioning(
companion,
"regional global companion",
)
def _validate_conditioning(value: object, name: str) -> Conditioning:
"""Validate the Comfy conditioning container used by the internal graph node."""
if not isinstance(value, list) or not value:
raise TypeError(f"{name} must be a non-empty CONDITIONING value.")
for index, item in enumerate(value):
if (
not isinstance(item, list | tuple)
or len(item) != 2
or not isinstance(item[1], dict)
):
raise ValueError(f"{name} item {index} must contain a tensor and metadata.")
return value
def _copy_conditioning(conditioning: Conditioning) -> Conditioning:
"""Copy conditioning containers and metadata without cloning tensors."""
return [[item[0], dict(item[1])] for item in conditioning]
@@ -0,0 +1,60 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Prepare regional LoRA hooks without needlessly cloning text encoders."""
from __future__ import annotations
import logging
from importlib import import_module
from typing import Any
LOGGER = logging.getLogger(__name__)
def prepare_regional_lora_clip(clip: Any, hooks: object) -> tuple[Any, object]:
"""Prepare CLIP only when at least one hook contains CLIP-compatible weights."""
comfy_hooks = import_module("comfy.hooks")
if not isinstance(hooks, comfy_hooks.HookGroup):
raise TypeError("regional LoRA hooks must be a Comfy HOOKS value.")
matching_hook_count = _count_clip_patch_hooks(clip, hooks, comfy_hooks)
if matching_hook_count == 0:
LOGGER.debug(
"Regional LoRA hook preparation preserved the original CLIP because "
"the hook group contains no matching text-encoder patches."
)
return clip, hooks
LOGGER.debug(
"Regional LoRA hook preparation created a scheduled CLIP for %d "
"text-encoder hook group entries.",
matching_hook_count,
)
prepared_clip = clip.clone(disable_dynamic=True)
prepared_clip.patcher.forced_hooks = hooks.clone()
prepared_clip.use_clip_schedule = True
prepared_clip.patcher.register_all_hook_patches(
hooks,
comfy_hooks.create_target_dict(comfy_hooks.EnumWeightTarget.Clip),
)
return prepared_clip, hooks
def _count_clip_patch_hooks(clip: Any, hooks: Any, comfy_hooks: Any) -> int:
"""Count enabled weight hooks that resolve to at least one CLIP model key."""
comfy_lora = import_module("comfy.lora")
key_map = comfy_lora.model_lora_keys_clip(clip.cond_stage_model, {})
matching_hook_count = 0
for hook in hooks.get_type(comfy_hooks.EnumHookType.Weight):
if hook._strength_clip == 0.0:
continue
if hook.need_weight_init:
loaded = comfy_lora.load_lora(hook.weights, key_map, log_missing=False)
else:
loaded = hook.weights_clip
if loaded:
matching_hook_count += 1
return matching_hook_count
@@ -0,0 +1,59 @@
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
# Copyright (C) 2026 Artificial Sweetener and contributors
# SPDX-License-Identifier: AGPL-3.0-or-later
"""Encode aligned standard positive and negative prompt batches."""
from __future__ import annotations
from typing import Any, Protocol
from ..domain.conditioning_batch import ConditioningBatch, split_prompt_batch
from ..domain.prompt_segment_alignment import build_prompt_segment_alignment
class PromptBatchEncoder(Protocol):
"""Encode ordered prompt text through a Comfy-compatible adapter."""
def encode_batch(
self,
clip: Any,
chunks: tuple[str, ...],
) -> ConditioningBatch:
"""Return one conditioning entry per ordered prompt chunk."""
class PromptBatchEncodingService:
"""Align two SEP prompt sides before delegating Comfy text encoding."""
def __init__(self, encoder: PromptBatchEncoder) -> None:
"""Store the runtime encoding adapter for one node execution."""
self._encoder = encoder
def encode(
self,
*,
clip: Any,
positive_prompt: str,
negative_prompt: str,
separator: str,
) -> tuple[ConditioningBatch, ConditioningBatch]:
"""Return positive and negative batches with matched segment counts."""
positive_chunks = split_prompt_batch(positive_prompt, separator)
negative_chunks = split_prompt_batch(negative_prompt, separator)
alignment = build_prompt_segment_alignment(
positive_count=len(positive_chunks),
negative_count=len(negative_chunks),
)
return (
self._encoder.encode_batch(
clip,
alignment.positive.materialize(positive_chunks),
),
self._encoder.encode_batch(
clip,
alignment.negative.materialize(negative_chunks),
),
)
@@ -9,10 +9,15 @@ from __future__ import annotations
from dataclasses import dataclass
from ..domain.prompt_control_prompt import (
PreparedPromptChunk,
PreparedPromptSide,
PromptSegmentHookPlan,
prepare_prompt_side,
)
from ..domain.prompt_segment_alignment import (
PromptSideAlignment,
build_prompt_segment_alignment,
)
@dataclass(frozen=True)
@@ -40,16 +45,21 @@ class PromptControlSegmentPlanningService:
negative_prompt: str,
separator: str,
) -> PromptControlSegmentPlan:
"""Return a deterministic plan without padding either prompt side."""
"""Return matched prompt sides and one shared hook plan per position."""
positive = prepare_prompt_side(positive_prompt, separator)
negative = prepare_prompt_side(negative_prompt, separator)
hook_count = max(len(positive.chunks), len(negative.chunks))
authored_positive = prepare_prompt_side(positive_prompt, separator)
authored_negative = prepare_prompt_side(negative_prompt, separator)
alignment = build_prompt_segment_alignment(
positive_count=len(authored_positive.chunks),
negative_count=len(authored_negative.chunks),
)
positive = self._align_side(authored_positive, alignment.positive)
negative = self._align_side(authored_negative, alignment.negative)
hooks = tuple(
PromptSegmentHookPlan(
lora_tags=self._combined_lora_tags(positive, negative, index)
)
for index in range(hook_count)
for index in range(alignment.segment_count)
)
return PromptControlSegmentPlan(
positive=positive,
@@ -57,6 +67,34 @@ class PromptControlSegmentPlanningService:
hooks=hooks,
)
def _align_side(
self,
side: PreparedPromptSide,
alignment: PromptSideAlignment,
) -> PreparedPromptSide:
"""Build global-context regions without repeating global LoRA tags."""
chunks = tuple(
self._effective_chunk(
authored=source.authored,
source=source.resolve(side.chunks),
)
for source in alignment.sources
)
return PreparedPromptSide(chunks=chunks)
def _effective_chunk(
self,
*,
authored: bool,
source: PreparedPromptChunk,
) -> PreparedPromptChunk:
"""Strip tags only when a chunk was synthesized from global text."""
if authored:
return source
return PreparedPromptChunk(text=source.text, lora_tags="")
def _combined_lora_tags(
self,
positive: PreparedPromptSide,
@@ -66,8 +104,8 @@ class PromptControlSegmentPlanningService:
"""Join existing positive then negative tags for one segment index."""
tags: list[str] = []
if index < len(positive.chunks) and positive.chunks[index].lora_tags:
if positive.chunks[index].lora_tags:
tags.append(positive.chunks[index].lora_tags)
if index < len(negative.chunks) and negative.chunks[index].lora_tags:
if negative.chunks[index].lora_tags:
tags.append(negative.chunks[index].lora_tags)
return "\n".join(tags)
@@ -16,10 +16,10 @@ from ..domain.regional_prompting import (
validate_regional_prompt_weight,
)
from ..masking.regional_prompt_masks import (
complementary_global_prompt_mask,
prepare_regional_mask_batch,
regional_mask,
)
from ..runtime.regional_conditioning_companion import detach_global_companion
from ..shared.logging import get_logger
Conditioning: TypeAlias = list[list[Any]]
@@ -90,34 +90,47 @@ class RegionalConditioningService:
entries[0],
input_name=f"{input_name} global",
)
if not plan.pairs or regional_prompt_weight == 0.0:
if not plan.pairs:
return self._copy_conditioning(global_conditioning)
global_mask = complementary_global_prompt_mask(
mask_batch,
tuple(pair.mask_index for pair in plan.pairs),
regional_prompt_weight,
)
assembled = self._with_mask(
global_conditioning,
global_mask,
mask_strength=1.0,
)
assembled = self._as_default(global_conditioning)
for pair in plan.pairs:
conditioning = self._validate_conditioning(
entries[pair.conditioning_index],
input_name=(f"{input_name} regional entry {pair.conditioning_index}"),
)
conditioning, global_companion = detach_global_companion(conditioning)
mask = regional_mask(mask_batch, pair.mask_index)
assembled.extend(
self._with_mask(
conditioning,
mask,
mask_strength=regional_prompt_weight,
if global_companion is not None and regional_prompt_weight < 1.0:
assembled.extend(
self._with_mask(
global_companion,
mask,
mask_strength=1.0 - regional_prompt_weight,
)
)
)
if regional_prompt_weight > 0.0:
assembled.extend(
self._with_mask(
conditioning,
mask,
mask_strength=regional_prompt_weight,
)
)
if len(assembled) == len(global_conditioning):
return self._copy_conditioning(global_conditioning)
return assembled
def _as_default(self, conditioning: Conditioning) -> Conditioning:
"""Mark global conditioning to fill Comfy's remaining regional weight."""
default_conditioning: Conditioning = []
for item in conditioning:
metadata = dict(item[1])
metadata["default"] = True
default_conditioning.append([item[0], metadata])
return default_conditioning
def _validate_conditioning(
self,
value: object,
+27 -4
View File
@@ -25,6 +25,7 @@ def test_encode_prompt_batch_contract() -> None:
)
assert EncodePromptBatch.RETURN_NAMES == ("positive", "negative")
assert EncodePromptBatch.CATEGORY == "SimpleSyrup/Conditioning"
assert "global" in EncodePromptBatch.DESCRIPTION.lower()
assert list(inputs["required"]) == [
"clip",
"positive_prompt",
@@ -35,12 +36,14 @@ def test_encode_prompt_batch_contract() -> None:
assert inputs["required"]["positive_prompt"][1]["default"] == ""
assert inputs["required"]["negative_prompt"][1]["default"] == ""
assert inputs["required"]["separator"][1]["default"] == "[SEP]"
assert "global" in inputs["required"]["positive_prompt"][1]["tooltip"].lower()
assert "global" in inputs["required"]["negative_prompt"][1]["tooltip"].lower()
def test_encode_prompt_batch_splits_and_encodes_each_chunk(
def test_encode_prompt_batch_aligns_missing_negative_chunks_to_global(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Standard encoder batches positive and negative chunks independently."""
"""Standard encoder creates matched negative regions from global text."""
monkeypatch.setattr(EncodePromptBatch, "encoder_class", _FakeEncoder)
@@ -54,7 +57,7 @@ def test_encode_prompt_batch_splits_and_encodes_each_chunk(
assert isinstance(positive, ConditioningBatch)
assert isinstance(negative, ConditioningBatch)
assert positive.entries == ("clip:face", "clip:hair")
assert negative.entries == ("clip:blur",)
assert negative.entries == ("clip:blur", "clip:blur")
def test_encode_prompt_batch_encodes_blank_and_empty_chunks(
@@ -73,10 +76,30 @@ def test_encode_prompt_batch_encodes_blank_and_empty_chunks(
assert isinstance(positive, ConditioningBatch)
assert isinstance(negative, ConditioningBatch)
assert positive.entries == ("clip:",)
assert positive.entries == ("clip:", "clip:")
assert negative.entries == ("clip:bad", "clip:")
def test_encode_prompt_batch_uses_global_positive_for_missing_regions(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Negative-authored regions receive matched global positive entries."""
monkeypatch.setattr(EncodePromptBatch, "encoder_class", _FakeEncoder)
positive, negative = EncodePromptBatch().encode(
clip="clip",
positive_prompt="subject",
negative_prompt="bad [SEP] hands [SEP] text",
separator="[SEP]",
)
assert isinstance(positive, ConditioningBatch)
assert isinstance(negative, ConditioningBatch)
assert positive.entries == ("clip:subject",) * 3
assert negative.entries == ("clip:bad", "clip:hands", "clip:text")
class _FakeEncoder:
"""Fake prompt encoder for node tests."""
@@ -20,6 +20,7 @@ def test_prompt_control_prompt_batch_node_schema() -> None:
assert schema.display_name == "Encode Prompt Batch w/ Prompt Control"
assert schema.enable_expand is True
assert schema.category == "SimpleSyrup/Conditioning"
assert "global" in schema.description.lower()
assert [output.io_type for output in schema.outputs] == [
"CONDITIONING_BATCH",
"CONDITIONING_BATCH",
@@ -48,3 +49,5 @@ def test_prompt_control_prompt_batch_input_types() -> None:
]
assert inputs["required"]["clip"][0] == "CLIP"
assert inputs["required"]["separator"][1]["default"] == "[SEP]"
assert "global" in inputs["required"]["positive_prompt"][1]["tooltip"].lower()
assert "global" in inputs["required"]["negative_prompt"][1]["tooltip"].lower()
+111 -5
View File
@@ -103,21 +103,81 @@ def test_prompt_control_batch_graph_attaches_segment_local_lora_hooks(
)
assert output.expand is not None
hook_nodes = [
parsed_hook_nodes = [
node
for node in output.expand.values()
if node["class_type"] == "PCLoraHooksFromText"
]
assert [node["inputs"]["text"] for node in hook_nodes] == [
assert [node["inputs"]["text"] for node in parsed_hook_nodes] == [
"<lora:a:1>\n<lora:c:1>",
"<lora:b:1>",
]
assert [call["text"] for call in calls] == ["face ", "hair ", "blur ", "noise"]
assert calls[0]["clip"] == calls[2]["clip"]
assert calls[1]["clip"] == calls[3]["clip"]
regional_hook_nodes = [
node
for node in output.expand.values()
if node["class_type"] == "SimpleSyrup.PrepareRegionalLoraHooks"
]
assert len(regional_hook_nodes) == 2
assert not any(
node["class_type"] == "SetClipHooks" for node in output.expand.values()
)
attachment_nodes = [
node
for node in output.expand.values()
if node["class_type"] == "ConditioningSetProperties"
]
assert len(attachment_nodes) == 6
companion_nodes = [
node
for node in output.expand.values()
if node["class_type"] == "SimpleSyrup.AttachRegionalGlobalConditioning"
]
assert len(companion_nodes) == 2
assert [call["text"] for call in calls] == [
"face ",
"hair ",
"face ",
"blur ",
"noise",
"blur ",
]
assert calls[0]["clip"] == calls[3]["clip"]
assert calls[1]["clip"] == calls[4]["clip"]
assert calls[0]["clip"] != calls[1]["clip"]
def test_prompt_control_batch_graph_matches_missing_negative_with_region_hooks(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A sole global negative is re-encoded under each regional hook plan."""
calls = _install_fake_prompt_control(monkeypatch)
graph_utils = import_module("comfy_execution.graph_utils")
graph_utils.GraphBuilder.set_default_prefix("REGRESSION", 0, 0)
PromptControlBatchGraphBuilder().build(
clip=[0, 0],
positive_prompt=("global [SEP] left <lora:regional:1> [SEP] right"),
negative_prompt="bad quality",
separator="[SEP]",
)
assert [call["text"] for call in calls] == [
"global",
"left ",
"right",
"global",
"bad quality",
"bad quality",
"bad quality",
"bad quality",
]
assert calls[1]["clip"] == calls[5]["clip"]
assert calls[3]["clip"] == calls[7]["clip"]
assert calls[0]["clip"] == calls[4]["clip"]
assert calls[2]["clip"] == calls[6]["clip"]
def test_prompt_control_batch_graph_reports_missing_prompt_control(
monkeypatch: pytest.MonkeyPatch,
) -> None:
@@ -152,6 +212,51 @@ def _install_fake_prompt_control(
nodes_lazy = ModuleType("prompt_control.nodes_lazy")
calls: list[dict[str, Any]] = []
class FakePCLazyLoraLoaderAdvanced:
"""Expand segment LoRAs through Comfy's native hook nodes."""
@staticmethod
def execute(
model: Any,
clip: Any,
text: str,
apply_hooks: bool,
tags: str,
start: float,
end: float,
num_steps: int,
) -> Any:
"""Return a hooked CLIP link built from native graph components."""
assert model is None
assert apply_hooks is True
assert tags == ""
assert start == 0.0
assert end == 1.0
assert num_steps == 0
graph_utils = import_module("comfy_execution.graph_utils")
io = import_module("comfy_api.latest").io
graph = graph_utils.GraphBuilder()
hooks = graph.node(
"CreateHookLora",
lora_name=text,
strength_model=1.0,
strength_clip=1.0,
)
hooked_clip = graph.node(
"SetClipHooks",
clip=clip,
hooks=hooks.out(0),
apply_to_conds=True,
schedule_clip=True,
)
return io.NodeOutput(
None,
hooked_clip.out(0),
hooks.out(0),
expand=graph.finalize(),
)
class FakePCLazyTextEncodeAdvanced:
"""Graph-expanding stand-in for Prompt Control's lazy text encoder."""
@@ -178,6 +283,7 @@ def _install_fake_prompt_control(
)
return io.NodeOutput(node.out(0), expand=graph.finalize())
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)
@@ -60,12 +60,12 @@ def test_schedule_encode_graph_builds_single_conditioning_outputs(
]
def test_schedule_encode_graph_packs_only_multichunk_sides(
def test_schedule_encode_graph_packs_both_sides_to_matched_segment_counts(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A side with multiple chunks becomes a conditioning batch."""
"""A missing negative region is encoded and packed from global text."""
_install_fake_prompt_control(monkeypatch)
calls = _install_fake_prompt_control(monkeypatch)
output = PromptControlScheduleEncodeGraphBuilder().build(
model=["model", 0],
@@ -75,7 +75,13 @@ def test_schedule_encode_graph_packs_only_multichunk_sides(
)
assert output.args[1] != ["encode_0", 0]
assert output.args[2] == ["encode_2", 0]
assert output.args[2] != ["encode_2", 0]
assert [call["text"] for call in calls["encode"]] == [
"face",
"hair",
"blur",
"blur",
]
assert output.expand is not None
pack_nodes = [
node
@@ -85,9 +91,41 @@ def test_schedule_encode_graph_packs_only_multichunk_sides(
assert [node["class_type"] for node in pack_nodes] == [
"SimpleSyrup.ConditioningBatchStart",
"SimpleSyrup.ConditioningBatchAppend",
"SimpleSyrup.ConditioningBatchStart",
"SimpleSyrup.ConditioningBatchAppend",
]
def test_schedule_encode_graph_matches_lora_region_on_both_cfg_sides(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Regional LoRA hooks apply to positive and synthesized negative entries."""
calls = _install_fake_prompt_control(monkeypatch)
PromptControlScheduleEncodeGraphBuilder().build(
model=["model", 0],
clip=["clip", 0],
positive_prompt="global [SEP] left <lora:regional:1> [SEP] right",
negative_prompt="bad quality",
)
assert [call["text"] for call in calls["encode"]] == [
"global",
"left ",
"right",
"global",
"bad quality",
"bad quality",
"bad quality",
"bad quality",
]
assert calls["encode"][1]["clip"] == calls["encode"][5]["clip"]
assert calls["encode"][3]["clip"] == calls["encode"][7]["clip"]
assert calls["encode"][0]["clip"] == calls["encode"][4]["clip"]
assert calls["encode"][2]["clip"] == calls["encode"][6]["clip"]
def test_schedule_encode_graph_keeps_loras_local_to_aligned_segments(
monkeypatch: pytest.MonkeyPatch,
) -> None:
@@ -105,23 +143,44 @@ def test_schedule_encode_graph_keeps_loras_local_to_aligned_segments(
assert output.args[0] == ["model", 0]
assert calls["lora"] == []
assert output.expand is not None
hook_nodes = [
parsed_hook_nodes = [
node
for node in output.expand.values()
if node["class_type"] == "PCLoraHooksFromText"
]
assert [node["inputs"]["text"] for node in hook_nodes] == [
assert [node["inputs"]["text"] for node in parsed_hook_nodes] == [
"<lora:a:1>\n<lora:c:1>",
"<lora:b:1>\n<lora:d:[0:1:0.5]>",
]
regional_hook_nodes = [
node
for node in output.expand.values()
if node["class_type"] == "SimpleSyrup.PrepareRegionalLoraHooks"
]
assert len(regional_hook_nodes) == 2
clip_nodes = [
node for node in output.expand.values() if node["class_type"] == "SetClipHooks"
]
assert len(clip_nodes) == 2
assert all(node["inputs"]["apply_to_conds"] is True for node in clip_nodes)
assert all(node["inputs"]["schedule_clip"] is True for node in clip_nodes)
assert calls["encode"][0]["clip"] == calls["encode"][2]["clip"]
assert calls["encode"][1]["clip"] == calls["encode"][3]["clip"]
assert clip_nodes == []
attachment_nodes = [
node
for node in output.expand.values()
if node["class_type"] == "ConditioningSetProperties"
]
assert len(attachment_nodes) == 6
assert all(node["inputs"]["strength"] == 1.0 for node in attachment_nodes)
assert all(
node["inputs"]["set_cond_area"] == "default" for node in attachment_nodes
)
companion_nodes = [
node
for node in output.expand.values()
if node["class_type"] == "SimpleSyrup.AttachRegionalGlobalConditioning"
]
assert len(companion_nodes) == 2
assert calls["encode"][0]["clip"] == calls["encode"][3]["clip"]
assert calls["encode"][1]["clip"] == calls["encode"][4]["clip"]
assert calls["encode"][2]["clip"] == calls["encode"][5]["clip"]
assert calls["encode"][0]["clip"] != calls["encode"][1]["clip"]
@@ -176,7 +235,11 @@ def _install_fake_prompt_control(
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": []}
calls: dict[str, list[dict[str, Any]]] = {
"lora": [],
"hook": [],
"encode": [],
}
class FakePCLazyLoraLoaderAdvanced:
"""Prompt-Control LoRA scheduler test double."""
@@ -199,6 +262,29 @@ def _install_fake_prompt_control(
assert start == 0.0
assert end == 1.0
assert num_steps == 0
if model is None:
calls["hook"].append({"clip": clip, "text": text})
graph_utils = import_module("comfy_execution.graph_utils")
graph = graph_utils.GraphBuilder()
hooks = graph.node(
"CreateHookLora",
lora_name=text,
strength_model=1.0,
strength_clip=1.0,
)
hooked_clip = graph.node(
"SetClipHooks",
clip=clip,
hooks=hooks.out(0),
apply_to_conds=True,
schedule_clip=True,
)
return io.NodeOutput(
None,
hooked_clip.out(0),
hooks.out(0),
expand=graph.finalize(),
)
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}"
@@ -38,8 +38,8 @@ def test_planner_combines_lora_tags_only_at_aligned_indexes() -> None:
]
def test_planner_preserves_explicit_empties_without_padding_shorter_side() -> None:
"""Empty chunks remain positional while unequal side lengths remain unequal."""
def test_planner_preserves_empties_and_fills_missing_side_from_global() -> None:
"""Authored empties remain while missing negative positions reuse global text."""
plan = PromptControlSegmentPlanningService().prepare(
positive_prompt="global [SEP] [SEP] right <lora:right:1>",
@@ -52,10 +52,62 @@ def test_planner_preserves_explicit_empties_without_padding_shorter_side() -> No
"",
"right ",
]
assert [chunk.text for chunk in plan.negative.chunks] == ["negative"]
assert [chunk.text for chunk in plan.negative.chunks] == [
"negative",
"negative",
"negative",
]
assert [hook.lora_tags for hook in plan.hooks] == ["", "", "<lora:right:1>"]
def test_planner_fallback_text_does_not_repeat_global_lora_tags() -> None:
"""Synthetic segments inherit global text but not global SEP-local tags."""
plan = PromptControlSegmentPlanningService().prepare(
positive_prompt="global [SEP] left <lora:left:1> [SEP] right",
negative_prompt="bad <lora:global-negative:1>",
separator="[SEP]",
)
assert [chunk.text for chunk in plan.negative.chunks] == ["bad ", "bad ", "bad "]
assert [chunk.lora_tags for chunk in plan.negative.chunks] == [
"<lora:global-negative:1>",
"",
"",
]
assert [hook.lora_tags for hook in plan.hooks] == [
"<lora:global-negative:1>",
"<lora:left:1>",
"",
]
def test_planner_fills_missing_positive_side_symmetrically() -> None:
"""Negative-authored regions reuse global positive text with local hooks."""
plan = PromptControlSegmentPlanningService().prepare(
positive_prompt="subject <lora:global-positive:1>",
negative_prompt="bad [SEP] hands <lora:hands:1> [SEP] text",
separator="[SEP]",
)
assert [chunk.text for chunk in plan.positive.chunks] == [
"subject ",
"subject ",
"subject ",
]
assert [chunk.lora_tags for chunk in plan.positive.chunks] == [
"<lora:global-positive:1>",
"",
"",
]
assert [hook.lora_tags for hook in plan.hooks] == [
"<lora:global-positive:1>",
"<lora:hands:1>",
"",
]
def test_planner_marks_single_segment_prompts_as_unbatched() -> None:
"""No-SEP prompts retain the compatibility path used by the scheduler."""
+76
View File
@@ -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
"""Tests for authored and fallback SEP segment alignment."""
from __future__ import annotations
import pytest
from simple_syrup.domain.prompt_segment_alignment import (
build_prompt_segment_alignment,
)
def test_alignment_reuses_global_value_for_missing_negative_segments() -> None:
"""Missing negative positions resolve to the authored global entry."""
plan = build_prompt_segment_alignment(positive_count=3, negative_count=1)
assert plan.segment_count == 3
assert plan.positive.materialize(("global", "left", "right")) == (
"global",
"left",
"right",
)
assert plan.negative.materialize(("bad",)) == ("bad", "bad", "bad")
assert [source.authored for source in plan.negative.sources] == [True, False, False]
assert [source.source_index for source in plan.negative.sources] == [0, 0, 0]
def test_alignment_is_symmetric_for_missing_positive_segments() -> None:
"""Missing positive positions use global positive rather than a prior region."""
plan = build_prompt_segment_alignment(positive_count=1, negative_count=3)
assert plan.positive.materialize(("global",)) == (
"global",
"global",
"global",
)
assert plan.negative.materialize(("bad", "hands", "text")) == (
"bad",
"hands",
"text",
)
def test_alignment_preserves_authored_empty_segments() -> None:
"""An authored empty position remains empty instead of falling back."""
plan = build_prompt_segment_alignment(positive_count=3, negative_count=1)
assert plan.positive.materialize(("global", "", "right")) == (
"global",
"",
"right",
)
assert plan.positive.sources[1].authored is True
@pytest.mark.parametrize(
("positive_count", "negative_count"),
[(0, 1), (1, 0), (-1, 1), (1, -1)],
)
def test_alignment_rejects_sides_without_a_global_segment(
positive_count: int,
negative_count: int,
) -> None:
"""Each side must expose index 0 before fallback can be planned."""
with pytest.raises(ValueError, match="at least one authored segment"):
build_prompt_segment_alignment(
positive_count=positive_count,
negative_count=negative_count,
)
@@ -0,0 +1,51 @@
# 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 internal regional global-conditioning companions."""
from __future__ import annotations
import pytest
from simple_syrup.runtime.regional_conditioning_companion import (
attach_global_companion,
detach_global_companion,
)
def test_attach_and_detach_preserve_conditioning_without_mutation() -> None:
"""The internal companion round trip copies containers and metadata."""
local = [["local tensor", {"hooks": "regional", "local": True}]]
global_conditioning = [["global tensor", {"hooks": "regional", "global": True}]]
attached = attach_global_companion(local, global_conditioning)
detached, companion = detach_global_companion(attached)
assert detached == local
assert companion == global_conditioning
assert attached is not local
assert detached is not attached
assert companion is not global_conditioning
assert "simple_syrup.regional_global_companion" not in local[0][1]
def test_detach_without_companion_returns_a_clean_copy() -> None:
"""Ordinary conditioning remains ordinary after companion inspection."""
conditioning = [["tensor", {"source": "ordinary"}]]
detached, companion = detach_global_companion(conditioning)
assert detached == conditioning
assert detached is not conditioning
assert companion is None
@pytest.mark.parametrize("value", [None, [], "conditioning"])
def test_attach_rejects_invalid_conditioning(value: object) -> None:
"""The internal node rejects empty or non-conditioning graph values."""
with pytest.raises(TypeError, match="CONDITIONING"):
attach_global_companion(value, [["global", {}]])
+73 -12
View File
@@ -10,6 +10,9 @@ import pytest
import torch
from simple_syrup.domain.conditioning_batch import ConditioningBatch
from simple_syrup.runtime.regional_conditioning_companion import (
attach_global_companion,
)
from simple_syrup.services.regional_conditioning_service import (
RegionalConditioningService,
)
@@ -75,16 +78,10 @@ def test_batches_pair_global_and_regions_independently() -> None:
"negative global",
"negative region 0",
]
assert torch.equal(
assembled_positive[0][1]["mask"],
torch.full((1, 3, 3), 0.25),
)
assert torch.equal(
assembled_negative[0][1]["mask"],
torch.ones((1, 3, 3)),
)
assert assembled_positive[0][1]["mask_strength"] == 1.0
assert assembled_negative[0][1]["mask_strength"] == 1.0
assert assembled_positive[0][1]["default"] is True
assert assembled_negative[0][1]["default"] is True
assert "mask" not in assembled_positive[0][1]
assert "mask" not in assembled_negative[0][1]
assert torch.equal(assembled_positive[1][1]["mask"], masks[0:1])
assert torch.equal(assembled_positive[2][1]["mask"], masks[1:2])
assert torch.equal(assembled_negative[1][1]["mask"], masks[0:1])
@@ -169,6 +166,42 @@ def test_mask_composition_preserves_segment_lora_hook_metadata() -> None:
assert assembled[1][1]["other"] == "region metadata"
@pytest.mark.parametrize(
("regional_prompt_weight", "expected_sources", "expected_strengths"),
[
(0.0, ["global", "hooked global"], [1.0]),
(0.5, ["global", "hooked global", "region"], [0.5, 0.5]),
(1.0, ["global", "region"], [1.0]),
],
)
def test_hooked_global_companion_keeps_lora_full_while_prompt_weight_changes(
regional_prompt_weight: float,
expected_sources: list[str],
expected_strengths: list[float],
) -> None:
"""Global and local prompt shares use one regional LoRA model state."""
regional = attach_global_companion(
[["region", {"hooks": "regional hooks"}]],
[["hooked global", {"hooks": "regional hooks"}]],
)
assembled, _ = RegionalConditioningService().assemble(
positive=ConditioningBatch((_conditioning("global"), regional)),
negative=_conditioning("negative"),
masks=torch.ones((1, 2, 2)),
regional_prompt_weight=regional_prompt_weight,
region_mask_feather=0,
)
assert [item[0] for item in assembled] == expected_sources
assert [item[1]["mask_strength"] for item in assembled[1:]] == expected_strengths
assert all(item[1]["hooks"] == "regional hooks" for item in assembled[1:])
assert all(
"simple_syrup.regional_global_companion" not in item[1] for item in assembled
)
def test_zero_regional_prompt_weight_returns_only_unchanged_global_entries() -> None:
"""The zero endpoint disables regional conditioning completely."""
@@ -200,12 +233,40 @@ def test_full_regional_prompt_weight_complements_global_inside_mask() -> None:
region_mask_feather=0,
)
assert torch.equal(positive[0][1]["mask"], 1.0 - mask)
assert positive[0][1]["mask_strength"] == 1.0
assert positive[0][1]["default"] is True
assert "mask" not in positive[0][1]
assert torch.equal(positive[1][1]["mask"], mask)
assert positive[1][1]["mask_strength"] == 1.0
@pytest.mark.parametrize("regional_prompt_weight", [0.25, 0.5, 1.0])
def test_matched_global_negative_retains_full_regional_influence(
regional_prompt_weight: float,
) -> None:
"""Global and fallback negative shares sum to full strength in a region."""
negative = ConditioningBatch(
(_conditioning("global negative"), _conditioning("global negative"))
)
_, assembled_negative = RegionalConditioningService().assemble(
positive=ConditioningBatch(
(_conditioning("global positive"), _conditioning("regional positive"))
),
negative=negative,
masks=torch.ones((1, 2, 2)),
regional_prompt_weight=regional_prompt_weight,
region_mask_feather=0,
)
assert assembled_negative[0][1]["default"] is True
assert assembled_negative[1][1]["mask_strength"] == regional_prompt_weight
assert torch.equal(
assembled_negative[1][1]["mask"],
torch.ones((1, 2, 2)),
)
@pytest.mark.parametrize("weight", [-0.01, 1.01, float("nan")])
def test_invalid_regional_prompt_weight_fails_before_mask_processing(
weight: float,
+149
View File
@@ -0,0 +1,149 @@
# 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 regional LoRA CLIP preparation."""
from __future__ import annotations
from types import SimpleNamespace
from typing import Any, cast
import comfy.hooks
import comfy.lora
import pytest
from simple_syrup.runtime.regional_lora_hooks import prepare_regional_lora_clip
class _FakeClip:
"""Record whether regional preparation clones the text encoder."""
def __init__(self) -> None:
"""Create a clip and its minimal patcher collaboration."""
self.cond_stage_model = object()
self.registrations: list[tuple[Any, Any]] = []
self.patcher = SimpleNamespace(
forced_hooks=None,
register_all_hook_patches=self._register_hooks,
)
self.use_clip_schedule = False
self.clone_calls: list[bool] = []
def clone(self, disable_dynamic: bool = False) -> _FakeClip:
"""Return a distinct clip while recording clone policy."""
self.clone_calls.append(disable_dynamic)
clone = _FakeClip()
clone.clone_calls = self.clone_calls
clone.registrations = self.registrations
clone.patcher.register_all_hook_patches = clone._register_hooks
return clone
def _register_hooks(self, hooks: Any, target: Any) -> None:
"""Record native hook registration on the prepared clip."""
self.registrations.append((hooks, target))
def test_model_only_hooks_preserve_original_clip(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A diffusion-only LoRA never clones or patches the text encoder."""
monkeypatch.setattr(comfy.lora, "model_lora_keys_clip", lambda model, keys: keys)
monkeypatch.setattr(
comfy.lora,
"load_lora",
lambda weights, key_map, log_missing: {},
)
hooks = comfy.hooks.create_hook_lora(
{"diffusion_model.layer.lora_A.weight": object()},
strength_model=1.0,
strength_clip=1.0,
)
clip = _FakeClip()
prepared_clip, prepared_hooks = prepare_regional_lora_clip(clip, hooks)
assert prepared_clip is clip
assert prepared_hooks is hooks
assert clip.clone_calls == []
def test_clip_hooks_prepare_a_scheduled_clip(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A LoRA with matching text-encoder weights retains native CLIP hooks."""
monkeypatch.setattr(
comfy.lora,
"model_lora_keys_clip",
lambda model, keys: {"clip_key": "source_key"},
)
monkeypatch.setattr(
comfy.lora,
"load_lora",
lambda weights, key_map, log_missing: {"clip.weight": object()},
)
hooks = cast(
Any,
comfy.hooks.create_hook_lora(
{"clip.source.lora_A.weight": object()},
strength_model=1.0,
strength_clip=0.75,
),
)
clip = _FakeClip()
prepared_clip, prepared_hooks = prepare_regional_lora_clip(clip, hooks)
assert prepared_clip is not clip
assert prepared_hooks is hooks
assert clip.clone_calls == [True]
assert prepared_clip.use_clip_schedule is True
assert prepared_clip.patcher.forced_hooks is not hooks
forced_hooks = cast(Any, prepared_clip.patcher.forced_hooks)
source_hooks = cast(Any, hooks)
forced_hook = cast(Any, forced_hooks.hooks[0])
source_hook = cast(Any, source_hooks.hooks[0])
assert forced_hook.hook_ref is source_hook.hook_ref
assert len(clip.registrations) == 1
assert clip.registrations[0][0] is hooks
def test_zero_strength_clip_hooks_preserve_original_clip(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Disabled CLIP strength avoids matching and cloning the text encoder."""
monkeypatch.setattr(
comfy.lora,
"model_lora_keys_clip",
lambda model, keys: {"clip_key": "source_key"},
)
def fail_load_lora(*args: object, **kwargs: object) -> object:
"""Fail if disabled CLIP hooks attempt to resolve weights."""
raise AssertionError("disabled CLIP hook should not load weights")
monkeypatch.setattr(comfy.lora, "load_lora", fail_load_lora)
hooks = comfy.hooks.create_hook_lora(
{"clip.source.lora_A.weight": object()},
strength_model=1.0,
strength_clip=0.0,
)
clip = _FakeClip()
prepared_clip, _ = prepare_regional_lora_clip(clip, hooks)
assert prepared_clip is clip
assert clip.clone_calls == []
def test_prepare_regional_lora_clip_rejects_non_hook_values() -> None:
"""Invalid graph values fail before CLIP or model registration."""
with pytest.raises(TypeError, match="Comfy HOOKS"):
prepare_regional_lora_clip(_FakeClip(), object())
+18 -22
View File
@@ -10,40 +10,36 @@ import pytest
import torch
from simple_syrup.masking.regional_prompt_masks import (
complementary_global_prompt_mask,
prepare_regional_mask_batch,
regional_mask,
)
def test_complementary_global_mask_uses_clamped_accumulated_coverage() -> None:
"""Overlapping regions accumulate while global coverage remains normalized."""
def test_prepare_regional_mask_batch_clamps_without_mutating_input() -> None:
"""Mask preparation normalizes authored values on a separate tensor."""
masks = torch.tensor(
[
[[1.0, 1.0, 0.0]],
[[0.0, 1.0, 1.0]],
]
)
masks = torch.tensor([[[-1.0, 0.5, 2.0]]])
original = masks.clone()
global_mask = complementary_global_prompt_mask(masks, (0, 1), 0.5)
regional_total = masks.sum(dim=0, keepdim=True) * 0.5
global_share = global_mask / (global_mask + regional_total)
prepared = prepare_regional_mask_batch(masks, feather=0)
assert torch.equal(global_mask, torch.tensor([[[0.5, 0.5, 0.5]]]))
assert torch.allclose(global_share, torch.tensor([[[0.5, 1.0 / 3.0, 0.5]]]))
assert torch.equal(prepared, torch.tensor([[[0.0, 0.5, 1.0]]]))
assert torch.equal(masks, original)
def test_complementary_global_mask_uses_only_paired_indices() -> None:
"""Extra authored masks do not reduce global prompt influence."""
def test_regional_mask_returns_one_ordered_batch_entry() -> None:
"""Positional selection preserves authored mask order and BHW shape."""
masks = torch.stack([torch.zeros((2, 2)), torch.ones((2, 2))])
global_mask = complementary_global_prompt_mask(masks, (0,), 1.0)
selected = regional_mask(masks, 1)
assert torch.equal(global_mask, torch.ones((1, 2, 2)))
assert selected.shape == (1, 2, 2)
assert torch.equal(selected, masks[1:2])
def test_complementary_global_mask_requires_a_regional_pair() -> None:
"""Coverage cannot be calculated without a paired regional prompt."""
def test_regional_mask_rejects_out_of_range_index() -> None:
"""Selection fails explicitly when prompt and mask planning diverge."""
with pytest.raises(ValueError, match="at least one regional mask"):
complementary_global_prompt_mask(torch.ones((1, 2, 2)), (), 0.5)
with pytest.raises(IndexError, match="out of range"):
regional_mask(torch.ones((1, 2, 2)), 1)
+5
View File
@@ -133,10 +133,15 @@ def test_sep_prompts_and_disk_masks_flow_into_regional_sampler(
assert [item[0] for item in assembled_negative] == [
"global negative",
"left negative",
"global negative",
]
assert torch.equal(assembled_positive[1][1]["mask"], masks[0:1])
assert torch.equal(assembled_positive[2][1]["mask"], masks[1:2])
assert assembled_positive[1][1]["mask_strength"] == 0.8
assert assembled_positive[2][1]["mask_strength"] == 0.8
assert assembled_negative[1][1]["mask_strength"] == 0.8
assert torch.equal(assembled_negative[2][1]["mask"], masks[1:2])
assert assembled_negative[2][1]["mask_strength"] == 0.8
assert assembled_negative[0][1]["default"] is True
assert "mask" not in assembled_negative[0][1]
assert call["latent_image"]["samples"].shape[0] == 2
+2
View File
@@ -60,7 +60,9 @@ BASE_NODE_IDS = [
]
PROMPT_CONTROL_NODE_IDS = [
"SimpleSyrup.AttachRegionalGlobalConditioning",
"SimpleSyrup.EncodePromptBatchWithPromptControl",
"SimpleSyrup.PrepareRegionalLoraHooks",
"SimpleSyrup.ScheduleAndEncodePromptsWithPromptControl",
]
@@ -140,6 +140,7 @@ def test_schedule_and_encode_prompt_control_node_schema() -> None:
assert schema.display_name == "Schedule & Encode Prompts"
assert schema.enable_expand is True
assert schema.category == "SimpleSyrup/Conditioning"
assert "global" in schema.description.lower()
assert [output.io_type for output in schema.outputs] == [
"MODEL",
"CONDITIONING,CONDITIONING_BATCH",