feat(segs): add regional batching and wd14 tagging nodes
This commit is contained in:
@@ -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.
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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)),)
|
||||
@@ -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),)
|
||||
@@ -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.")
|
||||
@@ -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."
|
||||
)
|
||||
|
||||
@@ -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())),)
|
||||
@@ -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),)
|
||||
@@ -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"))
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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 []
|
||||
@@ -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
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user