fix(regional): align prompt batches and LoRA hooks
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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."
|
||||
),
|
||||
},
|
||||
),
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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", {}]])
|
||||
@@ -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,
|
||||
|
||||
@@ -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())
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user