feat(detailing): add external llm segs tagging
This commit is contained in:
@@ -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 "
|
||||
|
||||
@@ -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
|
||||
@@ -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}."
|
||||
)
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
)
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user