diff --git a/simple_syrup/domain/prompt_segment_alignment.py b/simple_syrup/domain/prompt_segment_alignment.py new file mode 100644 index 0000000..ba3f105 --- /dev/null +++ b/simple_syrup/domain/prompt_segment_alignment.py @@ -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) + ), + ) diff --git a/simple_syrup/masking/regional_prompt_masks.py b/simple_syrup/masking/regional_prompt_masks.py index aee70eb..5215d70 100644 --- a/simple_syrup/masking/regional_prompt_masks.py +++ b/simple_syrup/masking/regional_prompt_masks.py @@ -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 diff --git a/simple_syrup/nodes/encode_prompt_batch.py b/simple_syrup/nodes/encode_prompt_batch.py index 4702e3c..08654e7 100644 --- a/simple_syrup/nodes/encode_prompt_batch.py +++ b/simple_syrup/nodes/encode_prompt_batch.py @@ -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, ) diff --git a/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py b/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py index c9c18f4..1d9266e 100644 --- a/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py +++ b/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py @@ -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." ), }, ), diff --git a/simple_syrup/nodes_v3/__init__.py b/simple_syrup/nodes_v3/__init__.py index aa2de4a..ad3da19 100644 --- a/simple_syrup/nodes_v3/__init__.py +++ b/simple_syrup/nodes_v3/__init__.py @@ -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, ] diff --git a/simple_syrup/nodes_v3/attach_regional_global_conditioning.py b/simple_syrup/nodes_v3/attach_regional_global_conditioning.py new file mode 100644 index 0000000..35a10e1 --- /dev/null +++ b/simple_syrup/nodes_v3/attach_regional_global_conditioning.py @@ -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),) diff --git a/simple_syrup/nodes_v3/encode_prompt_batch_with_prompt_control.py b/simple_syrup/nodes_v3/encode_prompt_batch_with_prompt_control.py index b009edf..6b4eb85 100644 --- a/simple_syrup/nodes_v3/encode_prompt_batch_with_prompt_control.py +++ b/simple_syrup/nodes_v3/encode_prompt_batch_with_prompt_control.py @@ -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( diff --git a/simple_syrup/nodes_v3/prepare_regional_lora_hooks.py b/simple_syrup/nodes_v3/prepare_regional_lora_hooks.py new file mode 100644 index 0000000..c534399 --- /dev/null +++ b/simple_syrup/nodes_v3/prepare_regional_lora_hooks.py @@ -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) diff --git a/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py b/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py index 994c2b0..ca4d234 100644 --- a/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py +++ b/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py @@ -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." ), ), ], diff --git a/simple_syrup/runtime/prompt_control_batch_graph.py b/simple_syrup/runtime/prompt_control_batch_graph.py index 43f806a..67fd9b7 100644 --- a/simple_syrup/runtime/prompt_control_batch_graph.py +++ b/simple_syrup/runtime/prompt_control_batch_graph.py @@ -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, diff --git a/simple_syrup/runtime/prompt_control_graph_adapter.py b/simple_syrup/runtime/prompt_control_graph_adapter.py index fca4347..fe4ce13 100644 --- a/simple_syrup/runtime/prompt_control_graph_adapter.py +++ b/simple_syrup/runtime/prompt_control_graph_adapter.py @@ -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]], diff --git a/simple_syrup/runtime/prompt_control_schedule_encode_graph.py b/simple_syrup/runtime/prompt_control_schedule_encode_graph.py index b2d7ced..03aa60b 100644 --- a/simple_syrup/runtime/prompt_control_schedule_encode_graph.py +++ b/simple_syrup/runtime/prompt_control_schedule_encode_graph.py @@ -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, diff --git a/simple_syrup/runtime/regional_conditioning_companion.py b/simple_syrup/runtime/regional_conditioning_companion.py new file mode 100644 index 0000000..24c9791 --- /dev/null +++ b/simple_syrup/runtime/regional_conditioning_companion.py @@ -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] diff --git a/simple_syrup/runtime/regional_lora_hooks.py b/simple_syrup/runtime/regional_lora_hooks.py new file mode 100644 index 0000000..4929f4b --- /dev/null +++ b/simple_syrup/runtime/regional_lora_hooks.py @@ -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 diff --git a/simple_syrup/services/prompt_batch_encoding_service.py b/simple_syrup/services/prompt_batch_encoding_service.py new file mode 100644 index 0000000..59ccb50 --- /dev/null +++ b/simple_syrup/services/prompt_batch_encoding_service.py @@ -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), + ), + ) diff --git a/simple_syrup/services/prompt_control_segment_planning_service.py b/simple_syrup/services/prompt_control_segment_planning_service.py index ae325c3..194ae2c 100644 --- a/simple_syrup/services/prompt_control_segment_planning_service.py +++ b/simple_syrup/services/prompt_control_segment_planning_service.py @@ -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) diff --git a/simple_syrup/services/regional_conditioning_service.py b/simple_syrup/services/regional_conditioning_service.py index 6bdf33c..df5f757 100644 --- a/simple_syrup/services/regional_conditioning_service.py +++ b/simple_syrup/services/regional_conditioning_service.py @@ -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, diff --git a/tests/test_encode_prompt_batch_node.py b/tests/test_encode_prompt_batch_node.py index 1974040..fc1e462 100644 --- a/tests/test_encode_prompt_batch_node.py +++ b/tests/test_encode_prompt_batch_node.py @@ -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.""" diff --git a/tests/test_encode_prompt_batch_with_prompt_control_node.py b/tests/test_encode_prompt_batch_with_prompt_control_node.py index 3314558..8bb31d5 100644 --- a/tests/test_encode_prompt_batch_with_prompt_control_node.py +++ b/tests/test_encode_prompt_batch_with_prompt_control_node.py @@ -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() diff --git a/tests/test_prompt_control_batch_graph.py b/tests/test_prompt_control_batch_graph.py index f9ace8d..d901adb 100644 --- a/tests/test_prompt_control_batch_graph.py +++ b/tests/test_prompt_control_batch_graph.py @@ -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] == [ "\n", "", ] - 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 [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) diff --git a/tests/test_prompt_control_schedule_encode_graph.py b/tests/test_prompt_control_schedule_encode_graph.py index 883e630..e4659bf 100644 --- a/tests/test_prompt_control_schedule_encode_graph.py +++ b/tests/test_prompt_control_schedule_encode_graph.py @@ -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 [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] == [ "\n", "\n", ] + 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}" diff --git a/tests/test_prompt_control_segment_planning_service.py b/tests/test_prompt_control_segment_planning_service.py index 69903fd..4b62549 100644 --- a/tests/test_prompt_control_segment_planning_service.py +++ b/tests/test_prompt_control_segment_planning_service.py @@ -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 ", @@ -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] == ["", "", ""] +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 [SEP] right", + negative_prompt="bad ", + separator="[SEP]", + ) + + assert [chunk.text for chunk in plan.negative.chunks] == ["bad ", "bad ", "bad "] + assert [chunk.lora_tags for chunk in plan.negative.chunks] == [ + "", + "", + "", + ] + assert [hook.lora_tags for hook in plan.hooks] == [ + "", + "", + "", + ] + + +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 ", + negative_prompt="bad [SEP] hands [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] == [ + "", + "", + "", + ] + assert [hook.lora_tags for hook in plan.hooks] == [ + "", + "", + "", + ] + + def test_planner_marks_single_segment_prompts_as_unbatched() -> None: """No-SEP prompts retain the compatibility path used by the scheduler.""" diff --git a/tests/test_prompt_segment_alignment.py b/tests/test_prompt_segment_alignment.py new file mode 100644 index 0000000..f35d7a4 --- /dev/null +++ b/tests/test_prompt_segment_alignment.py @@ -0,0 +1,76 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""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, + ) diff --git a/tests/test_regional_conditioning_companion.py b/tests/test_regional_conditioning_companion.py new file mode 100644 index 0000000..b538d9b --- /dev/null +++ b/tests/test_regional_conditioning_companion.py @@ -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", {}]]) diff --git a/tests/test_regional_conditioning_service.py b/tests/test_regional_conditioning_service.py index 08b8295..13e36ec 100644 --- a/tests/test_regional_conditioning_service.py +++ b/tests/test_regional_conditioning_service.py @@ -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, diff --git a/tests/test_regional_lora_hooks.py b/tests/test_regional_lora_hooks.py new file mode 100644 index 0000000..4b64a10 --- /dev/null +++ b/tests/test_regional_lora_hooks.py @@ -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()) diff --git a/tests/test_regional_prompt_masks.py b/tests/test_regional_prompt_masks.py index 31b1b99..71471e6 100644 --- a/tests/test_regional_prompt_masks.py +++ b/tests/test_regional_prompt_masks.py @@ -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) diff --git a/tests/test_regional_prompt_workflow.py b/tests/test_regional_prompt_workflow.py index 0875f3d..97dcbb9 100644 --- a/tests/test_regional_prompt_workflow.py +++ b/tests/test_regional_prompt_workflow.py @@ -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 diff --git a/tests/test_registration.py b/tests/test_registration.py index 5c2f4c6..bbd36c8 100644 --- a/tests/test_registration.py +++ b/tests/test_registration.py @@ -60,7 +60,9 @@ BASE_NODE_IDS = [ ] PROMPT_CONTROL_NODE_IDS = [ + "SimpleSyrup.AttachRegionalGlobalConditioning", "SimpleSyrup.EncodePromptBatchWithPromptControl", + "SimpleSyrup.PrepareRegionalLoraHooks", "SimpleSyrup.ScheduleAndEncodePromptsWithPromptControl", ] diff --git a/tests/test_schedule_and_encode_prompts_with_prompt_control_node.py b/tests/test_schedule_and_encode_prompts_with_prompt_control_node.py index 0c54e30..3f499bd 100644 --- a/tests/test_schedule_and_encode_prompts_with_prompt_control_node.py +++ b/tests/test_schedule_and_encode_prompts_with_prompt_control_node.py @@ -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",