feat(segs): add regional batching and wd14 tagging nodes

This commit is contained in:
Artificial Sweetener
2026-05-30 21:50:35 -04:00
parent 8e553f5662
commit ae2dbc7e1a
26 changed files with 2203 additions and 1 deletions
+499
View File
@@ -0,0 +1,499 @@
# Batch SEGS and WD14 SEGS Tagging Plan
## Goal
Add a detector-first regional detailing workflow:
1. Upstream detector nodes create `SEGS` payloads. These can come from multiple Ultralytics detector models, prompted SAM, or any other Impact-compatible SEGS producer.
2. A new **Batch SEGS** node combines those ordered `SEGS` payloads into one ordered `SEGS` payload.
3. A future **Tag SEGS w/ WD14** node runs WD14 over each segment crop and produces a `CONDITIONING_BATCH` aligned to the same segment order.
4. The existing **Detail SEGS as Regions** node receives the batched `SEGS` plus aligned regional conditioning and details all regions in one regional MultiDiffusion pass.
The immediate implementation target is **Batch SEGS**. The WD14 tagging node is included here so the Batch SEGS design supports the intended downstream workflow.
## Implementation Status
- [x] Domain batching policy added in `simple_syrup/domain/segs.py`.
- Landing note: `batch_segs(...)` returns native immutable SEGS and leaves Impact-compatible conversion to Comfy-facing code.
- [x] Domain tests added in `tests/test_segs_domain.py`.
- Landing note: tests cover order-preserving flattening, empty SEGS inputs, all-empty output, no-input rejection, and mismatched headers.
- [x] **Batch SEGS** v3 node added in `simple_syrup/nodes_v3/batch_segs.py`.
- Landing note: the node uses `io.Autogrow.TemplatePrefix` with `segs0` through `segs49`, matching Comfy core batch-node behavior.
- [x] **Batch SEGS** exported through `simple_syrup/nodes_v3/__init__.py`.
- Landing note: the node is intentionally v3-only and is not registered in legacy mappings.
- [x] V3 node and entrypoint tests added or updated.
- Landing note: `tests/test_batch_segs_v3_node.py` covers schema and execution; `tests/test_registration.py` covers v3 entrypoint visibility.
- [x] Focused verification passed.
- Landing note: `..\..\venv\Scripts\python.exe -m pytest -n auto -q tests\test_segs_domain.py tests\test_batch_segs_v3_node.py tests\test_registration.py` passed with 78 tests.
- [x] Format, lint, type check, and full test suite passed for **Batch SEGS**.
- Landing note: initial Batch SEGS gates passed before later workflow scope additions.
- [x] **Tag SEGS w/ WD14** implemented after scope correction.
- Landing note: `TagSEGSWithWD14Service` crops existing SEGS in order, runs WD14, prefixes prompts, CLIP-encodes aligned conditioning, and returns the original SEGS with `CONDITIONING_BATCH`.
- [x] **Tag SEGS w/ WD14** exported through legacy mappings and v3 entrypoint.
- Landing note: unlike **Batch SEGS**, this node is not v3-only because it does not need Autogrow and should be available through every repository-supported export path.
- [x] **Tag SEGS w/ WD14** focused verification passed.
- Landing note: `..\..\venv\Scripts\python.exe -m pytest -n auto -q tests\test_tag_segs_with_wd14_service.py tests\test_tag_segs_with_wd14_node.py tests\test_tag_segs_with_wd14_v3_node.py tests\test_registration.py tests\test_node_tooltips.py` passed with 57 tests.
- [x] Final gates rerun after **Tag SEGS w/ WD14**.
- Landing note: gates passed after both SEGS batching and WD14 tagging were implemented.
- [x] **Batch Region Conditioning** implemented after mixed auto/manual conditioning requirement.
- Landing note: v3-only Autogrow node accepts mixed `CONDITIONING` and `CONDITIONING_BATCH` sockets, flattening existing batches and independent conditionings into one ordered `CONDITIONING_BATCH`.
- [x] Final gates rerun after **Batch Region Conditioning**.
- Landing note: gates passed after the mixed conditioning batcher was added.
- [x] Legacy-visible **Batch SEGS** and **Batch Region Conditioning** added after Comfy menu visibility check.
- Landing note: the Autogrow v3 nodes did not appear in the maintainer's current Comfy menu. Legacy two-input chainable nodes were added and registered so both utilities appear through `NODE_CLASS_MAPPINGS`.
- [x] Final gates rerun after legacy-visible batch nodes.
- Landing note: gates passed after legacy Batch SEGS and Batch Region Conditioning were added.
- [x] Warning handling made explicit before commit.
- Landing note: pytest now treats unhandled warnings as errors, narrowly filters known third-party SWIG import deprecations, and asserts the intentional PyTorch nested tensor warning in the owning test. Final full suite result: 920 passed with no warning summary.
## Decisions From Maintainer Discussion
- The node name is **Batch SEGS**.
- **Batch SEGS** must be its own node, not folded into **Tile & Tag SEGS**.
- **Tile & Tag SEGS** remains for generated tile SEGS. The new workflow is for existing detector-produced SEGS.
- **Batch SEGS** must support expandable inputs like ComfyUI core nodes such as **Math Expression**, **Batch Images**, **Batch Masks**, and **Batch Latents**.
- Use Comfy v3 `io.Autogrow` for the expandable input UI.
- It is acceptable for **Batch SEGS** to be v3-only. Do not add an awkward fixed-input legacy fallback.
- **Batch SEGS** must accept existing SEGS batches and flatten them in order. If one input contains segments `1 2 3` and another input contains `4 5 6`, the output must contain `1 2 3 4 5 6`.
- Segment order matters because downstream conditioning is index-aligned to segments.
- All inputs must target the same source image dimensions. Reject mismatched `SEGS` headers.
## Existing Code Context
Important files:
- `simple_syrup/domain/segs.py`
- Owns SEGS coercion, Impact compatibility, sorting, and limiting.
- Existing functions include `coerce_segs`, `coerce_segs_group`, `to_impact_compatible_segs`, and `to_impact_compatible_segs_group`.
- `simple_syrup/nodes_v3/__init__.py`
- Exports v3 nodes through `get_nodes()`.
- `simple_syrup/nodes_v3/tile_and_tag_segs.py`
- Shows the local v3 wrapper pattern.
- `simple_syrup/nodes/__init__.py`
- Legacy mapping exports. **Batch SEGS** does not need to be added here if implemented as v3-only.
- `simple_syrup/services/tile_and_tag_segs_service.py`
- Existing WD14 crop-tag-encode orchestration for generated tile SEGS.
- `simple_syrup/nodes/detail_segs_as_regions.py`
- Existing downstream regional detailer.
- It consumes `SEGS` plus `CONDITIONING_BATCH`.
- `simple_syrup/services/detail_segs_as_regions_service.py`
- Validates and pairs region conditioning by segment index.
- `tests/test_segs_domain.py`
- Existing SEGS domain tests.
- `tests/test_registration.py`
- Existing registration and v3 entrypoint tests.
Comfy reference examples:
- `E:\ComfyUI\comfy_extras\nodes_math.py`
- `MathExpressionNode` uses `io.Autogrow.TemplateNames`.
- `E:\ComfyUI\comfy_extras\nodes_post_processing.py`
- `BatchImagesNode`, `BatchMasksNode`, and `BatchLatentsNode` use `io.Autogrow.TemplatePrefix`.
- `E:\ComfyUI\comfy_api\latest\_io.py`
- `Autogrow` implementation and dynamic input expansion.
## Architecture Constraints
- Keep Comfy-facing node code thin.
- Put merge policy in the domain layer, not in the node wrapper.
- Domain logic must not import ComfyUI modules.
- Runtime adapters own external system interaction. **Batch SEGS** should not need runtime adapters.
- Preserve public behavior of existing nodes.
- Do not rename existing node IDs, display names, categories, inputs, outputs, or return shapes.
- Do not add internal compatibility shims.
- New or changed code must have docstrings.
- New behavior must be covered by tests.
- Use the ComfyUI virtual environment two directories above the repository for verification.
## Batch SEGS Functional Specification
### Node
Display name:
```text
Batch SEGS
```
Suggested v3 node id:
```text
SimpleSyrup.BatchSEGS
```
Suggested category:
```text
SimpleSyrup/Detection
```
Search aliases:
```text
batch, merge, join, combine, segs
```
Inputs:
- Expandable Autogrow group named `segs_inputs`.
- Template input type: `SEGS`.
- Prefix: `segs`.
- Minimum visible/required inputs: `2`.
- Maximum inputs: `50`, matching Comfy core batch nodes.
Outputs:
- `SEGS`
- Return name: `segs`
### Behavior
Given:
```text
segs0 = ((height, width), [segment1, segment2, segment3])
segs1 = ((height, width), [segment4, segment5, segment6])
```
Return:
```text
((height, width), [segment1, segment2, segment3, segment4, segment5, segment6])
```
Rules:
- Coerce every input using existing SEGS compatibility handling.
- Preserve input order from the Autogrow dict values.
- Preserve order inside each input SEGS payload.
- Reject an empty Autogrow group. This should not happen with `min=2`, but the domain function must still fail clearly if called directly with no inputs.
- Reject headers that do not match exactly.
- Allow empty SEGS payloads. They contribute no segments.
- If every input has no segments, return an empty SEGS payload with the shared header.
- Return an Impact-compatible tuple/list shape.
Error examples:
- No inputs:
- `Batch SEGS requires one or more SEGS inputs.`
- Header mismatch:
- `Batch SEGS requires all SEGS inputs to use the same image size; input 2 is 1024x768 but input 1 is 768x1024.`
Use height-first wording because SEGS headers are `(height, width)`.
## Batch SEGS Implementation Steps
### 1. Add Domain Function
Edit `simple_syrup/domain/segs.py`.
Status: completed.
Add a function similar to:
```python
def batch_segs(values: Iterable[object]) -> NativeSegs:
"""Return one SEGS payload containing all segments in input order."""
```
Implementation notes:
- Convert `values` to a tuple once so emptiness and indexing are deterministic.
- Use `coerce_segs` for each input.
- Track the first header as the expected header.
- Compare every later header to the first header.
- Extend a local `list[Segment]`.
- Return `(expected_header, tuple(segments))`.
- Do not call `to_impact_compatible_segs` in the domain function. Keep the domain return type native.
- Add `batch_segs` to imports where needed. Updating `__all__` is not necessary because this module does not currently define one.
### 2. Add V3 Node
Create `simple_syrup/nodes_v3/batch_segs.py`.
Status: completed.
Use the existing v3 import pattern from `tile_and_tag_segs.py`:
- Use `TYPE_CHECKING`.
- Import `comfy_api.latest` lazily through `import_module`.
- Define a `_ComfyNodeBase` type-checking shim.
- Use `_comfy_io.SEGS`.
Schema sketch:
```python
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.",
),
],
)
```
Execute sketch:
```python
@classmethod
def execute(cls, segs_inputs: Any) -> tuple[object]:
"""Batch provided SEGS inputs in Autogrow order."""
native = batch_segs(segs_inputs.values())
return (to_impact_compatible_segs(native),)
```
Do not use `io.NodeOutput` unless existing SimpleSyrup v3 wrappers are migrated. Current local wrappers return normal tuples, so match local style.
### 3. Export V3 Node
Edit `simple_syrup/nodes_v3/__init__.py`.
Status: completed.
- Import `BatchSEGSV3` inside `get_nodes()`.
- Include `BatchSEGSV3` in both returned node lists:
- prompt-control unavailable branch
- prompt-control available branch
- Keep ordering sensible. Put it near `TileAndTagSEGSV3` because both are SEGS workflow utilities.
Do not edit `simple_syrup/nodes/__init__.py` for this v3-only node.
Do not edit root `__init__.py` unless tests prove the v3 entrypoint needs changes. It already delegates to `simple_syrup.nodes_v3.get_nodes()`.
## Batch SEGS Tests
### Domain Tests
Add tests to `tests/test_segs_domain.py`.
Status: completed.
Required cases:
1. Batches multiple SEGS payloads in order.
- Input one has labels `1`, `2`, `3`.
- Input two has labels `4`, `5`, `6`.
- Output labels are `1`, `2`, `3`, `4`, `5`, `6`.
2. Accepts empty SEGS inputs.
- Input one has labels `1`, `2`.
- Input two is empty.
- Input three has label `3`.
- Output labels are `1`, `2`, `3`.
3. Returns empty output when all inputs are empty.
- Same shared header.
- Output segments tuple is empty.
4. Rejects no inputs.
- `batch_segs(())` raises `ValueError`.
5. Rejects mismatched headers.
- First input header `(8, 16)`.
- Second input header `(16, 8)`.
- Error message mentions the mismatched input index and both sizes.
Use existing `_segment(...)` helper in the test file where possible.
### V3 Node Tests
Create `tests/test_batch_segs_v3_node.py`.
Status: completed.
Required cases:
1. Schema test:
- `node_id == "SimpleSyrup.BatchSEGS"`
- `display_name == "Batch SEGS"`
- category is `SimpleSyrup/Detection`
- one input with id `segs_inputs`
- output id is `segs`
- output type is `SEGS`
2. Execute test:
- Pass a dict like:
```python
{"segs0": first, "segs1": second}
```
- Assert returned labels preserve dict insertion order.
- Assert return shape is Impact-compatible, meaning segments are a list.
3. Mismatched header test:
- Call `BatchSEGSV3.execute(...)`.
- Assert the domain error is surfaced.
### Registration Tests
Update `tests/test_registration.py`.
Status: completed.
In both v3 entrypoint tests, update expected node class names to include:
```text
BatchSEGSV3
```
Expected ordering should match `get_nodes()`.
No legacy mapping registration test is needed because this node is intentionally v3-only.
## WD14 SEGS Tagging Node
Status: completed after scope correction. The original document treated this as a follow-on, but the requested workflow included both **Batch SEGS** and **Tag SEGS w/ WD14**.
Implemented node name:
```text
Tag SEGS w/ WD14
```
Inputs:
- `image`: `IMAGE`
- `segs`: `SEGS`
- `clip`: `CLIP`
- `wd14_tagger`: `WD14_TAGGER`
- `universal_positive`: `STRING`
- `threshold`: `FLOAT`
- `character_threshold`: `FLOAT`
- `replace_underscore`: `BOOLEAN`
- `trailing_comma`: `BOOLEAN`
- `exclude_tags`: `STRING`
Outputs:
- `segs`: `SEGS`
- `positive`: `CONDITIONING_BATCH`
Implemented behavior:
- Validate the input image is a single image.
- Coerce incoming `SEGS`.
- Validate the `SEGS` header matches image dimensions.
- Crop each segment using `segment.crop_region`.
- Run WD14 over the crops in segment order.
- Prefix each tag prompt with the quality/universal positive prompt.
- Encode the resulting prompts using `ComfyConditioningEncoder.encode_batch`.
- Return the original SEGS unchanged plus the aligned `CONDITIONING_BATCH`.
- Reject mismatched tag count or conditioning count.
Implementation files:
- `simple_syrup/services/tag_segs_with_wd14_service.py`
- `simple_syrup/nodes/tag_segs_with_wd14.py`
- `simple_syrup/nodes_v3/tag_segs_with_wd14.py`
- `tests/test_tag_segs_with_wd14_service.py`
- `tests/test_tag_segs_with_wd14_node.py`
- `tests/test_tag_segs_with_wd14_v3_node.py`
## Batch Region Conditioning Node
Status: completed after the workflow requirement expanded to mixing autotagged SEGS conditioning with hand-authored regional conditioning. A v3 Autogrow node and a legacy two-input chainable node are both provided.
Implemented node name:
```text
Batch Region Conditioning
```
Inputs:
- Expandable v3 Autogrow inputs named `conditioning0`, `conditioning1`, etc.
- Legacy inputs named `first` and `second` for Comfy sessions that show legacy mapping nodes.
- Each input accepts either `CONDITIONING` or `CONDITIONING_BATCH`.
- Minimum visible/required inputs: `2`.
- Maximum inputs: `50`.
Output:
- `batch`: `CONDITIONING_BATCH`
Implemented behavior:
- Preserve socket order.
- If an input is `CONDITIONING_BATCH`, append all entries in that batch.
- If an input is normal `CONDITIONING`, append it as one entry.
- Return one flattened `CONDITIONING_BATCH`.
Example:
```text
auto_positive: [auto1, auto2]
hand_positive: hand1
extra_positive: [auto3]
↓
Batch Region Conditioning
↓
[auto1, auto2, hand1, auto3]
```
Implementation files:
- `simple_syrup/domain/conditioning_batch.py`
- `simple_syrup/nodes/batch_region_conditioning.py`
- `simple_syrup/nodes_v3/batch_region_conditioning.py`
- `tests/test_conditioning_batch.py`
- `tests/test_batch_region_conditioning_node.py`
- `tests/test_batch_region_conditioning_v3_node.py`
## Verification Commands
Run all commands from repository root:
```powershell
..\..\venv\Scripts\python.exe -m pytest -n auto -q tests\test_segs_domain.py tests\test_batch_segs_v3_node.py tests\test_registration.py
..\..\venv\Scripts\ruff.exe format .
..\..\venv\Scripts\ruff.exe check .
..\..\venv\Scripts\mypy.exe --strict simple_syrup tests
..\..\venv\Scripts\python.exe -m pytest -n auto -q
```
If a required tool is missing from `..\..\venv`, install or update development dependencies in that environment. Do not use global Python or a repository-local virtual environment.
## Definition of Done
- [x] `batch_segs` domain behavior exists and is tested.
- [x] **Batch SEGS** v3 node exists and uses `io.Autogrow`.
- [x] **Batch SEGS** is exported through `get_nodes()`.
- [x] Tests prove order-preserving flattening.
- [x] Tests prove mismatched headers fail clearly.
- [x] Tests prove v3 entrypoint includes the node.
- [x] Legacy node mappings are not changed for this v3-only node.
- [x] Required gates pass in the ComfyUI virtual environment.
- [x] **Tag SEGS w/ WD14** service, legacy node, and v3 node exist and are tested.
- [x] **Tag SEGS w/ WD14** is exported through legacy mappings and v3 entrypoint.
- [x] **Tag SEGS w/ WD14** emits `CONDITIONING_BATCH` from `CLIP` plus WD14 tags aligned to SEGS order.
- [x] **Batch Region Conditioning** v3 node accepts mixed `CONDITIONING` and `CONDITIONING_BATCH` inputs.
- [x] **Batch SEGS** and **Batch Region Conditioning** appear through legacy mappings for normal Comfy menu visibility.
+5
View File
@@ -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",
]
+17
View File
@@ -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."""
+33
View File
@@ -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."""
+12
View File
@@ -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",
@@ -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)),)
+44
View File
@@ -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),)
+172
View File
@@ -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.")
+28
View File
@@ -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."
)
+9
View File
@@ -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,
@@ -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())),)
+71
View File
@@ -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),)
+136
View File
@@ -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,
)
@@ -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}."
)
@@ -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"))
@@ -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"))
+53
View File
@@ -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,
)
+89
View File
@@ -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,
)
+20
View File
@@ -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."""
+6
View File
@@ -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,
+42
View File
@@ -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",
+68
View File
@@ -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."""
+112
View File
@@ -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
+283
View File
@@ -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 []
+103
View File
@@ -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
+2 -1
View File
@@ -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")