diff --git a/pyproject.toml b/pyproject.toml index 48301fc..ea19593 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -65,3 +65,8 @@ ignore_missing_imports = true [tool.pytest.ini_options] pythonpath = [".", "../.."] testpaths = ["tests"] +filterwarnings = [ + "error", + "ignore:builtin type SwigPyPacked has no __module__ attribute:DeprecationWarning", + "ignore:builtin type SwigPyObject has no __module__ attribute:DeprecationWarning", +] diff --git a/simple_syrup/domain/conditioning_batch.py b/simple_syrup/domain/conditioning_batch.py index ceb6be4..1b69ba1 100644 --- a/simple_syrup/domain/conditioning_batch.py +++ b/simple_syrup/domain/conditioning_batch.py @@ -40,6 +40,23 @@ class ConditioningBatch: return ConditioningBatch((*self.entries, conditioning)) +def batch_conditioning( + values: tuple[Conditioning | ConditioningBatch, ...], +) -> ConditioningBatch: + """Flatten conditioning values and batches into one ordered batch.""" + + if not values: + raise ValueError("Batch Region Conditioning requires one or more inputs.") + + entries: list[Conditioning] = [] + for value in values: + if isinstance(value, ConditioningBatch): + entries.extend(value.entries) + else: + entries.append(value) + return ConditioningBatch(tuple(entries)) + + def split_prompt_batch(text: str, separator: str = "[SEP]") -> tuple[str, ...]: """Split prompt text into ordered chunks using a configurable separator.""" diff --git a/simple_syrup/domain/segs.py b/simple_syrup/domain/segs.py index 96daa6c..b68158f 100644 --- a/simple_syrup/domain/segs.py +++ b/simple_syrup/domain/segs.py @@ -184,6 +184,32 @@ def to_impact_compatible_segs_group(segs_group: NativeSegsGroup) -> list[ImpactS return [to_impact_compatible_segs(segs) for segs in segs_group] +def batch_segs(values: Iterable[object]) -> NativeSegs: + """Return one SEGS payload containing all segments in input order.""" + + raw_values = tuple(values) + if not raw_values: + raise ValueError("Batch SEGS requires one or more SEGS inputs.") + + expected_header: SegsHeader | None = None + batched_segments: list[Segment] = [] + for index, value in enumerate(raw_values, start=1): + header, segments = coerce_segs(value) + if expected_header is None: + expected_header = header + elif header != expected_header: + raise ValueError( + "Batch SEGS requires all SEGS inputs to use the same image size; " + f"input {index} is {_format_header(header)} but input 1 is " + f"{_format_header(expected_header)}." + ) + batched_segments.extend(segments) + + if expected_header is None: + raise ValueError("Batch SEGS requires one or more SEGS inputs.") + return expected_header, tuple(batched_segments) + + def limit_segs(segs: NativeSegs, keep_only: int, keep_by: str) -> NativeSegs: """Return SEGS limited by a user-facing ranking policy.""" @@ -284,6 +310,13 @@ def _coerce_header(value: object) -> SegsHeader: return height, width +def _format_header(header: SegsHeader) -> str: + """Return a height-first image size description.""" + + height, width = header + return f"{height}x{width}" + + def _looks_like_segs(value: object) -> bool: """Return whether a value has the outer shape of one SEGS payload.""" diff --git a/simple_syrup/nodes/__init__.py b/simple_syrup/nodes/__init__.py index 7934259..c06329c 100644 --- a/simple_syrup/nodes/__init__.py +++ b/simple_syrup/nodes/__init__.py @@ -7,6 +7,8 @@ from __future__ import annotations from ..runtime.prompt_control_availability import prompt_control_is_available +from .batch_region_conditioning import BatchRegionConditioning +from .batch_segs import BatchSEGS from .conditioning_batch_pack import ConditioningBatchAppend, ConditioningBatchStart from .detail_segs_as_regions import DetailSEGSAsRegions from .detail_segs_by_scale_factor import DetailSEGSByScaleFactor @@ -32,6 +34,7 @@ from .scale_factor import ScaleFactor from .seed import Seed from .simple_load_anima import SimpleLoadAnima from .simple_load_checkpoint import SimpleLoadCheckpoint +from .tag_segs_with_wd14 import TagSEGSWithWD14 from .tile_and_tag_segs import TileAndTagSEGS from .vae_options import VAEDecodeOptions, VAEEncodeOptions from .vitmatte_model_loader import ViTMatteModelLoader @@ -40,6 +43,8 @@ from .wd14_tagger_loader import WD14TaggerLoader _PROMPT_CONTROL_EXPORTS: list[str] = [] NODE_CLASS_MAPPINGS = { + "SimpleSyrup.BatchRegionConditioning": BatchRegionConditioning, + "SimpleSyrup.BatchSEGS": BatchSEGS, "SimpleSyrup.ConditioningBatchAppend": ConditioningBatchAppend, "SimpleSyrup.ConditioningBatchStart": ConditioningBatchStart, "SimpleSyrup.GroundedSAMModelInfo": GroundedSAMModelInfo, @@ -69,12 +74,15 @@ NODE_CLASS_MAPPINGS = { "SimpleSyrup.LoadUltralyticsModel": LoadUltralyticsModel, "SimpleSyrup.DetectSEGSWithUltralytics": DetectSEGSWithUltralytics, "SimpleSyrup.EncodePromptBatch": EncodePromptBatch, + "SimpleSyrup.TagSEGSWithWD14": TagSEGSWithWD14, "SimpleSyrup.TileAndTagSEGS": TileAndTagSEGS, "SimpleSyrup.ViTMatteModelLoader": ViTMatteModelLoader, "SimpleSyrup.WD14TaggerLoader": WD14TaggerLoader, } NODE_DISPLAY_NAME_MAPPINGS = { + "SimpleSyrup.BatchRegionConditioning": "Batch Region Conditioning", + "SimpleSyrup.BatchSEGS": "Batch SEGS", "SimpleSyrup.ConditioningBatchAppend": "Conditioning Batch Append", "SimpleSyrup.ConditioningBatchStart": "Conditioning Batch Start", "SimpleSyrup.GroundedSAMModelInfo": "Grounded SAM Model Info", @@ -106,6 +114,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "SimpleSyrup.LoadUltralyticsModel": "Load Ultralytics Model", "SimpleSyrup.DetectSEGSWithUltralytics": "Detect SEGS w/ Ultralytics", "SimpleSyrup.EncodePromptBatch": "Encode Prompt Batch", + "SimpleSyrup.TagSEGSWithWD14": "Tag SEGS w/ WD14", "SimpleSyrup.TileAndTagSEGS": "Tile & Tag SEGS", "SimpleSyrup.ViTMatteModelLoader": "ViTMatte Model Loader", "SimpleSyrup.WD14TaggerLoader": "Load WD14 Tagger", @@ -125,6 +134,8 @@ if prompt_control_is_available(): _PROMPT_CONTROL_EXPORTS.append("ScheduleAndEncodePromptsWithPromptControl") __all__ = [ + "BatchRegionConditioning", + "BatchSEGS", "ConditioningBatchAppend", "ConditioningBatchStart", "GroundedSAMModelInfo", @@ -152,6 +163,7 @@ __all__ = [ "SimpleLoadAnima", "SimpleLoadCheckpoint", "SimpleVAEEncode", + "TagSEGSWithWD14", "TileAndTagSEGS", "UpscaleLatentFromImage", "VAEDecodeOptions", diff --git a/simple_syrup/nodes/batch_region_conditioning.py b/simple_syrup/nodes/batch_region_conditioning.py new file mode 100644 index 0000000..3ffdd27 --- /dev/null +++ b/simple_syrup/nodes/batch_region_conditioning.py @@ -0,0 +1,54 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""ComfyUI node declaration for batching regional conditioning.""" + +from __future__ import annotations + +from typing import Any + +from ..domain.conditioning_batch import batch_conditioning +from ..nodes import tooltips + + +class BatchRegionConditioning: + """Combine conditioning and conditioning batches for regional detailing.""" + + RETURN_TYPES = ("CONDITIONING_BATCH",) + RETURN_NAMES = ("batch",) + OUTPUT_TOOLTIPS = (tooltips.BATCH_REGION_CONDITIONING_OUTPUT,) + FUNCTION = "batch" + CATEGORY = "SimpleSyrup/Conditioning" + DESCRIPTION = ( + "Combines two CONDITIONING or CONDITIONING_BATCH inputs into one ordered " + "regional conditioning batch. Chain this node to batch more sources." + ) + SEARCH_ALIASES = [ + "batch", + "conditioning batch", + "region conditioning", + "segs prompts", + ] + + @classmethod + def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]: + """Declare legacy ComfyUI inputs for regional conditioning batching.""" + + return { + "required": { + "first": ( + "CONDITIONING,CONDITIONING_BATCH", + {"tooltip": tooltips.BATCH_REGION_CONDITIONING_FIRST}, + ), + "second": ( + "CONDITIONING,CONDITIONING_BATCH", + {"tooltip": tooltips.BATCH_REGION_CONDITIONING_SECOND}, + ), + }, + } + + def batch(self, first: Any, second: Any) -> tuple[object]: + """Batch two conditioning inputs in input order.""" + + return (batch_conditioning((first, second)),) diff --git a/simple_syrup/nodes/batch_segs.py b/simple_syrup/nodes/batch_segs.py new file mode 100644 index 0000000..f0d56e6 --- /dev/null +++ b/simple_syrup/nodes/batch_segs.py @@ -0,0 +1,44 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""ComfyUI node declaration for batching SEGS.""" + +from __future__ import annotations + +from typing import Any + +from ..domain.segs import batch_segs, to_impact_compatible_segs +from ..nodes import tooltips + + +class BatchSEGS: + """Combine two SEGS payloads into one ordered SEGS payload.""" + + RETURN_TYPES = ("SEGS",) + RETURN_NAMES = ("segs",) + OUTPUT_TOOLTIPS = (tooltips.BATCH_SEGS_OUTPUT,) + FUNCTION = "batch" + CATEGORY = "SimpleSyrup/Detection" + DESCRIPTION = ( + "Combines two SEGS inputs into one ordered SEGS payload. Chain this node " + "to batch more than two SEGS sources." + ) + SEARCH_ALIASES = ["batch", "merge", "join", "combine", "segs"] + + @classmethod + def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]: + """Declare legacy ComfyUI inputs for SEGS batching.""" + + return { + "required": { + "first": ("SEGS", {"tooltip": tooltips.BATCH_SEGS_FIRST}), + "second": ("SEGS", {"tooltip": tooltips.BATCH_SEGS_SECOND}), + }, + } + + def batch(self, first: object, second: object) -> tuple[object]: + """Batch two SEGS payloads in input order.""" + + native = batch_segs((first, second)) + return (to_impact_compatible_segs(native),) diff --git a/simple_syrup/nodes/tag_segs_with_wd14.py b/simple_syrup/nodes/tag_segs_with_wd14.py new file mode 100644 index 0000000..d4e9e8a --- /dev/null +++ b/simple_syrup/nodes/tag_segs_with_wd14.py @@ -0,0 +1,172 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""ComfyUI node declaration for tagging existing SEGS with WD14.""" + +from __future__ import annotations + +from collections.abc import Callable +from typing import Any, ClassVar + +from ..nodes import tooltips +from ..runtime.wd14_tagger import WD14TagFormattingControls +from ..services.tag_segs_with_wd14_service import TagSEGSWithWD14Service +from .tile_and_tag_segs import DEFAULT_EXCLUDE_TAGS + + +class TagSEGSWithWD14: + """Create WD14 conditioning for existing SEGS.""" + + RETURN_TYPES = ("SEGS", "CONDITIONING_BATCH") + RETURN_NAMES = ("segs", "positive") + OUTPUT_TOOLTIPS = ( + tooltips.TAG_SEGS_SEGS_OUTPUT, + tooltips.TAG_SEGS_POSITIVE_OUTPUT, + ) + FUNCTION = "tag" + CATEGORY = "SimpleSyrup/Detailing" + DESCRIPTION = ( + "Tags existing SEGS crops with a connected WD14 tagger and returns " + "aligned conditioning for SEGS detailing." + ) + SEARCH_ALIASES = ["tag", "wd14", "segs", "detail", "regional"] + + service_class: ClassVar[Callable[[], TagSEGSWithWD14Service]] = ( + TagSEGSWithWD14Service + ) + + @classmethod + def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]: + """Declare ComfyUI inputs for WD14 tagging of existing SEGS.""" + + return { + "required": { + "image": ("IMAGE", {"tooltip": tooltips.TAG_SEGS_IMAGE}), + "segs": ("SEGS", {"tooltip": tooltips.TAG_SEGS_SEGS}), + "clip": ("CLIP", {"tooltip": tooltips.TAG_SEGS_CLIP}), + "wd14_tagger": ( + "WD14_TAGGER", + {"tooltip": tooltips.TAG_SEGS_WD14_TAGGER}, + ), + "universal_positive": ( + "STRING", + { + "default": "", + "multiline": False, + "tooltip": tooltips.TAG_SEGS_UNIVERSAL_POSITIVE, + }, + ), + "threshold": ( + "FLOAT", + { + "default": 0.35, + "min": 0.0, + "max": 1.0, + "step": 0.05, + "tooltip": tooltips.TILE_THRESHOLD, + }, + ), + "character_threshold": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 1.0, + "step": 0.05, + "tooltip": tooltips.TILE_CHARACTER_THRESHOLD, + }, + ), + "replace_underscore": ( + "BOOLEAN", + { + "default": True, + "tooltip": tooltips.TILE_REPLACE_UNDERSCORE, + }, + ), + "trailing_comma": ( + "BOOLEAN", + { + "default": False, + "tooltip": tooltips.TILE_TRAILING_COMMA, + }, + ), + "exclude_tags": ( + "STRING", + { + "default": DEFAULT_EXCLUDE_TAGS, + "multiline": False, + "tooltip": tooltips.TILE_EXCLUDE_TAGS, + }, + ), + }, + } + + def tag( + self, + image: object, + segs: object, + clip: Any, + wd14_tagger: object, + universal_positive: object, + threshold: object, + character_threshold: object, + replace_underscore: object, + trailing_comma: object, + exclude_tags: object, + ) -> tuple[object, object]: + """Tag existing SEGS and return aligned conditioning.""" + + tag_controls = WD14TagFormattingControls( + threshold=_float_input(threshold, "threshold"), + character_threshold=_float_input( + character_threshold, + "character_threshold", + ), + replace_underscore=_bool_input( + replace_underscore, + "replace_underscore", + ), + trailing_comma=_bool_input(trailing_comma, "trailing_comma"), + exclude_tags=_str_input(exclude_tags, "exclude_tags"), + ) + result = ( + type(self) + .service_class() + .tag( + image=image, + segs=segs, + clip=clip, + wd14_tagger=wd14_tagger, + tag_controls=tag_controls, + universal_positive=_str_input( + universal_positive, + "universal_positive", + ), + ) + ) + return result.segs, result.positive + + +def _float_input(value: object, name: str) -> float: + """Return a float node input.""" + + if isinstance(value, (int, float, str)): + return float(value) + raise TypeError(f"Tag SEGS w/ WD14 requires '{name}' to be a float.") + + +def _str_input(value: object, name: str) -> str: + """Return a string node input.""" + + if isinstance(value, str): + return value + raise TypeError(f"Tag SEGS w/ WD14 requires '{name}' to be a string.") + + +def _bool_input(value: object, name: str) -> bool: + """Return a boolean node input.""" + + if isinstance(value, bool): + return value + raise TypeError(f"Tag SEGS w/ WD14 requires '{name}' to be a boolean.") diff --git a/simple_syrup/nodes/tooltips.py b/simple_syrup/nodes/tooltips.py index da145b4..c0b857d 100644 --- a/simple_syrup/nodes/tooltips.py +++ b/simple_syrup/nodes/tooltips.py @@ -209,3 +209,31 @@ TILE_SEGS_OUTPUT = "Generated tile SEGS in the same order as the conditioning ba TILE_POSITIVE_OUTPUT = ( "Positive conditioning from WD14 tile tags, matched to SEGS order." ) + +TAG_SEGS_IMAGE = "Image that the incoming SEGS were detected from." +TAG_SEGS_SEGS = "Existing SEGS to crop, tag, and keep in their current order." +TAG_SEGS_CLIP = "CLIP model used to encode each generated SEGS prompt." +TAG_SEGS_WD14_TAGGER = "WD14 tagger that reads each SEG crop and suggests prompt tags." +TAG_SEGS_UNIVERSAL_POSITIVE = ( + "Positive prompt text added before every generated SEGS tag prompt." +) +TAG_SEGS_SEGS_OUTPUT = ( + "Original SEGS returned in the same order as the conditioning batch." +) +TAG_SEGS_POSITIVE_OUTPUT = ( + "Positive conditioning from WD14 SEGS tags, matched to SEGS order." +) + +BATCH_SEGS_FIRST = "First SEGS payload in the output order." +BATCH_SEGS_SECOND = "Second SEGS payload appended after the first." +BATCH_SEGS_OUTPUT = "Combined SEGS with all input segments in order." + +BATCH_REGION_CONDITIONING_FIRST = ( + "First conditioning or conditioning batch in the output order." +) +BATCH_REGION_CONDITIONING_SECOND = ( + "Second conditioning or conditioning batch appended after the first." +) +BATCH_REGION_CONDITIONING_OUTPUT = ( + "Conditioning batch containing all input entries in order." +) diff --git a/simple_syrup/nodes_v3/__init__.py b/simple_syrup/nodes_v3/__init__.py index efafa0f..00cdf91 100644 --- a/simple_syrup/nodes_v3/__init__.py +++ b/simple_syrup/nodes_v3/__init__.py @@ -12,8 +12,11 @@ from ..runtime.prompt_control_availability import prompt_control_is_available def get_nodes() -> list[type[object]]: """Return v3 nodes that can be advertised in this environment.""" + from .batch_region_conditioning import BatchRegionConditioningV3 + from .batch_segs import BatchSEGSV3 from .scale_factor import ScaleFactorV3 from .simple_load_checkpoint import SimpleLoadCheckpointV3 + from .tag_segs_with_wd14 import TagSEGSWithWD14V3 from .tile_and_tag_segs import TileAndTagSEGSV3 from .vae_decode_options import VAEDecodeOptionsV3 from .vae_encode_options import VAEEncodeOptionsV3 @@ -22,6 +25,9 @@ def get_nodes() -> list[type[object]]: if not prompt_control_is_available(): return [ WD14TaggerLoaderV3, + BatchSEGSV3, + BatchRegionConditioningV3, + TagSEGSWithWD14V3, TileAndTagSEGSV3, SimpleLoadCheckpointV3, ScaleFactorV3, @@ -38,6 +44,9 @@ def get_nodes() -> list[type[object]]: return [ WD14TaggerLoaderV3, + BatchSEGSV3, + BatchRegionConditioningV3, + TagSEGSWithWD14V3, TileAndTagSEGSV3, SimpleLoadCheckpointV3, ScaleFactorV3, diff --git a/simple_syrup/nodes_v3/batch_region_conditioning.py b/simple_syrup/nodes_v3/batch_region_conditioning.py new file mode 100644 index 0000000..749d6e1 --- /dev/null +++ b/simple_syrup/nodes_v3/batch_region_conditioning.py @@ -0,0 +1,87 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Comfy v3 node wrapper for batching regional conditioning.""" + +from __future__ import annotations + +from collections.abc import Mapping +from importlib import import_module +from typing import TYPE_CHECKING, Any, ClassVar + +from ..domain.conditioning_batch import batch_conditioning + +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 +ConditioningBatchIO: Any = ( + None if TYPE_CHECKING else _comfy_io.Custom("CONDITIONING_BATCH") +) + + +class BatchRegionConditioningV3(_ComfyNodeBase): + """Expose expandable mixed conditioning batching through Comfy's v3 API.""" + + @classmethod + def define_schema(cls) -> Any: + """Declare the Batch Region Conditioning v3 schema.""" + + conditioning_input = _comfy_io.MultiType.Input( + "conditioning", + [_comfy_io.Conditioning, ConditioningBatchIO], + tooltip="Conditioning or conditioning batch to append in socket order.", + ) + autogrow_template = _comfy_io.Autogrow.TemplatePrefix( + conditioning_input, + prefix="conditioning", + min=2, + max=50, + ) + return _comfy_io.Schema( + node_id="SimpleSyrup.BatchRegionConditioning", + display_name="Batch Region Conditioning", + category="SimpleSyrup/Conditioning", + description=( + "Combines CONDITIONING and CONDITIONING_BATCH inputs into one " + "ordered regional conditioning batch." + ), + search_aliases=[ + "batch", + "conditioning batch", + "region conditioning", + "segs prompts", + ], + inputs=[ + _comfy_io.Autogrow.Input( + "conditioning_inputs", + template=autogrow_template, + tooltip=( + "Expandable conditioning inputs flattened in socket order." + ), + ), + ], + outputs=[ + ConditioningBatchIO.Output( + "batch", + tooltip=( + "Conditioning batch containing all input entries in order." + ), + ), + ], + ) + + @classmethod + def execute(cls, conditioning_inputs: Mapping[str, object]) -> tuple[object]: + """Batch conditioning inputs in Autogrow order.""" + + return (batch_conditioning(tuple(conditioning_inputs.values())),) diff --git a/simple_syrup/nodes_v3/batch_segs.py b/simple_syrup/nodes_v3/batch_segs.py new file mode 100644 index 0000000..00f633d --- /dev/null +++ b/simple_syrup/nodes_v3/batch_segs.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 node wrapper for Batch SEGS.""" + +from __future__ import annotations + +from collections.abc import Mapping +from importlib import import_module +from typing import TYPE_CHECKING, Any, ClassVar + +from ..domain.segs import batch_segs, to_impact_compatible_segs + +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 BatchSEGSV3(_ComfyNodeBase): + """Expose expandable SEGS batching through Comfy's v3 API.""" + + @classmethod + def define_schema(cls) -> Any: + """Declare the Batch SEGS v3 schema.""" + + autogrow_template = _comfy_io.Autogrow.TemplatePrefix( + _comfy_io.SEGS.Input( + "segs", + tooltip="SEGS payload to append to the output batch.", + ), + prefix="segs", + min=2, + max=50, + ) + return _comfy_io.Schema( + node_id="SimpleSyrup.BatchSEGS", + display_name="Batch SEGS", + category="SimpleSyrup/Detection", + description="Combines multiple SEGS inputs into one ordered SEGS payload.", + search_aliases=["batch", "merge", "join", "combine", "segs"], + inputs=[ + _comfy_io.Autogrow.Input( + "segs_inputs", + template=autogrow_template, + tooltip="Expandable SEGS inputs joined in socket order.", + ), + ], + outputs=[ + _comfy_io.SEGS.Output( + "segs", + tooltip="Combined SEGS with all input segments in order.", + ), + ], + ) + + @classmethod + def execute(cls, segs_inputs: Mapping[str, object]) -> tuple[object]: + """Batch provided SEGS inputs in Autogrow order.""" + + native = batch_segs(segs_inputs.values()) + return (to_impact_compatible_segs(native),) diff --git a/simple_syrup/nodes_v3/tag_segs_with_wd14.py b/simple_syrup/nodes_v3/tag_segs_with_wd14.py new file mode 100644 index 0000000..c1507d8 --- /dev/null +++ b/simple_syrup/nodes_v3/tag_segs_with_wd14.py @@ -0,0 +1,136 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Comfy v3 node wrapper for Tag SEGS w/ WD14.""" + +from __future__ import annotations + +from importlib import import_module +from typing import TYPE_CHECKING, Any, ClassVar + +from ..nodes import tooltips +from ..nodes.tag_segs_with_wd14 import TagSEGSWithWD14 +from ..nodes.tile_and_tag_segs import DEFAULT_EXCLUDE_TAGS + +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 +ConditioningBatchIO: Any = ( + None if TYPE_CHECKING else _comfy_io.Custom("CONDITIONING_BATCH") +) +WD14TaggerIO: Any = None if TYPE_CHECKING else _comfy_io.Custom("WD14_TAGGER") + + +class TagSEGSWithWD14V3(_ComfyNodeBase): + """Expose WD14 tagging for existing SEGS through Comfy's v3 API.""" + + @classmethod + def define_schema(cls) -> Any: + """Declare the Tag SEGS w/ WD14 v3 schema.""" + + return _comfy_io.Schema( + node_id="SimpleSyrup.TagSEGSWithWD14", + display_name="Tag SEGS w/ WD14", + category="SimpleSyrup/Detailing", + description=( + "Tags existing SEGS crops with a connected WD14 tagger and " + "returns aligned conditioning for SEGS detailing." + ), + search_aliases=["tag", "wd14", "segs", "detail", "regional"], + inputs=[ + _comfy_io.Image.Input("image", tooltip=tooltips.TAG_SEGS_IMAGE), + _comfy_io.SEGS.Input("segs", tooltip=tooltips.TAG_SEGS_SEGS), + _comfy_io.Clip.Input("clip", tooltip=tooltips.TAG_SEGS_CLIP), + WD14TaggerIO.Input( + "wd14_tagger", + tooltip=tooltips.TAG_SEGS_WD14_TAGGER, + ), + _comfy_io.String.Input( + "universal_positive", + multiline=False, + default="", + tooltip=tooltips.TAG_SEGS_UNIVERSAL_POSITIVE, + ), + _comfy_io.Float.Input( + "threshold", + default=0.35, + min=0.0, + max=1.0, + step=0.05, + tooltip=tooltips.TILE_THRESHOLD, + ), + _comfy_io.Float.Input( + "character_threshold", + default=1.0, + min=0.0, + max=1.0, + step=0.05, + tooltip=tooltips.TILE_CHARACTER_THRESHOLD, + ), + _comfy_io.Boolean.Input( + "replace_underscore", + default=True, + tooltip=tooltips.TILE_REPLACE_UNDERSCORE, + ), + _comfy_io.Boolean.Input( + "trailing_comma", + default=False, + tooltip=tooltips.TILE_TRAILING_COMMA, + ), + _comfy_io.String.Input( + "exclude_tags", + multiline=False, + default=DEFAULT_EXCLUDE_TAGS, + tooltip=tooltips.TILE_EXCLUDE_TAGS, + ), + ], + outputs=[ + _comfy_io.SEGS.Output( + "segs", + tooltip=tooltips.TAG_SEGS_SEGS_OUTPUT, + ), + ConditioningBatchIO.Output( + "positive", + tooltip=tooltips.TAG_SEGS_POSITIVE_OUTPUT, + ), + ], + ) + + @classmethod + def execute( + cls, + image: object, + segs: object, + clip: Any, + wd14_tagger: object, + universal_positive: str, + threshold: float, + character_threshold: float, + replace_underscore: bool, + trailing_comma: bool, + exclude_tags: str, + ) -> tuple[object, object]: + """Run the legacy implementation behind the v3 schema.""" + + return TagSEGSWithWD14().tag( + image=image, + segs=segs, + clip=clip, + wd14_tagger=wd14_tagger, + universal_positive=universal_positive, + threshold=threshold, + character_threshold=character_threshold, + replace_underscore=replace_underscore, + trailing_comma=trailing_comma, + exclude_tags=exclude_tags, + ) diff --git a/simple_syrup/services/tag_segs_with_wd14_service.py b/simple_syrup/services/tag_segs_with_wd14_service.py new file mode 100644 index 0000000..a522fa0 --- /dev/null +++ b/simple_syrup/services/tag_segs_with_wd14_service.py @@ -0,0 +1,173 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Application service for WD14 tagging of existing SEGS.""" + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +from typing import Any, Protocol + +import torch + +from ..domain.conditioning_batch import ConditioningBatch +from ..domain.prompt_composition import prefix_prompt +from ..domain.segs import ( + ImpactSegs, + NativeSegs, + Segment, + coerce_segs, + to_impact_compatible_segs, +) +from ..masking.segs_mask_ops import crop_image, validate_single_image +from ..runtime.conditioning_encoding import ComfyConditioningEncoder +from ..runtime.loaded_models import LoadedWD14Tagger, unwrap_wd14_tagger +from ..runtime.progress import ProgressReporter, create_comfy_progress +from ..runtime.wd14_tagger import WD14TagFormattingControls, WD14Tagger +from ..shared.logging import get_logger + +LOGGER = get_logger(__name__) +OPERATION = "Tag SEGS w/ WD14" + + +class WD14TaggingBoundary(Protocol): + """Tag ordered image crops.""" + + def tag_images( + self, + loaded_tagger: LoadedWD14Tagger, + images: tuple[torch.Tensor, ...], + controls: WD14TagFormattingControls, + progress: ProgressReporter | None = None, + ) -> tuple[str, ...]: + """Return one tag string per image in input order.""" + + +class ConditioningEncodingBoundary(Protocol): + """Encode ordered prompts into a conditioning batch.""" + + def encode_batch(self, clip: Any, chunks: tuple[str, ...]) -> ConditioningBatch: + """Return conditioning entries in prompt order.""" + + +@dataclass(frozen=True) +class TagSEGSWithWD14Result: + """Return unchanged SEGS and aligned positive conditioning.""" + + segs: ImpactSegs + positive: ConditioningBatch + + +class TagSEGSWithWD14Service: + """Tag provided SEGS crops and encode aligned regional conditioning.""" + + def __init__( + self, + tagger: WD14TaggingBoundary | None = None, + encoder: ConditioningEncodingBoundary | None = None, + progress_factory: Callable[[int], ProgressReporter] | None = None, + ) -> None: + """Create the service with injectable collaborators for tests.""" + + self._tagger = tagger or WD14Tagger() + self._encoder = encoder or ComfyConditioningEncoder() + self._progress_factory = progress_factory or create_comfy_progress + + def tag( + self, + image: object, + segs: object, + clip: Any, + wd14_tagger: object, + tag_controls: WD14TagFormattingControls, + universal_positive: str, + ) -> TagSEGSWithWD14Result: + """Return original SEGS plus WD14-derived conditioning in segment order.""" + + image_tensor = validate_single_image(image, OPERATION) + native_segs = coerce_segs(segs) + self._validate_segs_target_image(native_segs, image_tensor) + _header, segments = native_segs + if not segments: + raise ValueError("No SEGS were provided for Tag SEGS w/ WD14.") + + loaded_tagger = unwrap_wd14_tagger(wd14_tagger) + progress = self._progress_factory(len(segments) + 2) + progress.update(1) + crops = tuple( + crop_image(image_tensor, segment.crop_region) for segment in segments + ) + tags = self._tagger.tag_images( + loaded_tagger, + crops, + tag_controls, + progress=progress, + ) + if len(tags) != len(segments): + raise ValueError( + f"WD14 tagger returned {len(tags)} tag(s) for {len(segments)} SEGS." + ) + + prompts = tuple(prefix_prompt(universal_positive, tag) for tag in tags) + positive = self._encoder.encode_batch(clip, prompts) + progress.update(1) + if len(positive.entries) != len(segments): + raise ValueError( + "Conditioning encoder returned " + f"{len(positive.entries)} entries for {len(segments)} SEGS." + ) + + LOGGER.info( + "Tag SEGS w/ WD14 pass completed", + extra={ + "operation": "tag_segs_with_wd14", + "segment_count": len(segments), + "wd14_model": loaded_tagger.model_id, + "threshold": tag_controls.threshold, + "character_threshold": tag_controls.character_threshold, + "universal_positive_present": bool(universal_positive.strip()), + }, + ) + return TagSEGSWithWD14Result( + segs=to_impact_compatible_segs(native_segs), + positive=positive, + ) + + def _validate_segs_target_image( + self, + segs: NativeSegs, + image: torch.Tensor, + ) -> None: + """Reject SEGS that cannot be cropped from the provided image.""" + + header, segments = segs + image_height = int(image.shape[1]) + image_width = int(image.shape[2]) + if header != (image_height, image_width): + raise ValueError( + f"{OPERATION} requires SEGS header dimensions to match the image: " + f"SEGS is {header[0]}x{header[1]}, image is " + f"{image_height}x{image_width}." + ) + for index, segment in enumerate(segments): + _validate_segment_crop(segment, index, image_height, image_width) + + +def _validate_segment_crop( + segment: Segment, + index: int, + image_height: int, + image_width: int, +) -> None: + """Reject a segment crop region that falls outside the image.""" + + region = segment.crop_region + if region.right <= image_width and region.bottom <= image_height: + return + raise ValueError( + f"{OPERATION} SEG {index} ('{segment.label}') crop_region must fit " + f"inside the image; got ({region.left}, {region.top}, {region.right}, " + f"{region.bottom}) for image {image_height}x{image_width}." + ) diff --git a/tests/test_batch_region_conditioning_node.py b/tests/test_batch_region_conditioning_node.py new file mode 100644 index 0000000..6f6ef8c --- /dev/null +++ b/tests/test_batch_region_conditioning_node.py @@ -0,0 +1,35 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for the Batch Region Conditioning legacy node.""" + +from __future__ import annotations + +from simple_syrup.domain.conditioning_batch import ConditioningBatch +from simple_syrup.nodes.batch_region_conditioning import BatchRegionConditioning + + +def test_batch_region_conditioning_contract() -> None: + """Batch Region Conditioning exposes a mixed-input legacy contract.""" + + inputs = BatchRegionConditioning.INPUT_TYPES() + + assert BatchRegionConditioning.RETURN_TYPES == ("CONDITIONING_BATCH",) + assert BatchRegionConditioning.RETURN_NAMES == ("batch",) + assert BatchRegionConditioning.FUNCTION == "batch" + assert BatchRegionConditioning.CATEGORY == "SimpleSyrup/Conditioning" + assert list(inputs["required"]) == ["first", "second"] + assert inputs["required"]["first"][0] == "CONDITIONING,CONDITIONING_BATCH" + assert inputs["required"]["second"][0] == "CONDITIONING,CONDITIONING_BATCH" + + +def test_batch_region_conditioning_node_flattens_mixed_inputs() -> None: + """The legacy node flattens batches and normal conditionings in order.""" + + auto = ConditioningBatch(("auto 1", "auto 2")) + hand = "hand 1" + + (batch,) = BatchRegionConditioning().batch(auto, hand) + + assert batch == ConditioningBatch(("auto 1", "auto 2", "hand 1")) diff --git a/tests/test_batch_region_conditioning_v3_node.py b/tests/test_batch_region_conditioning_v3_node.py new file mode 100644 index 0000000..c4c7416 --- /dev/null +++ b/tests/test_batch_region_conditioning_v3_node.py @@ -0,0 +1,50 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for the Batch Region Conditioning Comfy v3 wrapper.""" + +from __future__ import annotations + +from simple_syrup.domain.conditioning_batch import ConditioningBatch +from simple_syrup.nodes_v3.batch_region_conditioning import ( + BatchRegionConditioningV3, +) + + +def test_batch_region_conditioning_v3_schema_uses_mixed_autogrow_inputs() -> None: + """The v3 schema accepts conditioning values and conditioning batches.""" + + schema = BatchRegionConditioningV3.define_schema() + + assert schema.node_id == "SimpleSyrup.BatchRegionConditioning" + assert schema.display_name == "Batch Region Conditioning" + assert schema.category == "SimpleSyrup/Conditioning" + assert [input_item.id for input_item in schema.inputs] == ["conditioning_inputs"] + assert schema.inputs[0].io_type == "COMFY_AUTOGROW_V3" + assert schema.inputs[0].template.prefix == "conditioning" + assert schema.inputs[0].template.min == 2 + assert schema.inputs[0].template.max == 50 + assert schema.inputs[0].template.input.get_io_type() == ( + "CONDITIONING,CONDITIONING_BATCH" + ) + assert [output.id for output in schema.outputs] == ["batch"] + assert schema.outputs[0].io_type == "CONDITIONING_BATCH" + + +def test_batch_region_conditioning_v3_execute_flattens_inputs() -> None: + """The v3 wrapper batches mixed inputs in Autogrow insertion order.""" + + auto = ConditioningBatch(("auto 1", "auto 2")) + hand = "hand 1" + extra = ConditioningBatch(("auto 3",)) + + (batch,) = BatchRegionConditioningV3.execute( + { + "conditioning0": auto, + "conditioning1": hand, + "conditioning2": extra, + } + ) + + assert batch == ConditioningBatch(("auto 1", "auto 2", "hand 1", "auto 3")) diff --git a/tests/test_batch_segs_node.py b/tests/test_batch_segs_node.py new file mode 100644 index 0000000..eb25a1a --- /dev/null +++ b/tests/test_batch_segs_node.py @@ -0,0 +1,53 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for the Batch SEGS legacy node.""" + +from __future__ import annotations + +from typing import cast + +from simple_syrup.domain.segs import BoundingBox, CropRegion, ImpactSegs, Segment +from simple_syrup.nodes.batch_segs import BatchSEGS + + +def test_batch_segs_contract() -> None: + """Batch SEGS exposes a legacy two-input chainable contract.""" + + inputs = BatchSEGS.INPUT_TYPES() + + assert BatchSEGS.RETURN_TYPES == ("SEGS",) + assert BatchSEGS.RETURN_NAMES == ("segs",) + assert BatchSEGS.FUNCTION == "batch" + assert BatchSEGS.CATEGORY == "SimpleSyrup/Detection" + assert list(inputs["required"]) == ["first", "second"] + assert inputs["required"]["first"][0] == "SEGS" + assert inputs["required"]["second"][0] == "SEGS" + + +def test_batch_segs_node_batches_in_input_order() -> None: + """The legacy node returns Impact-compatible batched SEGS.""" + + first = ((8, 8), [_segment("1"), _segment("2")]) + second = ((8, 8), [_segment("3")]) + + (raw_segs,) = BatchSEGS().batch(first, second) + segs = cast(ImpactSegs, raw_segs) + + _header, segments = segs + assert isinstance(segments, list) + assert [segment.label for segment in segments] == ["1", "2", "3"] + + +def _segment(label: str) -> Segment: + """Create a small test segment.""" + + return Segment( + cropped_image=None, + cropped_mask="mask", + confidence=1.0, + crop_region=CropRegion(0, 0, 2, 2), + bbox=BoundingBox(0, 0, 2, 2), + label=label, + ) diff --git a/tests/test_batch_segs_v3_node.py b/tests/test_batch_segs_v3_node.py new file mode 100644 index 0000000..1157c47 --- /dev/null +++ b/tests/test_batch_segs_v3_node.py @@ -0,0 +1,89 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for the Batch SEGS Comfy v3 wrapper.""" + +from __future__ import annotations + +from typing import cast + +import pytest + +from simple_syrup.domain.segs import BoundingBox, CropRegion, Segment +from simple_syrup.nodes_v3.batch_segs import BatchSEGSV3 + + +def test_batch_segs_v3_schema_uses_autogrow_segs_inputs() -> None: + """The v3 schema exposes expandable SEGS inputs.""" + + schema = BatchSEGSV3.define_schema() + + assert schema.node_id == "SimpleSyrup.BatchSEGS" + assert schema.display_name == "Batch SEGS" + assert schema.category == "SimpleSyrup/Detection" + assert [input_item.id for input_item in schema.inputs] == ["segs_inputs"] + assert schema.inputs[0].io_type == "COMFY_AUTOGROW_V3" + assert schema.inputs[0].template.prefix == "segs" + assert schema.inputs[0].template.min == 2 + assert schema.inputs[0].template.max == 50 + assert schema.inputs[0].template.input.io_type == "SEGS" + assert [output.id for output in schema.outputs] == ["segs"] + assert schema.outputs[0].io_type == "SEGS" + + +def test_batch_segs_v3_execute_returns_impact_compatible_segs() -> None: + """The v3 wrapper batches SEGS in Autogrow insertion order.""" + + first = ( + (16, 16), + [ + _segment("1", CropRegion(0, 0, 2, 2)), + _segment("2", CropRegion(2, 0, 4, 2)), + _segment("3", CropRegion(4, 0, 6, 2)), + ], + ) + second = ( + (16, 16), + ( + _segment("4", CropRegion(0, 2, 2, 4)), + _segment("5", CropRegion(2, 2, 4, 4)), + _segment("6", CropRegion(4, 2, 6, 4)), + ), + ) + + (raw_segs,) = BatchSEGSV3.execute({"segs0": first, "segs1": second}) + segs = cast(tuple[tuple[int, int], list[Segment]], raw_segs) + + header, segments = segs + assert header == (16, 16) + assert isinstance(segments, list) + assert [segment.label for segment in segments] == ["1", "2", "3", "4", "5", "6"] + + +def test_batch_segs_v3_execute_surfaces_header_mismatch() -> None: + """The v3 wrapper keeps domain validation errors visible.""" + + first = ((8, 16), (_segment("first", CropRegion(0, 0, 2, 2)),)) + second = ((16, 8), (_segment("second", CropRegion(0, 0, 2, 2)),)) + + with pytest.raises(ValueError, match="input 2 is 16x8 but input 1 is 8x16"): + BatchSEGSV3.execute({"segs0": first, "segs1": second}) + + +def _segment(label: str, crop_region: CropRegion) -> Segment: + """Create a segment for Batch SEGS v3 tests.""" + + return Segment( + cropped_image=None, + cropped_mask="mask", + confidence=1.0, + crop_region=crop_region, + bbox=BoundingBox( + crop_region.left, + crop_region.top, + crop_region.right, + crop_region.bottom, + ), + label=label, + ) diff --git a/tests/test_conditioning_batch.py b/tests/test_conditioning_batch.py index 633dba2..0ec9a6a 100644 --- a/tests/test_conditioning_batch.py +++ b/tests/test_conditioning_batch.py @@ -10,6 +10,7 @@ import pytest from simple_syrup.domain.conditioning_batch import ( ConditioningBatch, + batch_conditioning, select_conditioning, split_prompt_batch, ) @@ -65,6 +66,25 @@ def test_conditioning_batch_rejects_negative_indexes() -> None: ConditioningBatch(("a",)).select(-1) +def test_batch_conditioning_flattens_batches_and_normal_conditioning() -> None: + """Mixed conditioning inputs become one ordered per-region batch.""" + + first = ConditioningBatch(("auto 1", "auto 2")) + hand = "hand 1" + second = ConditioningBatch(("auto 3",)) + + batch = batch_conditioning((first, hand, second)) + + assert batch.entries == ("auto 1", "auto 2", "hand 1", "auto 3") + + +def test_batch_conditioning_rejects_no_inputs() -> None: + """At least one input is needed to build a conditioning batch.""" + + with pytest.raises(ValueError, match="one or more inputs"): + batch_conditioning(()) + + def test_select_conditioning_broadcasts_normal_conditioning() -> None: """Normal conditionings pass through unchanged for any valid index.""" diff --git a/tests/test_node_tooltips.py b/tests/test_node_tooltips.py index 20713d8..7b9d8ea 100644 --- a/tests/test_node_tooltips.py +++ b/tests/test_node_tooltips.py @@ -14,6 +14,8 @@ from typing import Any, Protocol, cast import pytest from simple_syrup.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS +from simple_syrup.nodes_v3.batch_region_conditioning import BatchRegionConditioningV3 +from simple_syrup.nodes_v3.batch_segs import BatchSEGSV3 from simple_syrup.nodes_v3.encode_prompt_batch_with_prompt_control import ( EncodePromptBatchWithPromptControl, ) @@ -22,6 +24,7 @@ from simple_syrup.nodes_v3.schedule_and_encode_prompts_with_prompt_control impor ScheduleAndEncodePromptsWithPromptControl, ) from simple_syrup.nodes_v3.simple_load_checkpoint import SimpleLoadCheckpointV3 +from simple_syrup.nodes_v3.tag_segs_with_wd14 import TagSEGSWithWD14V3 from simple_syrup.nodes_v3.tile_and_tag_segs import TileAndTagSEGSV3 from simple_syrup.nodes_v3.vae_decode_options import VAEDecodeOptionsV3 from simple_syrup.nodes_v3.vae_encode_options import VAEEncodeOptionsV3 @@ -120,6 +123,9 @@ def test_legacy_named_outputs_provide_tooltips() -> None: [ SimpleLoadCheckpointV3, ScaleFactorV3, + BatchSEGSV3, + BatchRegionConditioningV3, + TagSEGSWithWD14V3, TileAndTagSEGSV3, VAEDecodeOptionsV3, VAEEncodeOptionsV3, diff --git a/tests/test_registration.py b/tests/test_registration.py index 350554f..1ad5a75 100644 --- a/tests/test_registration.py +++ b/tests/test_registration.py @@ -137,6 +137,29 @@ def test_ksampler_tiled_diffusion_node_is_registered() -> None: ) +def test_batch_segs_node_is_registered() -> None: + """Batch SEGS node maps to its class and display name.""" + + package = importlib.import_module("SimpleSyrup") + registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.BatchSEGS"] + + assert registered.__name__ == "BatchSEGS" + assert package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.BatchSEGS"] == "Batch SEGS" + + +def test_batch_region_conditioning_node_is_registered() -> None: + """Batch Region Conditioning node maps to its class and display name.""" + + package = importlib.import_module("SimpleSyrup") + registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.BatchRegionConditioning"] + + assert registered.__name__ == "BatchRegionConditioning" + assert ( + package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.BatchRegionConditioning"] + == "Batch Region Conditioning" + ) + + def test_latent_diagnostics_node_is_registered() -> None: """Latent Diagnostics node maps to its class and display name.""" @@ -392,6 +415,19 @@ def test_tile_and_tag_segs_node_is_registered() -> None: ) +def test_tag_segs_with_wd14_node_is_registered() -> None: + """Tag SEGS w/ WD14 node maps to its class and display name.""" + + package = importlib.import_module("SimpleSyrup") + registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.TagSEGSWithWD14"] + + assert registered.__name__ == "TagSEGSWithWD14" + assert ( + package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.TagSEGSWithWD14"] + == "Tag SEGS w/ WD14" + ) + + def test_conditioning_batch_nodes_are_registered() -> None: """Conditioning batch nodes map to their classes and display names.""" @@ -548,6 +584,9 @@ def test_v3_entrypoint_registers_tile_and_prompt_control_batch_nodes( assert [node.__name__ for node in nodes] == [ "WD14TaggerLoaderV3", + "BatchSEGSV3", + "BatchRegionConditioningV3", + "TagSEGSWithWD14V3", "TileAndTagSEGSV3", "SimpleLoadCheckpointV3", "ScaleFactorV3", @@ -574,6 +613,9 @@ def test_v3_entrypoint_keeps_tile_node_when_prompt_control_unavailable( assert [node.__name__ for node in nodes] == [ "WD14TaggerLoaderV3", + "BatchSEGSV3", + "BatchRegionConditioningV3", + "TagSEGSWithWD14V3", "TileAndTagSEGSV3", "SimpleLoadCheckpointV3", "ScaleFactorV3", diff --git a/tests/test_segs_domain.py b/tests/test_segs_domain.py index bbfce9f..f073ca5 100644 --- a/tests/test_segs_domain.py +++ b/tests/test_segs_domain.py @@ -16,6 +16,7 @@ from simple_syrup.domain.segs import ( BoundingBox, CropRegion, Segment, + batch_segs, coerce_segment, coerce_segs, coerce_segs_group, @@ -155,6 +156,73 @@ def test_impact_segs_group_conversion_returns_list_outputs() -> None: assert all(isinstance(segments, list) for _header, segments in output) +def test_batch_segs_preserves_input_and_segment_order() -> None: + """Batch SEGS flattens input payloads without reordering segments.""" + + first = ( + (16, 16), + ( + _segment("1", CropRegion(0, 0, 2, 2), 0.9), + _segment("2", CropRegion(2, 0, 4, 2), 0.8), + _segment("3", CropRegion(4, 0, 6, 2), 0.7), + ), + ) + second = ( + (16, 16), + [ + _segment("4", CropRegion(0, 2, 2, 4), 0.6), + _segment("5", CropRegion(2, 2, 4, 4), 0.5), + _segment("6", CropRegion(4, 2, 6, 4), 0.4), + ], + ) + + header, segments = batch_segs((first, second)) + + assert header == (16, 16) + assert [segment.label for segment in segments] == ["1", "2", "3", "4", "5", "6"] + + +def test_batch_segs_allows_empty_payloads() -> None: + """Empty SEGS inputs contribute no segments to the batched payload.""" + + first = ( + (16, 16), + ( + _segment("1", CropRegion(0, 0, 2, 2), 0.9), + _segment("2", CropRegion(2, 0, 4, 2), 0.8), + ), + ) + empty = ((16, 16), ()) + third = ((16, 16), (_segment("3", CropRegion(4, 0, 6, 2), 0.7),)) + + _header, segments = batch_segs((first, empty, third)) + + assert [segment.label for segment in segments] == ["1", "2", "3"] + + +def test_batch_segs_returns_empty_payload_when_all_inputs_are_empty() -> None: + """All-empty SEGS inputs keep the shared header and return no segments.""" + + assert batch_segs((((16, 16), ()), ((16, 16), []))) == ((16, 16), ()) + + +def test_batch_segs_rejects_no_inputs() -> None: + """Batch SEGS requires at least one payload for an output header.""" + + with pytest.raises(ValueError, match="one or more SEGS inputs"): + batch_segs(()) + + +def test_batch_segs_rejects_mismatched_headers() -> None: + """Batch SEGS refuses to merge regions targeting different image sizes.""" + + first = ((8, 16), (_segment("first", CropRegion(0, 0, 2, 2), 0.9),)) + second = ((16, 8), (_segment("second", CropRegion(0, 0, 2, 2), 0.8),)) + + with pytest.raises(ValueError, match="input 2 is 16x8 but input 1 is 8x16"): + batch_segs((first, second)) + + def test_sort_order_options_are_plain_english_and_ordered() -> None: """SEGS sort options match the detector node combo contract.""" diff --git a/tests/test_tag_segs_with_wd14_node.py b/tests/test_tag_segs_with_wd14_node.py new file mode 100644 index 0000000..30bf609 --- /dev/null +++ b/tests/test_tag_segs_with_wd14_node.py @@ -0,0 +1,112 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for the Tag SEGS w/ WD14 node contract.""" + +from __future__ import annotations + +from typing import Any, cast + +import torch + +from simple_syrup.domain.conditioning_batch import ConditioningBatch +from simple_syrup.domain.segs import BoundingBox, CropRegion, ImpactSegs, Segment +from simple_syrup.nodes.tag_segs_with_wd14 import TagSEGSWithWD14 +from simple_syrup.nodes.tile_and_tag_segs import DEFAULT_EXCLUDE_TAGS +from simple_syrup.runtime.wd14_tagger import WD14TagFormattingControls +from simple_syrup.services.tag_segs_with_wd14_service import TagSEGSWithWD14Result + + +def test_tag_segs_with_wd14_contract() -> None: + """Tag SEGS w/ WD14 exposes the agreed ComfyUI contract.""" + + inputs = TagSEGSWithWD14.INPUT_TYPES() + + assert TagSEGSWithWD14.RETURN_TYPES == ("SEGS", "CONDITIONING_BATCH") + assert TagSEGSWithWD14.RETURN_NAMES == ("segs", "positive") + assert TagSEGSWithWD14.FUNCTION == "tag" + assert TagSEGSWithWD14.CATEGORY == "SimpleSyrup/Detailing" + assert list(inputs["required"]) == [ + "image", + "segs", + "clip", + "wd14_tagger", + "universal_positive", + "threshold", + "character_threshold", + "replace_underscore", + "trailing_comma", + "exclude_tags", + ] + assert inputs["required"]["segs"][0] == "SEGS" + assert inputs["required"]["clip"][0] == "CLIP" + assert inputs["required"]["wd14_tagger"][0] == "WD14_TAGGER" + assert inputs["required"]["universal_positive"][0] == "STRING" + assert inputs["required"]["universal_positive"][1]["default"] == "" + assert inputs["required"]["threshold"][1]["default"] == 0.35 + assert inputs["required"]["character_threshold"][1]["default"] == 1.0 + assert inputs["required"]["replace_underscore"][1]["default"] is True + assert inputs["required"]["trailing_comma"][1]["default"] is False + assert inputs["required"]["exclude_tags"][1]["default"] == DEFAULT_EXCLUDE_TAGS + assert "optional" not in inputs + + +def test_tag_segs_with_wd14_delegates_to_service(monkeypatch: Any) -> None: + """The node delegates behavior and returns service outputs unchanged.""" + + service = _FakeService() + monkeypatch.setattr(TagSEGSWithWD14, "service_class", lambda: service) + image = torch.zeros((1, 8, 8, 3)) + segs: ImpactSegs = ((8, 8), []) + wd14_tagger = object() + + output_segs, positive = TagSEGSWithWD14().tag( + image=image, + segs=segs, + clip="clip", + wd14_tagger=wd14_tagger, + universal_positive="masterpiece", + threshold=0.35, + character_threshold=1.0, + replace_underscore=True, + trailing_comma=False, + exclude_tags=DEFAULT_EXCLUDE_TAGS, + ) + + assert output_segs is service.result.segs + assert positive is service.result.positive + assert service.call["image"] is image + assert service.call["segs"] is segs + assert service.call["clip"] == "clip" + assert service.call["wd14_tagger"] is wd14_tagger + assert service.call["universal_positive"] == "masterpiece" + tag_controls = cast(WD14TagFormattingControls, service.call["tag_controls"]) + assert tag_controls.threshold == 0.35 + + +class _FakeService: + """Capture node calls for delegation tests.""" + + def __init__(self) -> None: + """Create a fake service result.""" + + segment = Segment( + cropped_image=None, + cropped_mask=torch.ones((8, 8)), + confidence=1.0, + crop_region=CropRegion(0, 0, 8, 8), + bbox=BoundingBox(0, 0, 8, 8), + label="seg_001", + ) + self.result = TagSEGSWithWD14Result( + segs=((8, 8), [segment]), + positive=ConditioningBatch(("encoded",)), + ) + self.call: dict[str, object] = {} + + def tag(self, **kwargs: object) -> TagSEGSWithWD14Result: + """Return a fixed result and remember provided inputs.""" + + self.call = kwargs + return self.result diff --git a/tests/test_tag_segs_with_wd14_service.py b/tests/test_tag_segs_with_wd14_service.py new file mode 100644 index 0000000..24827f7 --- /dev/null +++ b/tests/test_tag_segs_with_wd14_service.py @@ -0,0 +1,283 @@ +# 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 WD14 tagging of existing SEGS.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import pytest +import torch + +from simple_syrup.domain.conditioning_batch import ConditioningBatch +from simple_syrup.domain.segs import BoundingBox, CropRegion, NativeSegs, Segment +from simple_syrup.runtime.loaded_models import LoadedWD14Tagger +from simple_syrup.runtime.wd14_tagger import ( + FloatArray, + WD14TagFormattingControls, + WD14TagRecord, +) +from simple_syrup.services.tag_segs_with_wd14_service import TagSEGSWithWD14Service + + +def test_service_preserves_existing_segs_tag_and_conditioning_order() -> None: + """Existing SEGS, crops, tags, and conditioning stay aligned by index.""" + + progress = _ProgressRecorder() + tagger = _FakeTagger(("tag first", "", "tag third")) + encoder = _FakeEncoder() + loaded_tagger = _loaded_tagger() + service = TagSEGSWithWD14Service( + tagger=tagger, + encoder=encoder, + progress_factory=lambda _total: progress, + ) + segs = _native_segs(("first", "second", "third")) + + result = service.tag( + image=_image(), + segs=segs, + clip="clip", + wd14_tagger=loaded_tagger, + tag_controls=_tag_controls(), + universal_positive="masterpiece", + ) + + assert [segment.label for segment in result.segs[1]] == [ + "first", + "second", + "third", + ] + assert [tuple(crop.shape) for crop in tagger.crops] == [ + (1, 2, 2, 3), + (1, 2, 2, 3), + (1, 2, 2, 3), + ] + assert tagger.loaded_tagger is loaded_tagger + assert encoder.chunks == ( + "masterpiece, tag first", + "masterpiece", + "masterpiece, tag third", + ) + assert result.positive.entries == ( + "clip:masterpiece, tag first", + "clip:masterpiece", + "clip:masterpiece, tag third", + ) + assert progress.updates == [1, 3, 1] + + +def test_service_rejects_empty_segs() -> None: + """Tagging empty SEGS would not produce a selectable conditioning batch.""" + + service = TagSEGSWithWD14Service( + tagger=_FakeTagger(()), + encoder=_FakeEncoder(), + ) + + with pytest.raises(ValueError, match="No SEGS"): + service.tag( + image=_image(), + segs=((4, 4), ()), + clip="clip", + wd14_tagger=_loaded_tagger(), + tag_controls=_tag_controls(), + universal_positive="", + ) + + +def test_service_rejects_segs_image_header_mismatch() -> None: + """SEGS must describe the image being cropped for tagging.""" + + service = TagSEGSWithWD14Service( + tagger=_FakeTagger(("tag",)), + encoder=_FakeEncoder(), + ) + + with pytest.raises(ValueError, match="SEGS is 8x4, image is 4x4"): + service.tag( + image=_image(), + segs=((8, 4), _native_segs(("first",))[1]), + clip="clip", + wd14_tagger=_loaded_tagger(), + tag_controls=_tag_controls(), + universal_positive="", + ) + + +def test_service_rejects_tagger_count_mismatch() -> None: + """Dropping a tag would break SEGS alignment and is rejected.""" + + service = TagSEGSWithWD14Service( + tagger=_FakeTagger(("only one",)), + encoder=_FakeEncoder(), + ) + + with pytest.raises(ValueError, match="returned 1 tag"): + service.tag( + image=_image(), + segs=_native_segs(("first", "second")), + clip="clip", + wd14_tagger=_loaded_tagger(), + tag_controls=_tag_controls(), + universal_positive="", + ) + + +def test_service_rejects_conditioning_count_mismatch() -> None: + """Dropping encoded conditioning would break SEGS alignment and is rejected.""" + + service = TagSEGSWithWD14Service( + tagger=_FakeTagger(("first", "second")), + encoder=_ShortEncoder(), + ) + + with pytest.raises(ValueError, match="returned 1 entries for 2 SEGS"): + service.tag( + image=_image(), + segs=_native_segs(("first", "second")), + clip="clip", + wd14_tagger=_loaded_tagger(), + tag_controls=_tag_controls(), + universal_positive="", + ) + + +class _FakeTagger: + """Return fixed tag strings for ordered crops.""" + + def __init__(self, tags: tuple[str, ...]) -> None: + """Store the fixed tags.""" + + self.tags = tags + self.crops: tuple[torch.Tensor, ...] = () + self.loaded_tagger: LoadedWD14Tagger | None = None + + def tag_images( + self, + loaded_tagger: LoadedWD14Tagger, + images: tuple[torch.Tensor, ...], + controls: WD14TagFormattingControls, + progress: object | None = None, + ) -> tuple[str, ...]: + """Return fixed tags and remember the crop order.""" + + _ = controls + if progress is not None: + progress.update(len(images)) # type: ignore[attr-defined] + self.loaded_tagger = loaded_tagger + self.crops = images + return self.tags + + +class _FakeEncoder: + """Return visible conditioning values for prompt chunks.""" + + def __init__(self) -> None: + """Initialize captured chunks.""" + + self.chunks: tuple[str, ...] = () + + def encode_batch(self, clip: Any, chunks: tuple[str, ...]) -> ConditioningBatch: + """Encode prompts as simple strings.""" + + self.chunks = chunks + return ConditioningBatch(tuple(f"{clip}:{chunk}" for chunk in chunks)) + + +class _ShortEncoder: + """Return too few conditioning entries for validation tests.""" + + def encode_batch(self, clip: Any, chunks: tuple[str, ...]) -> ConditioningBatch: + """Encode only the first prompt chunk.""" + + _ = clip + return ConditioningBatch((chunks[0],)) + + +class _ProgressRecorder: + """Record service progress updates.""" + + def __init__(self) -> None: + """Initialize captured update values.""" + + self.updates: list[int] = [] + + def update(self, value: int) -> None: + """Record one progress advance.""" + + self.updates.append(value) + + +def _native_segs(labels: tuple[str, ...]) -> NativeSegs: + """Create native SEGS with stable two-pixel crop regions.""" + + segments = tuple( + Segment( + cropped_image=None, + cropped_mask=torch.ones((2, 2)), + confidence=1.0, + crop_region=CropRegion(index, index, index + 2, index + 2), + bbox=BoundingBox(index, index, index + 2, index + 2), + label=label, + ) + for index, label in enumerate(labels) + ) + return (4, 4), segments + + +def _image() -> torch.Tensor: + """Return a small deterministic BHWC image.""" + + return torch.arange(4 * 4 * 3, dtype=torch.float32).reshape(1, 4, 4, 3) / 255.0 + + +def _tag_controls() -> WD14TagFormattingControls: + """Return valid WD14 controls for service tests.""" + + return WD14TagFormattingControls( + threshold=0.35, + character_threshold=1.0, + replace_underscore=True, + trailing_comma=False, + exclude_tags="", + ) + + +def _loaded_tagger() -> LoadedWD14Tagger: + """Return a reusable loaded WD14 tagger test container.""" + + return LoadedWD14Tagger( + model_id="wd-eva02-large-tagger-v3", + source="test", + onnx_path=Path("wd-eva02-large-tagger-v3.onnx"), + csv_path=Path("wd-eva02-large-tagger-v3.csv"), + providers=("CPUExecutionProvider",), + session=_FakeWD14Session(), + tags=(WD14TagRecord("blue_hair", "0"),), + ) + + +class _FakeWD14Session: + """Minimal WD14 session test double.""" + + def get_inputs(self) -> list[object]: + """Return no fake inputs.""" + + return [] + + def get_outputs(self) -> list[object]: + """Return no fake outputs.""" + + return [] + + def run( + self, output_names: list[str], feeds: dict[str, FloatArray] + ) -> list[object]: + """Return no fake outputs.""" + + _ = output_names, feeds + return [] diff --git a/tests/test_tag_segs_with_wd14_v3_node.py b/tests/test_tag_segs_with_wd14_v3_node.py new file mode 100644 index 0000000..b89b9ec --- /dev/null +++ b/tests/test_tag_segs_with_wd14_v3_node.py @@ -0,0 +1,103 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for the Tag SEGS w/ WD14 Comfy v3 wrapper.""" + +from __future__ import annotations + +from typing import Any + +import torch + +from simple_syrup.domain.conditioning_batch import ConditioningBatch +from simple_syrup.domain.segs import BoundingBox, CropRegion, ImpactSegs, Segment +from simple_syrup.nodes.tag_segs_with_wd14 import TagSEGSWithWD14 +from simple_syrup.nodes_v3.tag_segs_with_wd14 import TagSEGSWithWD14V3 +from simple_syrup.services.tag_segs_with_wd14_service import TagSEGSWithWD14Result + + +def test_tag_segs_with_wd14_v3_schema_includes_clip_and_wd14_tagger() -> None: + """The v3 schema exposes existing-SEGS WD14 tagging inputs.""" + + schema = TagSEGSWithWD14V3.define_schema() + + assert schema.node_id == "SimpleSyrup.TagSEGSWithWD14" + assert schema.display_name == "Tag SEGS w/ WD14" + assert [input_item.id for input_item in schema.inputs][:4] == [ + "image", + "segs", + "clip", + "wd14_tagger", + ] + assert schema.inputs[1].io_type == "SEGS" + assert schema.inputs[2].io_type == "CLIP" + assert schema.inputs[3].io_type == "WD14_TAGGER" + universal_positive = schema.inputs[4] + assert universal_positive.io_type == "STRING" + assert universal_positive.default == "" + assert universal_positive.multiline is False + assert [output.id for output in schema.outputs] == ["segs", "positive"] + assert [output.io_type for output in schema.outputs] == [ + "SEGS", + "CONDITIONING_BATCH", + ] + + +def test_tag_segs_with_wd14_v3_execute_forwards_to_legacy_node( + monkeypatch: Any, +) -> None: + """The v3 wrapper forwards execution to the legacy implementation.""" + + service = _FakeService() + monkeypatch.setattr(TagSEGSWithWD14, "service_class", lambda: service) + image = torch.zeros((1, 8, 8, 3)) + segs: ImpactSegs = ((8, 8), []) + wd14_tagger = object() + + output_segs, positive = TagSEGSWithWD14V3.execute( + image=image, + segs=segs, + clip="clip", + wd14_tagger=wd14_tagger, + universal_positive="masterpiece", + threshold=0.35, + character_threshold=1.0, + replace_underscore=True, + trailing_comma=False, + exclude_tags="", + ) + + assert output_segs is service.result.segs + assert positive is service.result.positive + assert service.call["segs"] is segs + assert service.call["clip"] == "clip" + assert service.call["wd14_tagger"] is wd14_tagger + assert service.call["universal_positive"] == "masterpiece" + + +class _FakeService: + """Capture v3 wrapper calls through the legacy node.""" + + def __init__(self) -> None: + """Create a fake service result.""" + + segment = Segment( + cropped_image=None, + cropped_mask=torch.ones((8, 8)), + confidence=1.0, + crop_region=CropRegion(0, 0, 8, 8), + bbox=BoundingBox(0, 0, 8, 8), + label="seg_001", + ) + self.result = TagSEGSWithWD14Result( + segs=((8, 8), [segment]), + positive=ConditioningBatch(("encoded",)), + ) + self.call: dict[str, object] = {} + + def tag(self, **kwargs: object) -> TagSEGSWithWD14Result: + """Return a fixed result and remember provided inputs.""" + + self.call = kwargs + return self.result diff --git a/tests/test_tiled_sampling_runtime.py b/tests/test_tiled_sampling_runtime.py index c71082d..c319319 100644 --- a/tests/test_tiled_sampling_runtime.py +++ b/tests/test_tiled_sampling_runtime.py @@ -54,7 +54,8 @@ def test_validate_latent_samples_rejects_non_tensor() -> None: def test_validate_tensor_shape_rejects_nested_tensor() -> None: """Nested tensors are rejected before spatial tiling.""" - samples = torch.nested.nested_tensor([torch.zeros((4, 8, 8))]) + with pytest.warns(UserWarning, match="nested tensors.*prototype stage"): + samples = torch.nested.nested_tensor([torch.zeros((4, 8, 8))]) with pytest.raises(ValueError, match="non-nested latent samples"): tiled_sampling.validate_tensor_shape(samples, sampler_label="TestSampler")