diff --git a/simple_syrup/nodes/tooltips.py b/simple_syrup/nodes/tooltips.py index f9505a5..3d970f3 100644 --- a/simple_syrup/nodes/tooltips.py +++ b/simple_syrup/nodes/tooltips.py @@ -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 " diff --git a/simple_syrup/nodes_v3/__init__.py b/simple_syrup/nodes_v3/__init__.py index 9983c25..10b1de1 100644 --- a/simple_syrup/nodes_v3/__init__.py +++ b/simple_syrup/nodes_v3/__init__.py @@ -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, diff --git a/simple_syrup/nodes_v3/tag_segs_with_external_llm.py b/simple_syrup/nodes_v3/tag_segs_with_external_llm.py new file mode 100644 index 0000000..0efc6c8 --- /dev/null +++ b/simple_syrup/nodes_v3/tag_segs_with_external_llm.py @@ -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 diff --git a/simple_syrup/runtime/external_llm_images.py b/simple_syrup/runtime/external_llm_images.py index 6a4c660..797aae9 100644 --- a/simple_syrup/runtime/external_llm_images.py +++ b/simple_syrup/runtime/external_llm_images.py @@ -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}" diff --git a/simple_syrup/services/external_llm_prompt_service.py b/simple_syrup/services/external_llm_prompt_service.py index 41cb8fa..75e82c1 100644 --- a/simple_syrup/services/external_llm_prompt_service.py +++ b/simple_syrup/services/external_llm_prompt_service.py @@ -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, diff --git a/simple_syrup/services/tag_segs_with_external_llm_service.py b/simple_syrup/services/tag_segs_with_external_llm_service.py new file mode 100644 index 0000000..2db45df --- /dev/null +++ b/simple_syrup/services/tag_segs_with_external_llm_service.py @@ -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}." + ) diff --git a/tests/test_external_llm_prompt_service.py b/tests/test_external_llm_prompt_service.py index f3ea272..83e8e03 100644 --- a/tests/test_external_llm_prompt_service.py +++ b/tests/test_external_llm_prompt_service.py @@ -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, diff --git a/tests/test_external_llm_segs_images.py b/tests/test_external_llm_segs_images.py new file mode 100644 index 0000000..5ceeb7d --- /dev/null +++ b/tests/test_external_llm_segs_images.py @@ -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 + ) diff --git a/tests/test_registration.py b/tests/test_registration.py index a15eda0..9591e92 100644 --- a/tests/test_registration.py +++ b/tests/test_registration.py @@ -44,6 +44,7 @@ BASE_NODE_IDS = [ "SimpleSyrup.SimpleLoadAnima", "SimpleSyrup.SimpleLoadCheckpoint", "SimpleSyrup.SimpleVAEEncode", + "SimpleSyrup.TagSEGSWithExternalLLM", "SimpleSyrup.TagSEGSWithWD14", "SimpleSyrup.TileAndTagSEGS", "SimpleSyrup.UpscaleLatentFromImage", diff --git a/tests/test_tag_segs_with_external_llm_service.py b/tests/test_tag_segs_with_external_llm_service.py new file mode 100644 index 0000000..c76be1d --- /dev/null +++ b/tests/test_tag_segs_with_external_llm_service.py @@ -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 diff --git a/tests/test_tag_segs_with_external_llm_v3_node.py b/tests/test_tag_segs_with_external_llm_v3_node.py new file mode 100644 index 0000000..364c93b --- /dev/null +++ b/tests/test_tag_segs_with_external_llm_v3_node.py @@ -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