feat(detailing): add external llm segs tagging

This commit is contained in:
Artificial Sweetener
2026-05-31 16:07:50 -04:00
parent 4c7f087ce5
commit acd668f7c6
11 changed files with 1334 additions and 5 deletions
+32
View File
@@ -78,6 +78,38 @@ EXTERNAL_LLM_IMAGE_INPUT = (
)
EXTERNAL_LLM_RESPONSE_OUTPUT = "Assistant response returned by the external model."
EXTERNAL_LLM_TAG_SEGS_IMAGE = "Source image that the incoming SEGS were detected from."
EXTERNAL_LLM_TAG_SEGS_SEGS = (
"Existing SEGS to describe and keep aligned with the conditioning batch."
)
EXTERNAL_LLM_TAG_SEGS_CLIP = "CLIP model used to encode each generated regional prompt."
EXTERNAL_LLM_TAG_SEGS_MODEL = "External vision model used to describe each SEG crop."
EXTERNAL_LLM_TAG_SEGS_SYSTEM_PROMPT = (
"Instruction text that controls how the model writes regional tags."
)
EXTERNAL_LLM_TAG_SEGS_USER_PROMPT = "Per-region request sent with each SEG crop image."
EXTERNAL_LLM_TAG_SEGS_UNIVERSAL_POSITIVE = (
"Positive prompt text added before every generated regional prompt."
)
EXTERNAL_LLM_TAG_SEGS_IMAGE_MODE = (
"How pixels outside each SEG mask are shown to the vision model."
)
EXTERNAL_LLM_TAG_SEGS_REPLACE_UNDERSCORE = (
"Replace underscores in generated tags before CLIP encoding."
)
EXTERNAL_LLM_TAG_SEGS_TRAILING_COMMA = (
"Add a final comma to each generated regional prompt."
)
EXTERNAL_LLM_TAG_SEGS_EXCLUDE_TAGS = (
"Comma-separated exact tags removed from generated regional prompts."
)
EXTERNAL_LLM_TAG_SEGS_SEGS_OUTPUT = (
"Original SEGS returned in the same order as the conditioning batch."
)
EXTERNAL_LLM_TAG_SEGS_POSITIVE_OUTPUT = (
"Positive conditioning from external LLM tags, matched to SEGS order."
)
SAMPLING_MODEL = "Diffusion model used to denoise the input latent."
SAMPLING_SEED = (
"Seed used to create sampling noise. Reusing it with matching settings makes "
+2
View File
@@ -43,6 +43,7 @@ def get_nodes() -> list[type[object]]:
)
from .scale_factor import ScaleFactorV3
from .simple_load_checkpoint import SimpleLoadCheckpointV3
from .tag_segs_with_external_llm import TagSEGSWithExternalLLMV3
from .tag_segs_with_wd14 import TagSEGSWithWD14V3
from .tile_and_tag_segs import TileAndTagSEGSV3
from .vae_decode_options import VAEDecodeOptionsV3
@@ -77,6 +78,7 @@ def get_nodes() -> list[type[object]]:
SimpleLoadAnimaV3,
SimpleLoadCheckpointV3,
SimpleVAEEncodeV3,
TagSEGSWithExternalLLMV3,
TagSEGSWithWD14V3,
TileAndTagSEGSV3,
UpscaleLatentFromImageV3,
@@ -0,0 +1,193 @@
# 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 for tagging existing SEGS with an external LLM."""
from __future__ import annotations
from importlib import import_module
from typing import TYPE_CHECKING, Any, ClassVar
from ..domain.external_llm import (
DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
EXTERNAL_LLM_REASONING_EFFORTS,
)
from ..nodes import tooltips
from ..runtime.external_llm_images import SEG_IMAGE_MODES
from ..services.tag_segs_with_external_llm_service import (
LLMTagFormattingControls,
TagSEGSWithExternalLLMService,
)
MAX_EXTERNAL_LLM_MAX_TOKENS = 32768
DEFAULT_SEG_IMAGE_MODE = "transparent mask"
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 TagSEGSWithExternalLLMV3(_ComfyNodeBase):
"""Expose external-LLM SEGS tagging through Comfy's v3 API."""
_service = TagSEGSWithExternalLLMService()
@classmethod
def define_schema(cls) -> Any:
"""Declare the Tag SEGS w/ External LLM v3 schema."""
model_choices = cls._service.model_choices()
return _comfy_io.Schema(
node_id="SimpleSyrup.TagSEGSWithExternalLLM",
display_name="Tag SEGS w/ External LLM",
category="SimpleSyrup/Detailing",
description=(
"Tags existing SEGS crops with a configured external vision LLM "
"and returns aligned conditioning for SEGS detailing."
),
search_aliases=[
"llm",
"vision",
"tag",
"segs",
"detail",
"regional",
"prompt",
],
inputs=[
_comfy_io.Image.Input(
"image",
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_IMAGE,
),
_comfy_io.SEGS.Input(
"segs",
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_SEGS,
),
_comfy_io.Clip.Input(
"clip",
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_CLIP,
),
_comfy_io.Combo.Input(
"model",
options=model_choices,
default=model_choices[0],
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_MODEL,
),
_comfy_io.String.Input(
"system_prompt",
multiline=False,
default="",
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_SYSTEM_PROMPT,
),
_comfy_io.String.Input(
"user_prompt",
multiline=False,
default="",
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_USER_PROMPT,
),
_comfy_io.String.Input(
"universal_positive",
multiline=False,
default="",
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_UNIVERSAL_POSITIVE,
),
_comfy_io.Combo.Input(
"seg_image_mode",
options=list(SEG_IMAGE_MODES),
default=DEFAULT_SEG_IMAGE_MODE,
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_IMAGE_MODE,
),
_comfy_io.Boolean.Input(
"replace_underscore",
default=True,
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_REPLACE_UNDERSCORE,
),
_comfy_io.Boolean.Input(
"trailing_comma",
default=False,
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_TRAILING_COMMA,
),
_comfy_io.String.Input(
"exclude_tags",
multiline=False,
default="",
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_EXCLUDE_TAGS,
),
_comfy_io.Int.Input(
"max_tokens",
default=DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
min=1,
max=MAX_EXTERNAL_LLM_MAX_TOKENS,
step=1,
tooltip=tooltips.EXTERNAL_LLM_MAX_TOKENS_INPUT,
),
_comfy_io.Combo.Input(
"reasoning_effort",
options=list(EXTERNAL_LLM_REASONING_EFFORTS),
default=DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
tooltip=tooltips.EXTERNAL_LLM_REASONING_EFFORT_INPUT,
),
],
outputs=[
_comfy_io.SEGS.Output(
"segs",
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_SEGS_OUTPUT,
),
ConditioningBatchIO.Output(
"positive",
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_POSITIVE_OUTPUT,
),
],
)
@classmethod
def execute(
cls,
image: object,
segs: object,
clip: Any,
model: str,
system_prompt: str,
user_prompt: str,
universal_positive: str,
seg_image_mode: str,
replace_underscore: bool,
trailing_comma: bool,
exclude_tags: str,
max_tokens: int = DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
reasoning_effort: str = DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
) -> tuple[object, object]:
"""Tag existing SEGS and return aligned conditioning."""
result = cls._service.tag(
image=image,
segs=segs,
clip=clip,
model=model,
system_prompt=system_prompt,
user_prompt=user_prompt,
universal_positive=universal_positive,
seg_image_mode=seg_image_mode,
formatting=LLMTagFormattingControls(
replace_underscore=replace_underscore,
trailing_comma=trailing_comma,
exclude_tags=exclude_tags,
),
max_tokens=max_tokens,
reasoning_effort=reasoning_effort,
)
return result.segs, result.positive
+103
View File
@@ -12,8 +12,12 @@ from io import BytesIO
import torch
from PIL import Image
from ..domain.segs import Segment
from ..masking.segs_mask_ops import crop_image, crop_mask, resize_mask
from ..shared.tensor_validation import validate_image_tensor
SEG_IMAGE_MODES = ("transparent mask", "black mask", "full crop")
class ExternalLLMImageEncoder:
"""Encode ComfyUI IMAGE tensors for OpenAI-compatible vision payloads."""
@@ -32,6 +36,45 @@ class ExternalLLMImageEncoder:
return f"data:image/png;base64,{encoded}"
class ExternalLLMSegsImageEncoder:
"""Encode SEG crops for OpenAI-compatible vision payloads."""
def encode_segment_as_data_url(
self,
image: torch.Tensor,
segment: Segment,
mode: str,
) -> str:
"""Return one SEG crop as a PNG data URL."""
if mode not in SEG_IMAGE_MODES:
choices = ", ".join(SEG_IMAGE_MODES)
raise ValueError(f"seg_image_mode must be one of: {choices}.")
validate_image_tensor(image)
if int(image.shape[0]) != 1:
raise ValueError("SEG crop image encoding requires one IMAGE item.")
crop = crop_image(image.float().clamp(0.0, 1.0), segment.crop_region)
if mode == "full crop":
return _pil_to_png_data_url(_image_tensor_to_rgb_pil(crop))
mask = _segment_mask_for_crop(
segment=segment,
image_height=int(image.shape[1]),
image_width=int(image.shape[2]),
)
if int(mask.shape[0]) != int(crop.shape[1]) or int(mask.shape[1]) != int(
crop.shape[2]
):
mask = resize_mask(mask, int(crop.shape[1]), int(crop.shape[2]))
if mode == "black mask":
masked_crop = crop * mask.unsqueeze(0).unsqueeze(-1)
return _pil_to_png_data_url(_image_tensor_to_rgb_pil(masked_crop))
return _pil_to_png_data_url(_crop_and_mask_to_rgba_pil(crop, mask))
def _image_tensor_to_rgb_pil(image: torch.Tensor) -> Image.Image:
"""Convert the first BHWC image tensor item to RGB PIL image."""
@@ -45,3 +88,63 @@ def _image_tensor_to_rgb_pil(image: torch.Tensor) -> Image.Image:
else:
array = np.repeat(array[..., :1], 3, axis=-1)
return Image.fromarray((array * 255.0).round().astype(np.uint8))
def _segment_mask_for_crop(
segment: Segment,
image_height: int,
image_width: int,
) -> torch.Tensor:
"""Return the segment mask normalized to the crop region."""
if not isinstance(segment.cropped_mask, torch.Tensor):
raise TypeError("SEG cropped_mask must be a torch.Tensor.")
mask = _normalize_mask_shape(segment.cropped_mask)
region = segment.crop_region
crop_height = region.height
crop_width = region.width
mask_height = int(mask.shape[0])
mask_width = int(mask.shape[1])
if mask_height == crop_height and mask_width == crop_width:
return mask.float().clamp(0.0, 1.0)
if mask_height == image_height and mask_width == image_width:
return crop_mask(mask, region).float().clamp(0.0, 1.0)
return resize_mask(mask, crop_height, crop_width).float().clamp(0.0, 1.0)
def _normalize_mask_shape(mask: torch.Tensor) -> torch.Tensor:
"""Return a SEG mask as an HW tensor."""
working = mask.detach().cpu().float()
if working.ndim == 2:
return working
if working.ndim == 3:
return working[0]
raise ValueError("SEG cropped_mask must be an HW or BHW tensor.")
def _crop_and_mask_to_rgba_pil(crop: torch.Tensor, mask: torch.Tensor) -> Image.Image:
"""Convert a BHWC crop and HW alpha mask to an RGBA PIL image."""
import numpy as np
rgb = crop[0].detach().cpu().float().clamp(0.0, 1.0).numpy()
if rgb.shape[-1] == 1:
rgb = np.repeat(rgb, 3, axis=-1)
elif rgb.shape[-1] >= 3:
rgb = rgb[..., :3]
else:
rgb = np.repeat(rgb[..., :1], 3, axis=-1)
alpha = mask.detach().cpu().float().clamp(0.0, 1.0).numpy()
rgba = np.concatenate((rgb, alpha[..., None]), axis=-1)
return Image.fromarray((rgba * 255.0).round().astype(np.uint8), mode="RGBA")
def _pil_to_png_data_url(image: Image.Image) -> str:
"""Return a PNG data URL for a PIL image."""
buffer = BytesIO()
image.save(buffer, format="PNG")
encoded = base64.b64encode(buffer.getvalue()).decode("ascii")
return f"data:image/png;base64,{encoded}"
@@ -228,6 +228,30 @@ class ExternalLLMPromptService:
) -> str:
"""Generate an assistant response for the supplied prompt pair."""
return self.generate_with_image_data_url(
model=model,
system_prompt=system_prompt,
user_prompt=user_prompt,
max_tokens=max_tokens,
reasoning_effort=reasoning_effort,
image_data_url=(
None
if image is None
else self._image_encoder.encode_first_image_as_data_url(image)
),
)
def generate_with_image_data_url(
self,
model: str,
system_prompt: str,
user_prompt: str,
max_tokens: int = DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
reasoning_effort: str = DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
image_data_url: str | None = None,
) -> str:
"""Generate an assistant response with a pre-encoded optional image."""
selected_model = self._resolve_model_for_execution(model)
external = self._settings_repository.load().external_llm
@@ -250,11 +274,7 @@ class ExternalLLMPromptService:
user_prompt=user_prompt,
max_tokens=max_tokens,
reasoning_effort=reasoning_effort,
image_data_url=(
None
if image is None
else self._image_encoder.encode_first_image_as_data_url(image)
),
image_data_url=image_data_url,
)
response = self._client.create_chat_completion(
external.base_url,
@@ -0,0 +1,295 @@
# 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 external-LLM 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.external_llm import (
DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
)
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 validate_single_image
from ..runtime.conditioning_encoding import ComfyConditioningEncoder
from ..runtime.external_llm_images import ExternalLLMSegsImageEncoder
from ..runtime.progress import ProgressReporter, create_comfy_progress
from ..shared.logging import get_logger
from .external_llm_prompt_service import ExternalLLMPromptService
LOGGER = get_logger(__name__)
OPERATION = "Tag SEGS w/ External LLM"
@dataclass(frozen=True)
class LLMTagFormattingControls:
"""Store formatting controls for external LLM tag responses."""
replace_underscore: bool = True
trailing_comma: bool = False
exclude_tags: str = ""
class ExternalLLMGenerationBoundary(Protocol):
"""External LLM provider execution boundary."""
def model_choices(self) -> list[str]:
"""Return cached provider model choices for Comfy dropdowns."""
def generate_with_image_data_url(
self,
model: str,
system_prompt: str,
user_prompt: str,
max_tokens: int = DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
reasoning_effort: str = DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
image_data_url: str | None = None,
) -> str:
"""Return one assistant response for a pre-encoded image."""
class SegmentImageEncodingBoundary(Protocol):
"""Encode ordered SEG crops as LLM vision images."""
def encode_segment_as_data_url(
self,
image: torch.Tensor,
segment: Segment,
mode: str,
) -> str:
"""Return one SEG crop image data URL."""
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 TagSEGSWithExternalLLMResult:
"""Return unchanged SEGS and aligned external-LLM conditioning."""
segs: ImpactSegs
positive: ConditioningBatch
class TagSEGSWithExternalLLMService:
"""Caption provided SEGS crops with an external LLM and encode conditioning."""
def __init__(
self,
llm: ExternalLLMGenerationBoundary | None = None,
image_encoder: SegmentImageEncodingBoundary | None = None,
conditioning_encoder: ConditioningEncodingBoundary | None = None,
progress_factory: Callable[[int], ProgressReporter] | None = None,
) -> None:
"""Create the service with injectable runtime boundaries."""
self._llm = llm or ExternalLLMPromptService()
self._image_encoder = image_encoder or ExternalLLMSegsImageEncoder()
self._conditioning_encoder = conditioning_encoder or ComfyConditioningEncoder()
self._progress_factory = progress_factory or create_comfy_progress
def model_choices(self) -> list[str]:
"""Return cached provider model choices for Comfy dropdowns."""
return self._llm.model_choices()
def tag(
self,
image: object,
segs: object,
clip: Any,
model: str,
system_prompt: str,
user_prompt: str,
universal_positive: str,
seg_image_mode: str,
formatting: LLMTagFormattingControls,
max_tokens: int = DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
reasoning_effort: str = DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
) -> TagSEGSWithExternalLLMResult:
"""Return original SEGS plus LLM-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/ External LLM.")
progress = self._progress_factory(len(segments) + 2)
progress.update(1)
prompts: list[str] = []
for segment in segments:
prompts.append(
self._prompt_for_segment(
image=image_tensor,
segment=segment,
model=model,
system_prompt=system_prompt,
user_prompt=user_prompt,
universal_positive=universal_positive,
seg_image_mode=seg_image_mode,
formatting=formatting,
max_tokens=max_tokens,
reasoning_effort=reasoning_effort,
)
)
progress.update(1)
positive = self._conditioning_encoder.encode_batch(clip, tuple(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/ External LLM pass completed",
extra={
"operation": "tag_segs_with_external_llm",
"segment_count": len(segments),
"external_llm_model": model,
"seg_image_mode": seg_image_mode,
"universal_positive_present": bool(universal_positive.strip()),
"replace_underscore": formatting.replace_underscore,
"trailing_comma": formatting.trailing_comma,
"exclude_tags_present": bool(formatting.exclude_tags.strip()),
},
)
return TagSEGSWithExternalLLMResult(
segs=to_impact_compatible_segs(native_segs),
positive=positive,
)
def _prompt_for_segment(
self,
image: torch.Tensor,
segment: Segment,
model: str,
system_prompt: str,
user_prompt: str,
universal_positive: str,
seg_image_mode: str,
formatting: LLMTagFormattingControls,
max_tokens: int,
reasoning_effort: str,
) -> str:
"""Generate, format, and prefix one segment prompt."""
image_data_url = self._image_encoder.encode_segment_as_data_url(
image,
segment,
seg_image_mode,
)
response = self._llm.generate_with_image_data_url(
model=model,
system_prompt=system_prompt,
user_prompt=user_prompt,
max_tokens=max_tokens,
reasoning_effort=reasoning_effort,
image_data_url=image_data_url,
)
prompt = format_external_llm_tags(response, formatting)
return prefix_prompt(universal_positive, prompt)
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 format_external_llm_tags(
response: str,
controls: LLMTagFormattingControls,
) -> str:
"""Format one external LLM response into comma-separated prompt tags."""
stripped = response.strip()
if not stripped:
raise ValueError("External LLM returned an empty response for a SEG.")
excluded = _excluded_tags(controls)
tags: list[str] = []
for chunk in stripped.split(","):
tag = _normalize_tag(chunk, controls.replace_underscore)
if not tag or tag.lower() in excluded:
continue
tags.append(tag)
if not tags:
raise ValueError("External LLM response had no usable tags after exclusions.")
prompt = ", ".join(tags)
if controls.trailing_comma and not prompt.endswith(","):
return f"{prompt},"
return prompt
def _excluded_tags(controls: LLMTagFormattingControls) -> set[str]:
"""Return normalized excluded tag names."""
return {
tag.lower()
for raw_tag in controls.exclude_tags.split(",")
if (tag := _normalize_tag(raw_tag, controls.replace_underscore))
}
def _normalize_tag(value: str, replace_underscore: bool) -> str:
"""Return the prompt-facing form of one tag-like response chunk."""
tag = value.strip()
if replace_underscore:
return tag.replace("_", " ")
return tag
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}."
)
+18
View File
@@ -401,6 +401,24 @@ def test_generate_attaches_encoded_image_when_supplied() -> None:
assert client.requests[0].image_data_url == "data:image/png;base64,abc"
def test_generate_with_image_data_url_forwards_preencoded_image() -> None:
"""SEG callers can supply a prebuilt image data URL."""
client = FakeClient()
service = configured_service(client=client)
assert (
service.generate_with_image_data_url(
model="model-a",
system_prompt="system",
user_prompt="user",
image_data_url="data:image/png;base64,seg",
)
== "assistant response"
)
assert client.requests[0].image_data_url == "data:image/png;base64,seg"
def configured_service(
client: FakeClient | None = None,
image_encoder: FakeImageEncoder | None = None,
+160
View File
@@ -0,0 +1,160 @@
# 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 external LLM SEG crop image encoding."""
from __future__ import annotations
import base64
from io import BytesIO
from typing import cast
import pytest
import torch
from PIL import Image
from simple_syrup.domain.segs import BoundingBox, CropRegion, Segment
from simple_syrup.runtime.external_llm_images import ExternalLLMSegsImageEncoder
def test_segs_image_encoder_returns_transparent_mask_png() -> None:
"""Transparent mode hides outside-mask pixels with PNG alpha."""
segment = _segment(torch.tensor([[1.0, 0.0], [0.0, 1.0]]))
encoded = ExternalLLMSegsImageEncoder().encode_segment_as_data_url(
_image(),
segment,
"transparent mask",
)
image = _decode_png(encoded)
assert image.mode == "RGBA"
assert image.size == (2, 2)
assert _rgba_pixel(image, 0, 0)[3] == 255
assert _rgba_pixel(image, 1, 0)[3] == 0
assert _rgba_pixel(image, 0, 1)[3] == 0
assert _rgba_pixel(image, 1, 1)[3] == 255
def test_segs_image_encoder_returns_black_mask_png() -> None:
"""Black mode zeros outside-mask pixels and keeps RGB output."""
segment = _segment(torch.tensor([[1.0, 0.0], [0.0, 1.0]]))
encoded = ExternalLLMSegsImageEncoder().encode_segment_as_data_url(
_image(),
segment,
"black mask",
)
image = _decode_png(encoded)
assert image.mode == "RGB"
assert image.size == (2, 2)
assert image.getpixel((1, 0)) == (0, 0, 0)
assert image.getpixel((0, 1)) == (0, 0, 0)
assert image.getpixel((0, 0)) != (0, 0, 0)
assert image.getpixel((1, 1)) != (0, 0, 0)
def test_segs_image_encoder_returns_full_crop_png() -> None:
"""Full crop mode preserves the whole crop rectangle."""
segment = _segment(torch.zeros((2, 2)))
encoded = ExternalLLMSegsImageEncoder().encode_segment_as_data_url(
_image(),
segment,
"full crop",
)
image = _decode_png(encoded)
assert image.mode == "RGB"
assert image.size == (2, 2)
assert image.getpixel((0, 0)) != (0, 0, 0)
assert image.getpixel((1, 0)) != (0, 0, 0)
assert image.getpixel((0, 1)) != (0, 0, 0)
assert image.getpixel((1, 1)) != (0, 0, 0)
def test_segs_image_encoder_accepts_full_image_masks() -> None:
"""Full-image SEG masks are cropped to the SEG crop region."""
full_mask = torch.zeros((4, 4), dtype=torch.float32)
full_mask[1, 1] = 1.0
full_mask[2, 2] = 1.0
segment = _segment(full_mask)
encoded = ExternalLLMSegsImageEncoder().encode_segment_as_data_url(
_image(),
segment,
"transparent mask",
)
image = _decode_png(encoded)
assert [_rgba_pixel(image, x, y)[3] for y in range(2) for x in range(2)] == [
255,
0,
0,
255,
]
def test_segs_image_encoder_rejects_invalid_masks() -> None:
"""SEG masks must be HW or BHW tensors."""
segment = _segment(torch.zeros((1, 1, 1, 1), dtype=torch.float32))
with pytest.raises(ValueError, match="HW or BHW"):
ExternalLLMSegsImageEncoder().encode_segment_as_data_url(
_image(),
segment,
"transparent mask",
)
def test_segs_image_encoder_rejects_unknown_mode() -> None:
"""SEG image mode is restricted to the node combo choices."""
with pytest.raises(ValueError, match="seg_image_mode"):
ExternalLLMSegsImageEncoder().encode_segment_as_data_url(
_image(),
_segment(torch.ones((2, 2), dtype=torch.float32)),
"white mask",
)
def _decode_png(data_url: str) -> Image.Image:
"""Decode a PNG data URL into a PIL image."""
assert data_url.startswith("data:image/png;base64,")
payload = data_url.removeprefix("data:image/png;base64,")
return Image.open(BytesIO(base64.b64decode(payload)))
def _rgba_pixel(image: Image.Image, x: int, y: int) -> tuple[int, int, int, int]:
"""Return one RGBA pixel with a precise test type."""
return cast(tuple[int, int, int, int], image.getpixel((x, y)))
def _segment(mask: torch.Tensor) -> Segment:
"""Return a SEG that crops a stable two-by-two image region."""
return Segment(
cropped_image=None,
cropped_mask=mask,
confidence=1.0,
crop_region=CropRegion(1, 1, 3, 3),
bbox=BoundingBox(1, 1, 3, 3),
label="seg",
)
def _image() -> torch.Tensor:
"""Return a deterministic BHWC image with no black crop pixels."""
return (
torch.arange(1, 4 * 4 * 3 + 1, dtype=torch.float32).reshape(1, 4, 4, 3) / 255.0
)
+1
View File
@@ -44,6 +44,7 @@ BASE_NODE_IDS = [
"SimpleSyrup.SimpleLoadAnima",
"SimpleSyrup.SimpleLoadCheckpoint",
"SimpleSyrup.SimpleVAEEncode",
"SimpleSyrup.TagSEGSWithExternalLLM",
"SimpleSyrup.TagSEGSWithWD14",
"SimpleSyrup.TileAndTagSEGS",
"SimpleSyrup.UpscaleLatentFromImage",
@@ -0,0 +1,360 @@
# 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 external-LLM tagging of existing SEGS."""
from __future__ import annotations
import logging
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.services.tag_segs_with_external_llm_service import (
LLMTagFormattingControls,
TagSEGSWithExternalLLMService,
)
DEFAULT_CROP_REGION = CropRegion(1, 1, 3, 3)
def test_service_preserves_segs_llm_prompt_and_conditioning_order(
caplog: pytest.LogCaptureFixture,
) -> None:
"""Existing SEGS, LLM calls, formatted prompts, and conditioning stay aligned."""
progress = _ProgressRecorder()
llm = _FakeLLM(("blue_hair, smile, bad_tag", "green_eyes"))
image_encoder = _FakeImageEncoder()
conditioning_encoder = _FakeConditioningEncoder()
service = TagSEGSWithExternalLLMService(
llm=llm,
image_encoder=image_encoder,
conditioning_encoder=conditioning_encoder,
progress_factory=lambda _total: progress,
)
segs = _native_segs(("first", "second"))
with caplog.at_level(logging.INFO, logger="simple_syrup"):
result = service.tag(
image=_image(),
segs=segs,
clip="clip",
model="vision-model",
system_prompt="system",
user_prompt="user",
universal_positive="masterpiece",
seg_image_mode="transparent mask",
formatting=LLMTagFormattingControls(
replace_underscore=True,
trailing_comma=True,
exclude_tags="bad tag",
),
max_tokens=64,
reasoning_effort="off",
)
assert [segment.label for segment in result.segs[1]] == ["first", "second"]
assert image_encoder.calls == (
("first", "transparent mask"),
("second", "transparent mask"),
)
assert [call["image_data_url"] for call in llm.calls] == [
"data:image/png;base64,first",
"data:image/png;base64,second",
]
assert [call["model"] for call in llm.calls] == ["vision-model", "vision-model"]
assert [call["max_tokens"] for call in llm.calls] == [64, 64]
assert [call["reasoning_effort"] for call in llm.calls] == ["off", "off"]
assert conditioning_encoder.chunks == (
"masterpiece, blue hair, smile,",
"masterpiece, green eyes,",
)
assert result.positive.entries == (
"clip:masterpiece, blue hair, smile,",
"clip:masterpiece, green eyes,",
)
assert progress.updates == [1, 1, 1, 1]
record = caplog.records[-1]
assert record.__dict__["operation"] == "tag_segs_with_external_llm"
assert record.__dict__["segment_count"] == 2
assert record.__dict__["external_llm_model"] == "vision-model"
assert record.__dict__["seg_image_mode"] == "transparent mask"
assert record.__dict__["universal_positive_present"] is True
def test_service_preserves_underscores_when_requested() -> None:
"""Formatting can keep booru-style underscores."""
conditioning_encoder = _FakeConditioningEncoder()
service = _service(
llm=_FakeLLM(("blue_hair, smile",)),
conditioning_encoder=conditioning_encoder,
)
service.tag(
image=_image(),
segs=_native_segs(("first",)),
clip="clip",
model="vision-model",
system_prompt="system",
user_prompt="user",
universal_positive="",
seg_image_mode="black mask",
formatting=LLMTagFormattingControls(replace_underscore=False),
)
assert conditioning_encoder.chunks == ("blue_hair, smile",)
def test_service_rejects_empty_segs() -> None:
"""External LLM tagging needs at least one SEG to tag."""
with pytest.raises(ValueError, match="No SEGS"):
_service().tag(
image=_image(),
segs=((4, 4), ()),
clip="clip",
model="vision-model",
system_prompt="system",
user_prompt="user",
universal_positive="",
seg_image_mode="transparent mask",
formatting=LLMTagFormattingControls(),
)
def test_service_rejects_segs_image_header_mismatch() -> None:
"""SEGS must describe the source image dimensions."""
with pytest.raises(ValueError, match="SEGS is 8x4, image is 4x4"):
_service().tag(
image=_image(),
segs=((8, 4), _native_segs(("first",))[1]),
clip="clip",
model="vision-model",
system_prompt="system",
user_prompt="user",
universal_positive="",
seg_image_mode="transparent mask",
formatting=LLMTagFormattingControls(),
)
def test_service_rejects_crop_regions_outside_image() -> None:
"""SEG crop regions must fit inside the connected source image."""
with pytest.raises(ValueError, match="crop_region must fit"):
_service().tag(
image=_image(),
segs=_native_segs(("outside",), crop_region=CropRegion(3, 3, 5, 5)),
clip="clip",
model="vision-model",
system_prompt="system",
user_prompt="user",
universal_positive="",
seg_image_mode="transparent mask",
formatting=LLMTagFormattingControls(),
)
def test_service_rejects_empty_llm_responses() -> None:
"""Empty provider responses cannot produce usable regional prompts."""
with pytest.raises(ValueError, match="empty response"):
_service(llm=_FakeLLM((" ",))).tag(
image=_image(),
segs=_native_segs(("first",)),
clip="clip",
model="vision-model",
system_prompt="system",
user_prompt="user",
universal_positive="",
seg_image_mode="transparent mask",
formatting=LLMTagFormattingControls(),
)
def test_service_rejects_responses_removed_by_exclusions() -> None:
"""Exclusion filtering must not silently create blank prompts."""
with pytest.raises(ValueError, match="no usable tags"):
_service(llm=_FakeLLM(("bad_tag",))).tag(
image=_image(),
segs=_native_segs(("first",)),
clip="clip",
model="vision-model",
system_prompt="system",
user_prompt="user",
universal_positive="",
seg_image_mode="transparent mask",
formatting=LLMTagFormattingControls(exclude_tags="bad tag"),
)
def test_service_rejects_conditioning_count_mismatch() -> None:
"""Conditioning output count must stay aligned to SEGS count."""
with pytest.raises(ValueError, match="returned 1 entries for 2 SEGS"):
_service(
llm=_FakeLLM(("first", "second")),
conditioning_encoder=_ShortConditioningEncoder(),
).tag(
image=_image(),
segs=_native_segs(("first", "second")),
clip="clip",
model="vision-model",
system_prompt="system",
user_prompt="user",
universal_positive="",
seg_image_mode="transparent mask",
formatting=LLMTagFormattingControls(),
)
class _FakeLLM:
"""Return ordered external LLM responses."""
def __init__(self, responses: tuple[str, ...] = ("tag",)) -> None:
"""Store fixed responses."""
self.responses = responses
self.calls: list[dict[str, object]] = []
def model_choices(self) -> list[str]:
"""Return deterministic model choices."""
return ["vision-model"]
def generate_with_image_data_url(
self,
model: str,
system_prompt: str,
user_prompt: str,
max_tokens: int = 1024,
reasoning_effort: str = "default",
image_data_url: str | None = None,
) -> str:
"""Capture one LLM call and return its configured response."""
self.calls.append(
{
"model": model,
"system_prompt": system_prompt,
"user_prompt": user_prompt,
"max_tokens": max_tokens,
"reasoning_effort": reasoning_effort,
"image_data_url": image_data_url,
}
)
return self.responses[len(self.calls) - 1]
class _FakeImageEncoder:
"""Return visible data URLs for SEG crops."""
def __init__(self) -> None:
"""Initialize captured calls."""
self.calls: tuple[tuple[str, str], ...] = ()
def encode_segment_as_data_url(
self,
image: torch.Tensor,
segment: Segment,
mode: str,
) -> str:
"""Record one SEG image encoding call."""
assert tuple(image.shape) == (1, 4, 4, 3)
self.calls = (*self.calls, (segment.label, mode))
return f"data:image/png;base64,{segment.label}"
class _FakeConditioningEncoder:
"""Return visible conditioning entries 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 _ShortConditioningEncoder:
"""Return too few conditioning entries."""
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 _service(
llm: _FakeLLM | None = None,
conditioning_encoder: _FakeConditioningEncoder
| _ShortConditioningEncoder
| None = None,
) -> TagSEGSWithExternalLLMService:
"""Create a service with fake boundaries."""
return TagSEGSWithExternalLLMService(
llm=llm or _FakeLLM(),
image_encoder=_FakeImageEncoder(),
conditioning_encoder=conditioning_encoder or _FakeConditioningEncoder(),
)
def _native_segs(
labels: tuple[str, ...],
crop_region: CropRegion = DEFAULT_CROP_REGION,
) -> NativeSegs:
"""Create native SEGS with stable crop regions."""
segments = tuple(
Segment(
cropped_image=None,
cropped_mask=torch.ones((crop_region.height, crop_region.width)),
confidence=1.0,
crop_region=crop_region,
bbox=BoundingBox(
crop_region.left,
crop_region.top,
crop_region.right,
crop_region.bottom,
),
label=label,
)
for label in 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
@@ -0,0 +1,145 @@
# 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/ External LLM Comfy v3 node."""
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_v3.tag_segs_with_external_llm import (
TagSEGSWithExternalLLMV3,
)
from simple_syrup.services.tag_segs_with_external_llm_service import (
LLMTagFormattingControls,
TagSEGSWithExternalLLMResult,
)
def test_tag_segs_with_external_llm_v3_schema(monkeypatch: Any) -> None:
"""The v3 schema exposes the external LLM SEGS tagging contract."""
monkeypatch.setattr(TagSEGSWithExternalLLMV3, "_service", _FakeService())
schema = TagSEGSWithExternalLLMV3.define_schema()
assert schema.node_id == "SimpleSyrup.TagSEGSWithExternalLLM"
assert schema.display_name == "Tag SEGS w/ External LLM"
assert schema.category == "SimpleSyrup/Detailing"
assert [input_item.id for input_item in schema.inputs] == [
"image",
"segs",
"clip",
"model",
"system_prompt",
"user_prompt",
"universal_positive",
"seg_image_mode",
"replace_underscore",
"trailing_comma",
"exclude_tags",
"max_tokens",
"reasoning_effort",
]
assert schema.inputs[0].io_type == "IMAGE"
assert schema.inputs[1].io_type == "SEGS"
assert schema.inputs[2].io_type == "CLIP"
assert schema.inputs[3].options == ["vision-model"]
assert schema.inputs[4].default == ""
assert schema.inputs[5].default == ""
seg_image_mode = schema.inputs[7]
assert seg_image_mode.io_type == "COMBO"
assert seg_image_mode.options == ["transparent mask", "black mask", "full crop"]
assert seg_image_mode.default == "transparent mask"
assert schema.inputs[8].default is True
assert schema.inputs[9].default is False
assert schema.inputs[10].default == ""
assert schema.inputs[11].default == 1024
assert schema.inputs[12].options == ["default", "high", "medium", "low", "off"]
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_external_llm_v3_execute_forwards_to_service(
monkeypatch: Any,
) -> None:
"""The v3 node delegates execution to the service."""
service = _FakeService()
monkeypatch.setattr(TagSEGSWithExternalLLMV3, "_service", service)
image = torch.zeros((1, 8, 8, 3))
segs: ImpactSegs = ((8, 8), [])
output_segs, positive = TagSEGSWithExternalLLMV3.execute(
image=image,
segs=segs,
clip="clip",
model="vision-model",
system_prompt="system",
user_prompt="user",
universal_positive="masterpiece",
seg_image_mode="black mask",
replace_underscore=False,
trailing_comma=True,
exclude_tags="bad tag",
max_tokens=128,
reasoning_effort="off",
)
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["model"] == "vision-model"
assert service.call["system_prompt"] == "system"
assert service.call["user_prompt"] == "user"
assert service.call["universal_positive"] == "masterpiece"
assert service.call["seg_image_mode"] == "black mask"
assert service.call["max_tokens"] == 128
assert service.call["reasoning_effort"] == "off"
formatting = service.call["formatting"]
assert isinstance(formatting, LLMTagFormattingControls)
assert formatting.replace_underscore is False
assert formatting.trailing_comma is True
assert formatting.exclude_tags == "bad tag"
class _FakeService:
"""Capture v3 node service calls."""
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 = TagSEGSWithExternalLLMResult(
segs=((8, 8), [segment]),
positive=ConditioningBatch(("encoded",)),
)
self.call: dict[str, object] = {}
def model_choices(self) -> list[str]:
"""Return deterministic model choices."""
return ["vision-model"]
def tag(self, **kwargs: object) -> TagSEGSWithExternalLLMResult:
"""Return a fixed result and remember provided inputs."""
self.call = kwargs
return self.result