diff --git a/simple_syrup/domain/segs_tiled_diffusion.py b/simple_syrup/domain/segs_tiled_diffusion.py new file mode 100644 index 0000000..89e5a4a --- /dev/null +++ b/simple_syrup/domain/segs_tiled_diffusion.py @@ -0,0 +1,505 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Build irregular, SEGS-guided latent tiles for tiled diffusion sampling.""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch +import torch.nn.functional as functional + +from .segs import NativeSegs, Segment, coerce_segs +from .tiled_diffusion import ( + LatentTile, + TiledDiffusionPlan, + batch_latent_tiles, + build_tiled_diffusion_plan, +) + + +@dataclass(frozen=True) +class _OwnershipCore: + """Represent a non-overlapping latent ownership region before window placement.""" + + mask: torch.Tensor + bounds: tuple[int, int, int, int] + area: int + + +def build_segs_guided_tiled_diffusion_plan( + *, + segs: object, + latent_width: int, + latent_height: int, + tile_width: int, + tile_height: int, + overlap: int, + tile_batch_size: int, +) -> TiledDiffusionPlan: + """Build bounded sampling windows whose irregular cores follow supplied SEGS. + + Every latent pixel receives exactly one ownership core. Each core is sampled + through a rectangular window, while its local blend mask retains the irregular + boundary and shares a feathered overlap with neighboring cores. + """ + + base_plan = build_tiled_diffusion_plan( + latent_width=latent_width, + latent_height=latent_height, + tile_width=tile_width, + tile_height=tile_height, + overlap=overlap, + tile_batch_size=tile_batch_size, + ) + native_segs = coerce_segs(segs) + _validate_aspect_ratio(native_segs, latent_height, latent_width) + ownership = _build_ownership_cores( + native_segs, + latent_height=latent_height, + latent_width=latent_width, + ) + max_core_width = max(1, base_plan.tile_width - base_plan.overlap) + max_core_height = max(1, base_plan.tile_height - base_plan.overlap) + split_cores = tuple( + split_core + for core in ownership + for split_core in _split_core( + core, + max_width=max_core_width, + max_height=max_core_height, + ) + ) + merged_cores = _merge_small_cores( + split_cores, + max_width=max_core_width, + max_height=max_core_height, + ) + tiles = tuple( + sorted( + ( + _tile_for_core( + core, + latent_width=latent_width, + latent_height=latent_height, + tile_width=base_plan.tile_width, + tile_height=base_plan.tile_height, + overlap=base_plan.overlap, + ) + for core in merged_cores + ), + key=lambda tile: (tile.y, tile.x), + ) + ) + batches, effective_batch_size = batch_latent_tiles(tiles, tile_batch_size) + return TiledDiffusionPlan( + latent_width=latent_width, + latent_height=latent_height, + tile_width=base_plan.tile_width, + tile_height=base_plan.tile_height, + overlap=base_plan.overlap, + requested_tile_batch_size=tile_batch_size, + tile_batch_size=effective_batch_size, + tiles=tiles, + batches=batches, + ) + + +def _validate_aspect_ratio( + segs: NativeSegs, + latent_height: int, + latent_width: int, +) -> None: + """Reject SEGS that cannot describe the sampled latent's image proportions.""" + + source_height, source_width = segs[0] + source_ratio = source_width / source_height + latent_ratio = latent_width / latent_height + if abs(source_ratio - latent_ratio) / source_ratio <= 0.02: + return + raise ValueError( + "SEGS-guided tiled diffusion requires SEGS to match the latent image " + f"aspect ratio; SEGS is {source_height}x{source_width}, latent is " + f"{latent_height}x{latent_width}." + ) + + +def _build_ownership_cores( + segs: NativeSegs, + *, + latent_height: int, + latent_width: int, +) -> tuple[_OwnershipCore, ...]: + """Resolve overlapping SEGS into one deterministic latent ownership partition.""" + + source_height, source_width = segs[0] + segment_masks = tuple( + _segment_mask_to_latent( + segment, + source_height=source_height, + source_width=source_width, + latent_height=latent_height, + latent_width=latent_width, + ) + for segment in segs[1] + ) + ranked_indexes = sorted( + range(len(segment_masks)), + key=lambda index: ( + int(segment_masks[index].sum().item()), + -float(segs[1][index].confidence), + index, + ), + ) + occupied = torch.zeros((latent_height, latent_width), dtype=torch.bool) + cores: list[_OwnershipCore] = [] + for index in ranked_indexes: + owned = torch.logical_and(segment_masks[index], torch.logical_not(occupied)) + if bool(owned.any()): + cores.append(_core_from_mask(owned)) + occupied = torch.logical_or(occupied, segment_masks[index]) + background = torch.logical_not(occupied) + if bool(background.any()): + cores.append(_core_from_mask(background)) + if cores: + return tuple(cores) + return ( + _core_from_mask(torch.ones((latent_height, latent_width), dtype=torch.bool)), + ) + + +def _segment_mask_to_latent( + segment: Segment, + *, + source_height: int, + source_width: int, + latent_height: int, + latent_width: int, +) -> torch.Tensor: + """Restore one crop-local SEG mask and map it to a latent-space mask.""" + + crop = segment.crop_region + if ( + crop.left < 0 + or crop.top < 0 + or crop.right > source_width + or crop.bottom > source_height + or crop.width < 1 + or crop.height < 1 + ): + raise ValueError( + "SEGS-guided tiled diffusion requires every SEG crop_region to fit " + "inside the SEGS header dimensions." + ) + local_mask = ( + torch.as_tensor(segment.cropped_mask, dtype=torch.float32).detach().cpu() + ) + if local_mask.ndim == 3 and int(local_mask.shape[0]) == 1: + local_mask = local_mask.squeeze(0) + if local_mask.shape != (crop.height, crop.width): + raise ValueError( + "SEGS-guided tiled diffusion requires each cropped_mask to match its " + "crop_region." + ) + latent_top, latent_bottom = _latent_sample_range( + crop.top, + crop.bottom, + source_height, + latent_height, + ) + latent_left, latent_right = _latent_sample_range( + crop.left, + crop.right, + source_width, + latent_width, + ) + latent_mask = torch.zeros((latent_height, latent_width), dtype=torch.bool) + if latent_bottom <= latent_top or latent_right <= latent_left: + return latent_mask + sampled_rows = ( + torch.div( + torch.arange(latent_top, latent_bottom) * source_height, + latent_height, + rounding_mode="floor", + ) + - crop.top + ) + sampled_columns = ( + torch.div( + torch.arange(latent_left, latent_right) * source_width, + latent_width, + rounding_mode="floor", + ) + - crop.left + ) + sampled_mask = ( + local_mask.clamp(0.0, 1.0) + .index_select( + 0, + sampled_rows, + ) + .index_select(1, sampled_columns) + ) + latent_mask[latent_top:latent_bottom, latent_left:latent_right] = ( + sampled_mask >= 0.5 + ) + return latent_mask + + +def _split_core( + core: _OwnershipCore, + *, + max_width: int, + max_height: int, +) -> tuple[_OwnershipCore, ...]: + """Recursively divide a core into balanced pieces that fit its tile budget.""" + + left, top, right, bottom = core.bounds + width = right - left + height = bottom - top + if width <= max_width and height <= max_height: + return (core,) + split_x = width / max_width >= height / max_height + first, second = _split_mask_at_balanced_axis(core.mask, core.bounds, split_x) + return _split_core( + _core_from_mask(first), max_width=max_width, max_height=max_height + ) + _split_core(_core_from_mask(second), max_width=max_width, max_height=max_height) + + +def _split_mask_at_balanced_axis( + mask: torch.Tensor, + bounds: tuple[int, int, int, int], + split_x: bool, +) -> tuple[torch.Tensor, torch.Tensor]: + """Split one non-empty mask near its active-pixel median on one axis.""" + + left, top, right, bottom = bounds + counts = ( + mask[top:bottom, left:right].sum(dim=0) + if split_x + else mask[top:bottom, left:right].sum(dim=1) + ) + cumulative = torch.cumsum(counts, dim=0) + midpoint = int(torch.searchsorted(cumulative, cumulative[-1] / 2, right=False)) + axis_start = left if split_x else top + axis_end = right if split_x else bottom + split_at = min(axis_end - 1, max(axis_start + 1, axis_start + midpoint + 1)) + first = mask.clone() + second = mask.clone() + if split_x: + first[:, split_at:] = False + second[:, :split_at] = False + else: + first[split_at:, :] = False + second[:split_at, :] = False + if not bool(first.any()) or not bool(second.any()): + raise ValueError("Unable to split an oversized SEGS-guided tile core.") + return first, second + + +def _merge_small_cores( + cores: tuple[_OwnershipCore, ...], + *, + max_width: int, + max_height: int, +) -> tuple[_OwnershipCore, ...]: + """Greedily combine small nearby cores when one bounded window can hold both.""" + + pending = list(cores) + minimum_area = max(1, (max_width * max_height) // 4) + merged = True + while merged: + merged = False + for index, core in enumerate(tuple(pending)): + if core.area >= minimum_area: + continue + candidate_index = _best_merge_candidate_index( + core, + pending, + excluded_index=index, + max_width=max_width, + max_height=max_height, + ) + if candidate_index is None: + continue + candidate = pending[candidate_index] + pending[index] = _OwnershipCore( + mask=torch.logical_or(core.mask, candidate.mask), + bounds=_union_bounds(core.bounds, candidate.bounds), + area=core.area + candidate.area, + ) + pending.pop(candidate_index) + merged = True + break + return tuple(pending) + + +def _best_merge_candidate_index( + core: _OwnershipCore, + candidates: list[_OwnershipCore], + *, + excluded_index: int, + max_width: int, + max_height: int, +) -> int | None: + """Return a candidate index whose combined bounds fit one ownership budget.""" + + eligible: list[tuple[int, int, int]] = [] + for index, candidate in enumerate(candidates): + if index == excluded_index: + continue + bounds = _union_bounds(core.bounds, candidate.bounds) + left, top, right, bottom = bounds + width = right - left + height = bottom - top + if width > max_width or height > max_height: + continue + distance = _bounds_distance(core.bounds, candidate.bounds) + eligible.append((width * height, distance, index)) + if not eligible: + return None + return min(eligible, key=lambda item: (item[0], item[1]))[2] + + +def _bounds_distance( + first: tuple[int, int, int, int], + second: tuple[int, int, int, int], +) -> int: + """Return the axis-aligned gap between two mask bounding boxes.""" + + left, top, right, bottom = first + other_left, other_top, other_right, other_bottom = second + horizontal = max(0, other_left - right, left - other_right) + vertical = max(0, other_top - bottom, top - other_bottom) + return horizontal + vertical + + +def _union_bounds( + first: tuple[int, int, int, int], + second: tuple[int, int, int, int], +) -> tuple[int, int, int, int]: + """Return the tight rectangle containing both ownership-core bounds.""" + + return ( + min(first[0], second[0]), + min(first[1], second[1]), + max(first[2], second[2]), + max(first[3], second[3]), + ) + + +def _tile_for_core( + core: _OwnershipCore, + *, + latent_width: int, + latent_height: int, + tile_width: int, + tile_height: int, + overlap: int, +) -> LatentTile: + """Place one bounded sampling window around an irregular ownership core.""" + + left, top, right, bottom = core.bounds + center_x = (left + right) / 2.0 + center_y = (top + bottom) / 2.0 + x = _clamp_window_start(center_x, tile_width, latent_width) + y = _clamp_window_start(center_y, tile_height, latent_height) + weight_mask = _feathered_tile_weight( + core.mask, + x=x, + y=y, + width=tile_width, + height=tile_height, + overlap=overlap, + ) + if not bool((weight_mask > 0).any()): + raise ValueError("SEGS-guided tiled diffusion generated an empty tile weight.") + return LatentTile(x, y, tile_width, tile_height, weight_mask) + + +def _clamp_window_start(center: float, window_size: int, limit: int) -> int: + """Center a fixed sampling window while keeping it inside the latent bounds.""" + + desired = round(center - window_size / 2.0) + return min(max(0, desired), limit - window_size) + + +def _feathered_tile_weight( + mask: torch.Tensor, + *, + x: int, + y: int, + width: int, + height: int, + overlap: int, +) -> torch.Tensor: + """Build one feathered tile weight without blurring the full latent mask.""" + + if overlap == 0: + return mask[y : y + height, x : x + width].float().contiguous() + radius = max(1, overlap // 2) + source_left = max(0, x - radius) + source_top = max(0, y - radius) + source_right = min(int(mask.shape[1]), x + width + radius) + source_bottom = min(int(mask.shape[0]), y + height + radius) + local_weight = ( + functional.avg_pool2d( + mask[source_top:source_bottom, source_left:source_right] + .float() + .unsqueeze(0) + .unsqueeze(0), + kernel_size=radius * 2 + 1, + stride=1, + padding=radius, + count_include_pad=False, + ) + .squeeze(0) + .squeeze(0) + ) + local_y = y - source_top + local_x = x - source_left + return local_weight[ + local_y : local_y + height, local_x : local_x + width + ].contiguous() + + +def _latent_sample_range( + source_start: int, + source_end: int, + source_limit: int, + latent_limit: int, +) -> tuple[int, int]: + """Return latent coordinates whose nearest samples fall in a source interval.""" + + start = (source_start * latent_limit + source_limit - 1) // source_limit + end = (source_end * latent_limit + source_limit - 1) // source_limit + return max(0, min(latent_limit, start)), max(0, min(latent_limit, end)) + + +def _core_from_mask(mask: torch.Tensor) -> _OwnershipCore: + """Build one core with bounds and area computed exactly once.""" + + bounds = _mask_bounds(mask) + if bounds is None: + raise ValueError("SEGS-guided tiled diffusion cannot use an empty core.") + return _OwnershipCore( + mask=mask, + bounds=bounds, + area=int(mask.sum().item()), + ) + + +def _mask_bounds(mask: torch.Tensor) -> tuple[int, int, int, int] | None: + """Return left, top, right, bottom bounds for one non-empty boolean mask.""" + + y_coords, x_coords = torch.where(mask) + if y_coords.numel() == 0: + return None + return ( + int(x_coords.min().item()), + int(y_coords.min().item()), + int(x_coords.max().item()) + 1, + int(y_coords.max().item()) + 1, + ) diff --git a/simple_syrup/domain/tiled_diffusion.py b/simple_syrup/domain/tiled_diffusion.py index e9e2951..be5714e 100644 --- a/simple_syrup/domain/tiled_diffusion.py +++ b/simple_syrup/domain/tiled_diffusion.py @@ -20,12 +20,13 @@ TILED_DIFFUSION_MODES = ("multidiffusion", "mixture_of_diffusers") @dataclass(frozen=True) class LatentTile: - """Describe one rectangular latent-space tile.""" + """Describe one rectangular latent-space tile and optional local blend weights.""" x: int y: int width: int height: int + weight_mask: torch.Tensor | None = None @property def slicer(self) -> tuple[slice, slice, slice, slice]: @@ -83,7 +84,7 @@ def build_tiled_diffusion_plan( tile_height=effective_tile_height, overlap=effective_overlap, ) - batches, effective_tile_batch_size = _batch_tiles(tiles, tile_batch_size) + batches, effective_tile_batch_size = batch_latent_tiles(tiles, tile_batch_size) return TiledDiffusionPlan( latent_width=latent_width, latent_height=latent_height, @@ -214,7 +215,7 @@ def _split_tiles( return tuple(tiles) -def _batch_tiles( +def batch_latent_tiles( tiles: tuple[LatentTile, ...], requested_tile_batch_size: int, ) -> tuple[tuple[tuple[LatentTile, ...], ...], int]: diff --git a/simple_syrup/nodes/ksampler_tiled_diffusion.py b/simple_syrup/nodes/ksampler_tiled_diffusion.py index ada4761..4454e3a 100644 --- a/simple_syrup/nodes/ksampler_tiled_diffusion.py +++ b/simple_syrup/nodes/ksampler_tiled_diffusion.py @@ -153,7 +153,18 @@ class KSamplerTiledDiffusion: "tooltip": tooltips.LATENT_TILE_BATCH_SIZE, }, ), - } + }, + "optional": { + "segs": ( + "SEGS", + { + "tooltip": ( + "Optional image regions that guide irregular tile " + "boundaries while preserving the configured overlap." + ), + }, + ), + }, } def sample( @@ -173,6 +184,7 @@ class KSamplerTiledDiffusion: latent_tile_height: int = 128, latent_tile_overlap: int = 16, latent_tile_batch_size: int = 4, + segs: object | None = None, ) -> tuple[Latent]: """Sample a latent with the selected tiled diffusion method.""" @@ -193,5 +205,6 @@ class KSamplerTiledDiffusion: latent_tile_overlap=latent_tile_overlap, latent_tile_batch_size=latent_tile_batch_size, preview_context=None, + segs=segs, ) return (output,) diff --git a/simple_syrup/nodes/segs_from_sam_output.py b/simple_syrup/nodes/segs_from_sam_output.py new file mode 100644 index 0000000..c6d1e80 --- /dev/null +++ b/simple_syrup/nodes/segs_from_sam_output.py @@ -0,0 +1,120 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""ComfyUI node declaration for automatic SEGS from a SAM model.""" + +from __future__ import annotations + +from collections.abc import Callable +from typing import Any, ClassVar + +from ..masking.segs_mask_ops import iter_single_images, validate_image_batch +from ..runtime.progress import PhaseProgressReporter, create_comfy_phase_progress +from ..services.segs_from_sam_output_service import SEGSFromSAMOutputService + + +class SEGSFromSAMOutput: + """Generate reusable unprompted SEGS from a connected SAM model.""" + + service_class: ClassVar[type[SEGSFromSAMOutputService]] = SEGSFromSAMOutputService + progress_factory: ClassVar[Callable[..., PhaseProgressReporter]] = ( + create_comfy_phase_progress + ) + + RETURN_TYPES = ("SEGS",) + RETURN_NAMES = ("segs",) + OUTPUT_IS_LIST = (True,) + OUTPUT_TOOLTIPS = ( + "Automatic image regions as SEGS for detailing, masking, or tiled diffusion.", + ) + FUNCTION = "generate" + CATEGORY = "SimpleSyrup/Detection" + DESCRIPTION = "Creates automatic, unprompted image SEGS from a connected SAM model." + SEARCH_ALIASES = ["sam", "automatic", "segment", "segmentation", "segs"] + + @classmethod + def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]: + """Declare inputs for automatic SAM-to-SEGS conversion.""" + + return { + "required": { + "image": ( + "IMAGE", + {"tooltip": "Image whose automatic regions become SEGS."}, + ), + "sam_model": ( + "SAM_MODEL", + {"tooltip": "SAM model used to find unprompted image regions."}, + ), + "segmentation_resolution": ( + "INT", + { + "default": 640, + "min": 64, + "max": 8192, + "step": 64, + "tooltip": ( + "Maximum long edge in pixels used for segmentation. " + "Lower values run faster and omit smaller details." + ), + }, + ), + "minimum_region_area": ( + "INT", + { + "default": 0, + "min": 0, + "max": 268435456, + "step": 1, + "tooltip": ( + "Discard masks smaller than this many pixels in the " + "original image." + ), + }, + ), + } + } + + def generate( + self, + image: object, + sam_model: object, + segmentation_resolution: int = 640, + minimum_region_area: int = 0, + ) -> tuple[list[object]]: + """Return one automatic SEGS payload for each image batch item.""" + + image_batch = validate_image_batch(image, "SEGS from SAM Output") + service = self.service_class() + phase_progress = type(self).progress_factory( + operation="segs_from_sam_output", + subject=_sam_model_subject(sam_model), + total_phases=int(image_batch.shape[0]) * 3 + 1, + ) + outputs: list[object] = [] + try: + for single_image in iter_single_images(image_batch): + outputs.append( + service.build( + image=single_image, + sam_model=sam_model, + segmentation_resolution=segmentation_resolution, + minimum_region_area=minimum_region_area, + phase_progress=phase_progress, + ) + ) + except Exception: + phase_progress.advance("failed") + raise + phase_progress.advance("completed") + return (outputs,) + + +def _sam_model_subject(sam_model: object) -> str: + """Return a concise model identity for Comfy progress diagnostics.""" + + model_id = getattr(sam_model, "model_id", None) + if isinstance(model_id, str) and model_id: + return model_id + return type(sam_model).__name__ diff --git a/simple_syrup/nodes_v3/__init__.py b/simple_syrup/nodes_v3/__init__.py index ad3da19..3bbfa89 100644 --- a/simple_syrup/nodes_v3/__init__.py +++ b/simple_syrup/nodes_v3/__init__.py @@ -39,6 +39,7 @@ def get_nodes() -> list[type[object]]: ResizeImageToTargetV3, SAMModelLoaderV3, SeedV3, + SEGSFromSAMOutputV3, SimpleLoadAnimaV3, SimpleVAEEncodeV3, UpscaleLatentFromImageV3, @@ -83,6 +84,7 @@ def get_nodes() -> list[type[object]]: PromptSEGSWithSAMV3, ResizeImageToTargetV3, SAMModelLoaderV3, + SEGSFromSAMOutputV3, ScaleFactorV3, SeedV3, SimpleLoadAnimaV3, diff --git a/simple_syrup/nodes_v3/legacy_node_wrappers.py b/simple_syrup/nodes_v3/legacy_node_wrappers.py index b99801e..1c29167 100644 --- a/simple_syrup/nodes_v3/legacy_node_wrappers.py +++ b/simple_syrup/nodes_v3/legacy_node_wrappers.py @@ -37,6 +37,7 @@ from ..nodes.prompt_segs_with_sam import PromptSEGSWithSAM from ..nodes.provenance_latent import SimpleVAEEncode, UpscaleLatentFromImage from ..nodes.sam_model_loader import SAMModelLoader from ..nodes.seed import Seed +from ..nodes.segs_from_sam_output import SEGSFromSAMOutput from ..nodes.simple_load_anima import SimpleLoadAnima from ..nodes.vitmatte_model_loader import ViTMatteModelLoader @@ -254,6 +255,14 @@ class SAMModelLoaderV3(LegacyNodeV3Adapter): DISPLAY_NAME = "SAM Model Loader" +class SEGSFromSAMOutputV3(LegacyNodeV3Adapter): + """Expose automatic SAM-to-SEGS generation through Comfy v3 only.""" + + LEGACY_NODE_CLASS = SEGSFromSAMOutput + NODE_ID = "SimpleSyrup.SEGSFromSAMOutput" + DISPLAY_NAME = "SEGS from SAM Output" + + class SeedV3(LegacyNodeV3Adapter): """Expose Seed through Comfy v3 only.""" @@ -505,6 +514,7 @@ __all__ = [ "PromptSEGSWithSAMV3", "ResizeImageToTargetV3", "SAMModelLoaderV3", + "SEGSFromSAMOutputV3", "SeedV3", "SimpleLoadAnimaV3", "SimpleVAEEncodeV3", diff --git a/simple_syrup/runtime/mixture_of_diffusers_sampling.py b/simple_syrup/runtime/mixture_of_diffusers_sampling.py index fda3501..249a9a4 100644 --- a/simple_syrup/runtime/mixture_of_diffusers_sampling.py +++ b/simple_syrup/runtime/mixture_of_diffusers_sampling.py @@ -31,6 +31,7 @@ from .tiled_sampling import ( ApplyModel, Latent, ModelFunctionWrapper, + SemanticTileWeightCache, make_tiled_model_args, new_spatial_weight_buffer, reject_unsupported_conditioning, @@ -63,6 +64,7 @@ def sample_mixture_of_diffusers( preview_context: DetailPreviewContext | None = None, differential_diffusion: bool = False, allow_full_context_masks: bool = False, + tiled_plan: TiledDiffusionPlan | None = None, ) -> Latent: """Sample a latent with a cloned model patched for Mixture of Diffusers.""" @@ -117,6 +119,7 @@ def sample_mixture_of_diffusers( overlap=latent_tile_overlap, tile_batch_size=latent_tile_batch_size, differential_diffusion=differential_diffusion, + tiled_plan=tiled_plan, ) batch_inds = latent_image["batch_index"] if "batch_index" in latent_image else None @@ -173,10 +176,11 @@ def clone_model_with_mixture_of_diffusers( overlap: int, tile_batch_size: int, differential_diffusion: bool = False, + tiled_plan: TiledDiffusionPlan | None = None, ) -> tuple[Any, TiledDiffusionPlan]: """Return a model clone patched with a pre-CFG Mixture wrapper.""" - plan = build_tiled_diffusion_plan( + plan = tiled_plan or build_tiled_diffusion_plan( latent_width=latent_width, latent_height=latent_height, tile_width=tile_width, @@ -184,6 +188,7 @@ def clone_model_with_mixture_of_diffusers( overlap=overlap, tile_batch_size=tile_batch_size, ) + _validate_supplied_plan(plan, latent_width, latent_height) cloned_model = model.clone() if differential_diffusion: install_differential_diffusion(cloned_model) @@ -213,6 +218,7 @@ class MixtureOfDiffusersModelWrapper: self._plan = plan self._existing_wrapper = existing_wrapper self._tile_weights_2d: torch.Tensor | None = None + self._semantic_tile_weights = SemanticTileWeightCache(plan.tiles) def __call__( self, @@ -245,7 +251,6 @@ class MixtureOfDiffusersModelWrapper: output_buffer = torch.zeros_like(x) weight_buffer = new_spatial_weight_buffer(x, self._plan) input_batch_size = int(x.shape[0]) - weights = self._weights_for(x) for batch in self._plan.batches: tiled_args = self._make_tiled_args( @@ -254,14 +259,22 @@ class MixtureOfDiffusersModelWrapper: input_batch_size=input_batch_size, ) tile_output = self._call_original(apply_model, tiled_args) + weights = self._weights_for(tile_output) + accumulation_weights = weights.to(dtype=weight_buffer.dtype) + semantic_weights = self._semantic_tile_weights.for_output(tile_output) for index, tile in enumerate(batch): tile_slice = spatial_tile_slicer(tile, x.ndim) start = index * input_batch_size end = start + input_batch_size - output_buffer[tile_slice] += tile_output[start:end] * weights.to( - dtype=tile_output.dtype + model_weight, accumulation_weight = ( + self._semantic_tile_weights.for_tile( + semantic_weights, + tile, + ) ) - weight_buffer[tile_slice] += weights.to(dtype=weight_buffer.dtype) + tile_weight = weights * model_weight + output_buffer[tile_slice] += tile_output[start:end] * tile_weight + weight_buffer[tile_slice] += accumulation_weights * accumulation_weight return output_buffer / weight_buffer.to(dtype=output_buffer.dtype) @@ -315,6 +328,21 @@ class MixtureOfDiffusersModelWrapper: ) +def _validate_supplied_plan( + plan: TiledDiffusionPlan, + latent_width: int, + latent_height: int, +) -> None: + """Reject a semantic tile plan that belongs to another latent shape.""" + + if (plan.latent_width, plan.latent_height) == (latent_width, latent_height): + return + raise ValueError( + "Mixture of Diffusers tiled plan dimensions must match the sampled latent " + "shape." + ) + + def _comfy_sample() -> ModuleType: """Import ComfyUI sample helpers lazily.""" diff --git a/simple_syrup/runtime/model_catalog.py b/simple_syrup/runtime/model_catalog.py index f8a8159..5c79ad0 100644 --- a/simple_syrup/runtime/model_catalog.py +++ b/simple_syrup/runtime/model_catalog.py @@ -190,6 +190,25 @@ SAM_ENTRIES: tuple[ModelEntry, ...] = ( ), ), ), + ModelEntry( + entry_id="fast_sam_s", + display_name="FastSAM-s (23MB)", + family=ModelFamily.SAM, + model_type="fast_sam", + source_repo="ultralytics/assets", + artifacts=( + ModelArtifact( + artifact_id="fast_sam_s_checkpoint", + filename="FastSAM-s.pt", + folder_name="sams", + source_url=( + "https://github.com/ultralytics/assets/releases/latest/download/" + "FastSAM-s.pt" + ), + description="FastSAM-s checkpoint", + ), + ), + ), ) GROUNDING_DINO_ENTRIES: tuple[ModelEntry, ...] = ( diff --git a/simple_syrup/runtime/multidiffusion_sampling.py b/simple_syrup/runtime/multidiffusion_sampling.py index c9a732a..5db98e7 100644 --- a/simple_syrup/runtime/multidiffusion_sampling.py +++ b/simple_syrup/runtime/multidiffusion_sampling.py @@ -30,6 +30,7 @@ from .tiled_sampling import ( ApplyModel, Latent, ModelFunctionWrapper, + SemanticTileWeightCache, make_tiled_model_args, new_spatial_weight_buffer, reject_unsupported_conditioning, @@ -63,6 +64,7 @@ def sample_multidiffusion( preview_context: DetailPreviewContext | None = None, differential_diffusion: bool = False, allow_full_context_masks: bool = False, + tiled_plan: TiledDiffusionPlan | None = None, ) -> Latent: """Sample a latent with a cloned model patched for MultiDiffusion.""" @@ -118,6 +120,7 @@ def sample_multidiffusion( overlap=latent_tile_overlap, tile_batch_size=latent_tile_batch_size, differential_diffusion=differential_diffusion, + tiled_plan=tiled_plan, ) batch_inds = latent_image["batch_index"] if "batch_index" in latent_image else None @@ -175,10 +178,11 @@ def clone_model_with_multidiffusion( overlap: int, tile_batch_size: int, differential_diffusion: bool = False, + tiled_plan: TiledDiffusionPlan | None = None, ) -> tuple[Any, TiledDiffusionPlan]: """Return a model clone patched with a pre-CFG MultiDiffusion wrapper.""" - plan = build_tiled_diffusion_plan( + plan = tiled_plan or build_tiled_diffusion_plan( latent_width=latent_width, latent_height=latent_height, tile_width=tile_width, @@ -186,6 +190,7 @@ def clone_model_with_multidiffusion( overlap=overlap, tile_batch_size=tile_batch_size, ) + _validate_supplied_plan(plan, latent_width, latent_height) cloned_model = model.clone() if differential_diffusion: install_differential_diffusion(cloned_model) @@ -214,6 +219,7 @@ class MultiDiffusionModelWrapper: self._plan = plan self._existing_wrapper = existing_wrapper + self._semantic_tile_weights = SemanticTileWeightCache(plan.tiles) def __call__( self, @@ -254,12 +260,19 @@ class MultiDiffusionModelWrapper: input_batch_size=input_batch_size, ) tile_output = self._call_original(apply_model, tiled_args) + tile_weights = self._semantic_tile_weights.for_output(tile_output) for index, tile in enumerate(batch): tile_slice = spatial_tile_slicer(tile, x.ndim) start = index * input_batch_size end = start + input_batch_size - output_buffer[tile_slice] += tile_output[start:end] - weight_buffer[tile_slice] += 1.0 + model_weight, accumulation_weight = ( + self._semantic_tile_weights.for_tile( + tile_weights, + tile, + ) + ) + output_buffer[tile_slice] += tile_output[start:end] * model_weight + weight_buffer[tile_slice] += accumulation_weight return output_buffer / weight_buffer.to(dtype=output_buffer.dtype) @@ -302,6 +315,20 @@ def _reject_unipc_sampler(sampler_name: str) -> None: raise ValueError("MultiDiffusion is not compatible with UniPC samplers.") +def _validate_supplied_plan( + plan: TiledDiffusionPlan, + latent_width: int, + latent_height: int, +) -> None: + """Reject a semantic tile plan that belongs to another latent shape.""" + + if (plan.latent_width, plan.latent_height) == (latent_width, latent_height): + return + raise ValueError( + "MultiDiffusion tiled plan dimensions must match the sampled latent shape." + ) + + def _sampling_callback( model: Any, steps: int, diff --git a/simple_syrup/runtime/sam_automatic_segmenter.py b/simple_syrup/runtime/sam_automatic_segmenter.py new file mode 100644 index 0000000..853cfb0 --- /dev/null +++ b/simple_syrup/runtime/sam_automatic_segmenter.py @@ -0,0 +1,253 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Run unprompted automatic segmentation with supported SAM model families.""" + +from __future__ import annotations + +import importlib +from dataclasses import dataclass +from typing import Any, Protocol, cast + +import numpy as np +import torch + +from .loaded_models import LoadedSAMModel, unwrap_sam_model +from .model_device_manager import TorchModelDeviceManager, external_model_inference +from .sam_loader import SAM_HQ_RUNTIME_PACKAGE + + +@dataclass(frozen=True) +class AutomaticSAMMask: + """Describe one unprompted mask emitted by a SAM-compatible runtime.""" + + mask: torch.Tensor + confidence: float + label: str | None = None + + +class SAMAutomaticSegmenter(Protocol): + """Generate unprompted masks from one SAM-compatible model and image.""" + + def segment_all( + self, + sam_model: object, + image: torch.Tensor, + execution_device: str = "auto", + ) -> tuple[AutomaticSAMMask, ...]: + """Return source-image masks with optional confidence and labels.""" + + +class SAMModelAutomaticSegmenter: + """Adapt SimpleSyrup SAM models to their unprompted segmentation APIs.""" + + def segment_all( + self, + sam_model: object, + image: torch.Tensor, + execution_device: str = "auto", + ) -> tuple[AutomaticSAMMask, ...]: + """Return automatic masks for one single-image BHWC tensor.""" + + image_array = _tensor_to_rgb_array(image) + model = unwrap_sam_model(sam_model) + if _is_fast_sam(sam_model, model): + return self._segment_fast_sam( + sam_model=sam_model, + model=model, + image_array=image_array, + execution_device=execution_device, + ) + return self._segment_segment_anything( + sam_model=sam_model, + model=model, + image_array=image_array, + execution_device=execution_device, + ) + + def _segment_fast_sam( + self, + *, + sam_model: object, + model: object, + image_array: np.ndarray[Any, Any], + execution_device: str, + ) -> tuple[AutomaticSAMMask, ...]: + """Run FastSAM's unprompted everything-mask path.""" + + if ( + isinstance(sam_model, LoadedSAMModel) + and sam_model.managed_model is not None + ): + with TorchModelDeviceManager().inference( + sam_model.managed_model, + execution_device, + ) as loaded: + return _fast_sam_masks(loaded.model, image_array, loaded.device) + with external_model_inference(model, execution_device) as loaded: + return _fast_sam_masks(loaded.model, image_array, loaded.device) + + def _segment_segment_anything( + self, + *, + sam_model: object, + model: object, + image_array: np.ndarray[Any, Any], + execution_device: str, + ) -> tuple[AutomaticSAMMask, ...]: + """Run Segment Anything automatic-mask generation under device management.""" + + if ( + isinstance(sam_model, LoadedSAMModel) + and sam_model.managed_model is not None + ): + with TorchModelDeviceManager().inference( + sam_model.managed_model, + execution_device, + ) as loaded: + return _segment_anything_masks( + loaded.model, + image_array, + use_hq_generator=_uses_sam_hq_generator(sam_model), + ) + with external_model_inference(model, execution_device) as loaded: + return _segment_anything_masks( + loaded.model, + image_array, + use_hq_generator=_uses_sam_hq_generator(model), + ) + + +def _tensor_to_rgb_array(image: torch.Tensor) -> np.ndarray[Any, Any]: + """Convert one BHWC ComfyUI image into a uint8 RGB array.""" + + if image.ndim != 4 or int(image.shape[0]) != 1: + raise ValueError( + "SAM automatic segmentation requires one BHWC image at a time." + ) + sample = image[0].detach().cpu().float().clamp(0.0, 1.0).numpy() + channels = int(sample.shape[-1]) + if channels == 1: + sample = np.repeat(sample, 3, axis=-1) + elif channels >= 3: + sample = sample[..., :3] + else: + raise ValueError( + "SAM automatic segmentation requires at least one image channel." + ) + return (sample * 255.0).round().astype(np.uint8) + + +def _is_fast_sam(container: object, model: object) -> bool: + """Return whether a loaded or external model uses FastSAM's API.""" + + if isinstance(container, LoadedSAMModel): + return container.model_id.startswith("fast_sam") + return type(model).__name__ == "FastSAM" + + +def _fast_sam_masks( + model: object, + image_array: np.ndarray[Any, Any], + device: torch.device, +) -> tuple[AutomaticSAMMask, ...]: + """Extract FastSAM masks and scores from one Ultralytics result.""" + + predict = getattr(model, "predict", None) + if not callable(predict): + raise TypeError("FastSAM model does not expose the required predict method.") + results = predict( + image_array, + imgsz=max(image_array.shape[:2]), + retina_masks=True, + verbose=False, + device=str(device), + ) + if not isinstance(results, list | tuple) or not results: + return () + result = results[0] + masks_container = getattr(result, "masks", None) + mask_data = getattr(masks_container, "data", None) + if not isinstance(mask_data, torch.Tensor): + return () + boxes = getattr(result, "boxes", None) + confidences = getattr(boxes, "conf", None) + entries: list[AutomaticSAMMask] = [] + for index, mask in enumerate(mask_data.detach().cpu()): + confidence = 1.0 + if isinstance(confidences, torch.Tensor) and index < int(confidences.numel()): + confidence = float(confidences[index].detach().cpu().item()) + entries.append( + AutomaticSAMMask( + mask=mask.float().clamp(0.0, 1.0), + confidence=_clamp_confidence(confidence), + ) + ) + return tuple(entries) + + +def _segment_anything_masks( + model: object, + image_array: np.ndarray[Any, Any], + *, + use_hq_generator: bool, +) -> tuple[AutomaticSAMMask, ...]: + """Extract binary masks and predicted quality from automatic generators.""" + + generator_class = _automatic_generator_class(use_hq_generator) + generated = generator_class(model).generate(image_array) + if not isinstance(generated, list): + raise TypeError("SAM automatic mask generator returned an invalid result.") + masks: list[AutomaticSAMMask] = [] + for entry in generated: + if not isinstance(entry, dict) or "segmentation" not in entry: + raise ValueError("SAM automatic mask generator returned an invalid mask.") + raw_confidence = entry.get("predicted_iou", entry.get("stability_score", 1.0)) + confidence = ( + float(raw_confidence) if isinstance(raw_confidence, int | float) else 1.0 + ) + masks.append( + AutomaticSAMMask( + mask=torch.as_tensor(entry["segmentation"], dtype=torch.float32) + .detach() + .cpu() + .clamp(0.0, 1.0), + confidence=_clamp_confidence(confidence), + ) + ) + return tuple(masks) + + +def _automatic_generator_class(use_hq_generator: bool) -> type[Any]: + """Return the automatic generator matching the loaded SAM family.""" + + try: + if use_hq_generator: + automatic = importlib.import_module(f"{SAM_HQ_RUNTIME_PACKAGE}.automatic") + return cast(type[Any], automatic.SamAutomaticMaskGeneratorHQ) + segment_anything = importlib.import_module("segment_anything") + return cast(type[Any], segment_anything.SamAutomaticMaskGenerator) + except ImportError as error: + raise RuntimeError( + "SAM automatic segmentation requires the matching Segment Anything " + "runtime. " + f"Import failed: {error}." + ) from error + + +def _uses_sam_hq_generator(model: object) -> bool: + """Return whether a model container or raw object requires SAM-HQ generation.""" + + if isinstance(model, LoadedSAMModel): + return model.model_id.startswith("sam_hq") or model.model_id == "mobile_sam" + model_name = getattr(model, "model_name", "") + return isinstance(model_name, str) and ( + model_name.startswith("sam_hq") or model_name == "mobile_sam" + ) + + +def _clamp_confidence(value: float) -> float: + """Return one confidence value constrained to the public SEGS range.""" + + return max(0.0, min(1.0, value)) diff --git a/simple_syrup/runtime/sam_loader.py b/simple_syrup/runtime/sam_loader.py index 06b78ea..5a515d5 100644 --- a/simple_syrup/runtime/sam_loader.py +++ b/simple_syrup/runtime/sam_loader.py @@ -11,7 +11,7 @@ from collections.abc import MutableMapping from dataclasses import dataclass from pathlib import Path from types import ModuleType -from typing import Any +from typing import Any, cast from ..shared.logging import get_logger from .loaded_models import LoadedSAMModel @@ -118,7 +118,7 @@ class SAMLoaderService: """Load and wrap a SAM model after artifact resolution and cache lookup.""" phase_progress.advance("loading_checkpoint") - model = self._load_segment_anything_model(entry, checkpoint_path) + model = self._load_model(entry, checkpoint_path) phase_progress.advance("registering_device_management") managed_model = self._device_manager.manage( model, @@ -184,12 +184,15 @@ class SAMLoaderService: paths.append(result.path) return paths - def _load_segment_anything_model( + def _load_model( self, entry: ModelEntry, checkpoint_path: Path, ) -> object: - """Load a SAM model from the segment-anything registry.""" + """Load one SAM-compatible model from its owning runtime.""" + + if entry.model_type == "fast_sam": + return self._load_fast_sam_model(checkpoint_path) try: importlib.invalidate_caches() @@ -206,6 +209,19 @@ class SAMLoaderService: model.model_name = checkpoint_path.name return model + def _load_fast_sam_model(self, checkpoint_path: Path) -> object: + """Load FastSAM without delegating checkpoint download to Ultralytics.""" + + try: + ultralytics = importlib.import_module("ultralytics") + fast_sam_class = cast(Any, ultralytics).FastSAM + except (ImportError, AttributeError) as error: + raise RuntimeError( + "FastSAM support requires the installed ultralytics package. " + f"Import failed: {error}." + ) from error + return fast_sam_class(str(checkpoint_path)) + def _registry_module_name(model_type: str) -> str: """Return the registry module that owns one SAM-compatible model type.""" @@ -230,4 +246,6 @@ def _registry_import_error_message(model_type: str, error: ImportError) -> str: "runtime and its dependencies. Reinstall SimpleSyrup or restore " f"{SAM_HQ_RUNTIME_PACKAGE}. Import failed: {error}." ) + if model_type == "fast_sam": + return "FastSAM support requires the installed ultralytics package." return f"segment-anything is required to load SAM models. Import failed: {error}." diff --git a/simple_syrup/runtime/tiled_sampling.py b/simple_syrup/runtime/tiled_sampling.py index 1c1503e..fde1bbf 100644 --- a/simple_syrup/runtime/tiled_sampling.py +++ b/simple_syrup/runtime/tiled_sampling.py @@ -11,6 +11,7 @@ from __future__ import annotations from collections.abc import Callable, Sequence +from dataclasses import dataclass from typing import Any, TypeAlias import torch @@ -24,6 +25,72 @@ ModelFunctionWrapper: TypeAlias = Callable[[ApplyModel, dict[str, Any]], torch.T UNSUPPORTED_CONDITIONING_KEYS = frozenset({"area", "control", "gligen"}) +@dataclass(frozen=True) +class CachedTileWeights: + """Store model-output and float32 accumulation weights for one tensor layout.""" + + model: tuple[torch.Tensor, ...] + accumulation: tuple[torch.Tensor, ...] + + +class SemanticTileWeightCache: + """Keep semantic tile weights resident on the active sampling device.""" + + def __init__(self, tiles: Sequence[LatentTile]) -> None: + """Create an empty cache associated with one immutable tile plan.""" + + self._tiles = tuple(tiles) + self._tile_indexes = {id(tile): index for index, tile in enumerate(self._tiles)} + self._cache_key: tuple[torch.device, torch.dtype, int] | None = None + self._weights: CachedTileWeights | None = None + + def for_output(self, output: torch.Tensor) -> CachedTileWeights: + """Return tile weights shaped and typed for one model output tensor.""" + + cache_key = (output.device, output.dtype, output.ndim) + if self._cache_key == cache_key and self._weights is not None: + return self._weights + output_weights: list[torch.Tensor] = [] + accumulation_weights: list[torch.Tensor] = [] + for tile in self._tiles: + shape = (1,) * (output.ndim - 2) + (tile.height, tile.width) + if tile.weight_mask is None: + accumulation_weight = torch.ones( + shape, + device=output.device, + dtype=torch.float32, + ) + else: + if tuple(tile.weight_mask.shape) != (tile.height, tile.width): + raise ValueError( + "Semantic tile weight must match its tile dimensions." + ) + accumulation_weight = tile.weight_mask.to( + device=output.device, + dtype=torch.float32, + ).reshape(shape) + accumulation_weights.append(accumulation_weight) + output_weights.append(accumulation_weight.to(dtype=output.dtype)) + self._cache_key = cache_key + self._weights = CachedTileWeights( + model=tuple(output_weights), + accumulation=tuple(accumulation_weights), + ) + return self._weights + + def for_tile( + self, + weights: CachedTileWeights, + tile: LatentTile, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Return cached model and accumulation weights for one planned tile.""" + + index = self._tile_indexes.get(id(tile)) + if index is None: + raise ValueError("Semantic tile weight cache received an unknown tile.") + return weights.model[index], weights.accumulation[index] + + def validate_sampling_controls( *, steps: int, diff --git a/simple_syrup/services/segs_from_sam_output_service.py b/simple_syrup/services/segs_from_sam_output_service.py new file mode 100644 index 0000000..403f8d4 --- /dev/null +++ b/simple_syrup/services/segs_from_sam_output_service.py @@ -0,0 +1,372 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Convert unprompted SAM masks into image-associated SEGS.""" + +from __future__ import annotations + +from dataclasses import dataclass +from math import ceil, floor +from time import perf_counter + +import torch +import torch.nn.functional as functional + +from ..domain.segs import BoundingBox, CropRegion, NativeSegs, Segment +from ..masking.segs_mask_ops import validate_single_image +from ..runtime.progress import NullPhaseProgressReporter, PhaseProgressReporter +from ..runtime.sam_automatic_segmenter import ( + AutomaticSAMMask, + SAMAutomaticSegmenter, + SAMModelAutomaticSegmenter, +) +from ..shared.logging import get_logger + +LOGGER = get_logger(__name__) + + +@dataclass(frozen=True) +class SAMAutoSegsSettings: + """Validated controls for converting automatic SAM masks to SEGS.""" + + segmentation_resolution: int + minimum_region_area: int + + +@dataclass(frozen=True) +class _GuideMaskCandidate: + """Keep a retained SAM mask in compact guide-image coordinates.""" + + mask: torch.Tensor + bbox: BoundingBox + confidence: float + label: str | None + + +class SEGSFromSAMOutputService: + """Build reusable SEGS from a SAM model's unprompted masks.""" + + def __init__(self, segmenter: SAMAutomaticSegmenter | None = None) -> None: + """Create the service with an injectable automatic segmentation runtime.""" + + self._segmenter = segmenter or SAMModelAutomaticSegmenter() + + def build( + self, + *, + image: object, + sam_model: object, + segmentation_resolution: int, + minimum_region_area: int, + phase_progress: PhaseProgressReporter | None = None, + ) -> NativeSegs: + """Return source-sized SEGS from unprompted SAM masks.""" + + operation_started_at = perf_counter() + reporter = phase_progress or NullPhaseProgressReporter() + source_image = validate_single_image(image, "SEGS from SAM Output") + settings = _validate_settings( + segmentation_resolution=segmentation_resolution, + minimum_region_area=minimum_region_area, + ) + image_height = int(source_image.shape[1]) + image_width = int(source_image.shape[2]) + reporter.advance("preparing_segmentation_image") + guide_image = _resize_for_segmentation( + source_image, + maximum_long_edge=settings.segmentation_resolution, + ) + reporter.advance("generating_automatic_masks") + generated = self._segmenter.segment_all(sam_model, guide_image) + masks_generated_at = perf_counter() + reporter.advance("building_segs") + candidates = _normalize_masks( + generated, + source_height=image_height, + source_width=image_width, + minimum_region_area=settings.minimum_region_area, + ) + segments = tuple( + _build_segment( + image=source_image, + candidate=candidate, + source_height=image_height, + source_width=image_width, + confidence=candidate.confidence, + label=candidate.label or f"segment_{index:03d}", + ) + for index, candidate in enumerate(candidates, start=1) + ) + LOGGER.info( + "Built SEGS from SAM output", + extra={ + "operation": "segs_from_sam_output", + "segmentation_resolution": settings.segmentation_resolution, + "minimum_region_area": settings.minimum_region_area, + "source_height": image_height, + "source_width": image_width, + "generated_mask_count": len(generated), + "segment_count": len(segments), + "mask_generation_ms": round( + (masks_generated_at - operation_started_at) * 1000.0, + 2, + ), + "segs_construction_ms": round( + (perf_counter() - masks_generated_at) * 1000.0, + 2, + ), + }, + ) + return (image_height, image_width), segments + + +def _validate_settings( + *, + segmentation_resolution: int, + minimum_region_area: int, +) -> SAMAutoSegsSettings: + """Validate public automatic-SEGS configuration before model execution.""" + + if segmentation_resolution < 64 or segmentation_resolution % 64 != 0: + raise ValueError( + "segmentation_resolution must be at least 64 and divisible by 64." + ) + if minimum_region_area < 0: + raise ValueError("minimum_region_area must be greater than or equal to 0.") + return SAMAutoSegsSettings( + segmentation_resolution=segmentation_resolution, + minimum_region_area=minimum_region_area, + ) + + +def _resize_for_segmentation( + image: torch.Tensor, + *, + maximum_long_edge: int, +) -> torch.Tensor: + """Downscale an image to the requested guide resolution without upscaling it.""" + + source_height = int(image.shape[1]) + source_width = int(image.shape[2]) + source_long_edge = max(source_height, source_width) + if source_long_edge <= maximum_long_edge: + return image + scale = maximum_long_edge / source_long_edge + guide_height = max(1, round(source_height * scale)) + guide_width = max(1, round(source_width * scale)) + resized = functional.interpolate( + image.movedim(-1, 1), + size=(guide_height, guide_width), + mode="bilinear", + align_corners=False, + ) + return resized.movedim(1, -1).clamp(0.0, 1.0) + + +def _normalize_masks( + masks: tuple[AutomaticSAMMask, ...], + *, + source_height: int, + source_width: int, + minimum_region_area: int, +) -> tuple[_GuideMaskCandidate, ...]: + """Filter and deduplicate masks while they remain at guide resolution.""" + + normalized: list[_GuideMaskCandidate] = [] + for candidate in masks: + binary = _binary_guide_mask(candidate.mask) + bbox = _bbox_from_mask(binary) + if bbox is None: + continue + if ( + _projected_source_area( + active_pixels=int(binary.sum().item()), + guide_height=int(binary.shape[0]), + guide_width=int(binary.shape[1]), + source_height=source_height, + source_width=source_width, + ) + < minimum_region_area + ): + continue + normalized_candidate = _GuideMaskCandidate( + mask=binary, + bbox=bbox, + confidence=candidate.confidence, + label=candidate.label, + ) + duplicate_index = _duplicate_mask_index(normalized, normalized_candidate) + if duplicate_index is None: + normalized.append(normalized_candidate) + elif normalized_candidate.confidence > normalized[duplicate_index].confidence: + normalized[duplicate_index] = normalized_candidate + return tuple(normalized) + + +def _binary_guide_mask(mask: torch.Tensor) -> torch.Tensor: + """Return one validated binary guide-space mask on the CPU.""" + + working = mask.detach().cpu().float() + if working.ndim != 2: + raise ValueError("SAM automatic segmentation returned a mask that is not HW.") + if int(working.shape[0]) == 0 or int(working.shape[1]) == 0: + raise ValueError("SAM automatic segmentation returned an empty mask.") + return working >= 0.5 + + +def _projected_source_area( + *, + active_pixels: int, + guide_height: int, + guide_width: int, + source_height: int, + source_width: int, +) -> float: + """Estimate source-pixel area directly from one guide-space mask.""" + + return ( + active_pixels * source_height * source_width / float(guide_height * guide_width) + ) + + +def _duplicate_mask_index( + candidates: list[_GuideMaskCandidate], + candidate: _GuideMaskCandidate, +) -> int | None: + """Return a near-identical guide mask index without broad pairwise scans.""" + + for index, existing in enumerate(candidates): + if ( + tuple(existing.mask.shape) != tuple(candidate.mask.shape) + or _bbox_iou(existing.bbox, candidate.bbox) < 0.98 + ): + continue + union = torch.logical_or(existing.mask, candidate.mask).sum() + if int(union.item()) == 0: + continue + intersection = torch.logical_and(existing.mask, candidate.mask).sum() + if float(intersection.item()) / float(union.item()) >= 0.98: + return index + return None + + +def _build_segment( + *, + image: torch.Tensor, + candidate: _GuideMaskCandidate, + source_height: int, + source_width: int, + confidence: float, + label: str, +) -> Segment: + """Build one crop-local SEG without ever expanding a full source mask.""" + + crop_region = _source_crop_region( + candidate.bbox, + guide_height=int(candidate.mask.shape[0]), + guide_width=int(candidate.mask.shape[1]), + source_height=source_height, + source_width=source_width, + ) + guide_crop = candidate.mask[ + candidate.bbox.top : candidate.bbox.bottom, + candidate.bbox.left : candidate.bbox.right, + ] + local_mask = _resize_local_mask( + guide_crop, + height=crop_region.height, + width=crop_region.width, + ) + return Segment( + cropped_image=image[ + :, + crop_region.top : crop_region.bottom, + crop_region.left : crop_region.right, + :, + ] + .detach() + .clone(), + cropped_mask=local_mask.detach().clone(), + confidence=max(0.0, min(1.0, float(confidence))), + crop_region=crop_region, + bbox=BoundingBox( + crop_region.left, + crop_region.top, + crop_region.right, + crop_region.bottom, + ), + label=label, + ) + + +def _bbox_from_mask(mask: torch.Tensor) -> BoundingBox | None: + """Return the tight bounding box for one active guide-space mask.""" + + y_coords, x_coords = torch.where(mask) + if y_coords.numel() == 0: + return None + return BoundingBox( + left=int(x_coords.min().item()), + top=int(y_coords.min().item()), + right=int(x_coords.max().item()) + 1, + bottom=int(y_coords.max().item()) + 1, + ) + + +def _bbox_iou(first: BoundingBox, second: BoundingBox) -> float: + """Return the intersection-over-union of two rectangular bounds.""" + + overlap_width = max( + 0, min(first.right, second.right) - max(first.left, second.left) + ) + overlap_height = max( + 0, min(first.bottom, second.bottom) - max(first.top, second.top) + ) + intersection = overlap_width * overlap_height + union = first.width * first.height + second.width * second.height - intersection + return 0.0 if union == 0 else intersection / float(union) + + +def _source_crop_region( + guide_bbox: BoundingBox, + *, + guide_height: int, + guide_width: int, + source_height: int, + source_width: int, +) -> CropRegion: + """Map a guide-space bounding box to a conservative source-image crop.""" + + left = _map_lower_bound(guide_bbox.left, guide_width, source_width) + top = _map_lower_bound(guide_bbox.top, guide_height, source_height) + right = _map_upper_bound(guide_bbox.right, guide_width, source_width) + bottom = _map_upper_bound(guide_bbox.bottom, guide_height, source_height) + return CropRegion(left, top, right, bottom) + + +def _map_lower_bound(value: int, guide_limit: int, source_limit: int) -> int: + """Map a guide coordinate down while keeping it within source bounds.""" + + return min(source_limit - 1, max(0, floor(value * source_limit / guide_limit))) + + +def _map_upper_bound(value: int, guide_limit: int, source_limit: int) -> int: + """Map a guide coordinate up while keeping a non-empty source extent.""" + + return max(1, min(source_limit, ceil(value * source_limit / guide_limit))) + + +def _resize_local_mask(mask: torch.Tensor, *, height: int, width: int) -> torch.Tensor: + """Scale only a retained mask's tight crop into source pixel coordinates.""" + + return ( + functional.interpolate( + mask.unsqueeze(0).unsqueeze(0).float(), + size=(height, width), + mode="nearest", + ) + .squeeze(0) + .squeeze(0) + .float() + ) diff --git a/simple_syrup/services/tiled_diffusion_sampling_service.py b/simple_syrup/services/tiled_diffusion_sampling_service.py index fc3b177..a0fe716 100644 --- a/simple_syrup/services/tiled_diffusion_sampling_service.py +++ b/simple_syrup/services/tiled_diffusion_sampling_service.py @@ -11,7 +11,9 @@ from typing import Any import torch from ..domain.conditioning_batch import ConditioningBatch, select_conditioning -from ..domain.tiled_diffusion import validate_tiled_diffusion_mode +from ..domain.segs import coerce_segs_group +from ..domain.segs_tiled_diffusion import build_segs_guided_tiled_diffusion_plan +from ..domain.tiled_diffusion import TiledDiffusionPlan, validate_tiled_diffusion_mode from ..runtime import mixture_of_diffusers_sampling, multidiffusion_sampling from ..runtime.detail_previews import DetailPreviewContext @@ -42,10 +44,33 @@ class TiledDiffusionSamplingService: preview_context: DetailPreviewContext | None = None, differential_diffusion: bool = False, allow_full_context_masks: bool = False, + segs: object | None = None, ) -> Latent: """Sample a latent with the selected tiled diffusion method.""" validate_tiled_diffusion_mode(diffusion_mode) + if segs is not None: + return self._sample_segs_guided( + diffusion_mode=diffusion_mode, + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=positive, + negative=negative, + latent_image=latent_image, + denoise=denoise, + latent_tile_width=latent_tile_width, + latent_tile_height=latent_tile_height, + latent_tile_overlap=latent_tile_overlap, + latent_tile_batch_size=latent_tile_batch_size, + preview_context=preview_context, + differential_diffusion=differential_diffusion, + allow_full_context_masks=allow_full_context_masks, + segs=segs, + ) if self._uses_conditioning_batch(positive, negative): return self._sample_conditioning_batch( diffusion_mode=diffusion_mode, @@ -67,6 +92,52 @@ class TiledDiffusionSamplingService: differential_diffusion=differential_diffusion, allow_full_context_masks=allow_full_context_masks, ) + return self._sample_single( + diffusion_mode=diffusion_mode, + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=positive, + negative=negative, + latent_image=latent_image, + denoise=denoise, + latent_tile_width=latent_tile_width, + latent_tile_height=latent_tile_height, + latent_tile_overlap=latent_tile_overlap, + latent_tile_batch_size=latent_tile_batch_size, + preview_context=preview_context, + differential_diffusion=differential_diffusion, + allow_full_context_masks=allow_full_context_masks, + ) + + def _sample_single( + self, + *, + diffusion_mode: str, + model: Any, + seed: int, + steps: int, + cfg: float, + sampler_name: str, + scheduler: str, + positive: Any, + negative: Any, + latent_image: Latent, + denoise: float, + latent_tile_width: int, + latent_tile_height: int, + latent_tile_overlap: int, + latent_tile_batch_size: int, + preview_context: DetailPreviewContext | None, + differential_diffusion: bool, + allow_full_context_masks: bool, + tiled_plan: TiledDiffusionPlan | None = None, + ) -> Latent: + """Route one single-latent tiled sample to its selected runtime.""" + if diffusion_mode == "multidiffusion": return multidiffusion_sampling.sample_multidiffusion( model=model, @@ -86,6 +157,7 @@ class TiledDiffusionSamplingService: preview_context=preview_context, differential_diffusion=differential_diffusion, allow_full_context_masks=allow_full_context_masks, + tiled_plan=tiled_plan, ) return mixture_of_diffusers_sampling.sample_mixture_of_diffusers( model=model, @@ -105,8 +177,100 @@ class TiledDiffusionSamplingService: preview_context=preview_context, differential_diffusion=differential_diffusion, allow_full_context_masks=allow_full_context_masks, + tiled_plan=tiled_plan, ) + def _sample_segs_guided( + self, + *, + diffusion_mode: str, + model: Any, + seed: int, + steps: int, + cfg: float, + sampler_name: str, + scheduler: str, + positive: Any, + negative: Any, + latent_image: Latent, + denoise: float, + latent_tile_width: int, + latent_tile_height: int, + latent_tile_overlap: int, + latent_tile_batch_size: int, + preview_context: DetailPreviewContext | None, + differential_diffusion: bool, + allow_full_context_masks: bool, + segs: object, + ) -> Latent: + """Sample every latent batch item using its connected SEGS guide.""" + + segs_group = coerce_segs_group(segs) + latent_samples = latent_image.get("samples") + if not isinstance(latent_samples, torch.Tensor): + raise TypeError("Tiled diffusion latent samples must be a torch.Tensor.") + batch_size = int(latent_samples.shape[0]) + if len(segs_group) not in (1, batch_size): + raise ValueError( + "SEGS-guided tiled diffusion requires one SEGS payload or one per " + f"latent batch item; received {len(segs_group)} SEGS payloads for " + f"batch size {batch_size}." + ) + + outputs: list[torch.Tensor] = [] + for index in range(batch_size): + item_latent = self._single_item_latent(latent_image, index) + samples = item_latent["samples"] + if not isinstance(samples, torch.Tensor): + raise TypeError( + "Tiled diffusion latent samples must be a torch.Tensor." + ) + segs_for_item = segs_group[0 if len(segs_group) == 1 else index] + plan = build_segs_guided_tiled_diffusion_plan( + segs=segs_for_item, + latent_width=int(samples.shape[-1]), + latent_height=int(samples.shape[-2]), + tile_width=latent_tile_width, + tile_height=latent_tile_height, + overlap=latent_tile_overlap, + tile_batch_size=latent_tile_batch_size, + ) + output = self._sample_single( + diffusion_mode=diffusion_mode, + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=select_conditioning(positive, index) + if isinstance(positive, ConditioningBatch) + else positive, + negative=select_conditioning(negative, index) + if isinstance(negative, ConditioningBatch) + else negative, + latent_image=item_latent, + denoise=denoise, + latent_tile_width=latent_tile_width, + latent_tile_height=latent_tile_height, + latent_tile_overlap=latent_tile_overlap, + latent_tile_batch_size=latent_tile_batch_size, + preview_context=preview_context, + differential_diffusion=differential_diffusion, + allow_full_context_masks=allow_full_context_masks, + tiled_plan=plan, + ) + output_samples = output["samples"] + if not isinstance(output_samples, torch.Tensor): + raise TypeError( + "Tiled diffusion output samples must be a torch.Tensor." + ) + outputs.append(output_samples) + result = latent_image.copy() + result.pop("downscale_ratio_spacial", None) + result["samples"] = torch.cat(outputs, dim=0) + return result + def _sample_conditioning_batch( self, *, diff --git a/tests/test_ksampler_tiled_diffusion_node.py b/tests/test_ksampler_tiled_diffusion_node.py index bd771b5..46d62c9 100644 --- a/tests/test_ksampler_tiled_diffusion_node.py +++ b/tests/test_ksampler_tiled_diffusion_node.py @@ -27,6 +27,7 @@ def test_input_types_match_tiled_diffusion_contract( lambda: ("normal",), ) required = KSamplerTiledDiffusion.INPUT_TYPES()["required"] + optional = KSamplerTiledDiffusion.INPUT_TYPES()["optional"] assert tuple(required) == ( "model", @@ -58,6 +59,7 @@ def test_input_types_match_tiled_diffusion_contract( assert required["latent_tile_height"][1]["max"] == 512 assert required["latent_tile_overlap"][1]["default"] == 16 assert required["latent_tile_batch_size"][1]["default"] == 4 + assert optional["segs"][0] == "SEGS" def test_node_metadata_matches_contract() -> None: @@ -118,6 +120,7 @@ def test_sample_delegates_to_shared_service( assert call["latent_tile_overlap"] == 24 assert call["latent_tile_batch_size"] == 3 assert call["preview_context"] is None + assert call["segs"] is None def test_invalid_diffusion_mode_fails_before_runtime_sampling() -> None: @@ -171,6 +174,7 @@ class _FakeTiledDiffusionSamplingService: latent_tile_overlap: int, latent_tile_batch_size: int, preview_context: Any | None = None, + segs: object | None = None, ) -> dict[str, Any]: """Record sampling arguments and return a fixed latent.""" @@ -192,6 +196,7 @@ class _FakeTiledDiffusionSamplingService: "latent_tile_overlap": latent_tile_overlap, "latent_tile_batch_size": latent_tile_batch_size, "preview_context": preview_context, + "segs": segs, } ) return self.output diff --git a/tests/test_registration.py b/tests/test_registration.py index bbd36c8..a127c4f 100644 --- a/tests/test_registration.py +++ b/tests/test_registration.py @@ -44,6 +44,7 @@ BASE_NODE_IDS = [ "SimpleSyrup.PromptSEGSWithSAM", "SimpleSyrup.ResizeImageToTarget", "SimpleSyrup.SAMModelLoader", + "SimpleSyrup.SEGSFromSAMOutput", "SimpleSyrup.ScaleFactor", "SimpleSyrup.Seed", "SimpleSyrup.SimpleLoadAnima", diff --git a/tests/test_sam_automatic_segmenter.py b/tests/test_sam_automatic_segmenter.py new file mode 100644 index 0000000..94c3599 --- /dev/null +++ b/tests/test_sam_automatic_segmenter.py @@ -0,0 +1,131 @@ +# 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 unprompted SAM-family automatic segmentation runtime adapters.""" + +from __future__ import annotations + +import sys +from types import ModuleType + +import pytest +import torch + +from simple_syrup.runtime.sam_automatic_segmenter import SAMModelAutomaticSegmenter + + +def test_fast_sam_adapter_extracts_everything_masks_and_confidences() -> None: + """FastSAM returns its direct mask results without text or CLIP prompting.""" + + class _Masks: + """Expose Ultralytics-style mask data.""" + + data = torch.tensor( + [ + [[1.0, 0.0], [0.0, 1.0]], + [[0.0, 1.0], [1.0, 0.0]], + ] + ) + + class _Boxes: + """Expose Ultralytics-style detection confidences.""" + + conf = torch.tensor([0.75, 0.5]) + + class _Result: + """Expose one Ultralytics segmentation result.""" + + masks = _Masks() + boxes = _Boxes() + + class FastSAM: + """Minimal FastSAM-like external model.""" + + def __init__(self) -> None: + """Create an empty prediction call recorder.""" + + self.calls: list[dict[str, object]] = [] + + def to(self, device: torch.device) -> None: + """Accept temporary CPU inference movement.""" + + del device + + def eval(self) -> None: + """Accept evaluation mode.""" + + def predict(self, image: object, **kwargs: object) -> list[_Result]: + """Record direct everything-mode prediction options.""" + + self.calls.append({"image": image, **kwargs}) + return [_Result()] + + model = FastSAM() + + masks = SAMModelAutomaticSegmenter().segment_all( + model, + torch.zeros((1, 2, 2, 3)), + execution_device="cpu", + ) + + assert [mask.confidence for mask in masks] == [0.75, 0.5] + assert torch.equal(masks[0].mask, _Masks.data[0]) + assert model.calls[0]["retina_masks"] is True + assert model.calls[0]["verbose"] is False + assert model.calls[0]["device"] == "cpu" + + +def test_segment_anything_adapter_preserves_predicted_iou( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Standard SAM automatic generation maps quality metadata to SEGS confidence.""" + + class _Generator: + """Return a fixed native automatic-mask response.""" + + def __init__(self, model: object) -> None: + """Accept the raw SAM model.""" + + del model + + def generate(self, image: object) -> list[dict[str, object]]: + """Return one native SAM automatic mask.""" + + del image + return [ + { + "segmentation": torch.tensor([[True, False], [False, True]]), + "predicted_iou": 0.92, + } + ] + + segment_anything = ModuleType("segment_anything") + segment_anything.SamAutomaticMaskGenerator = _Generator # type: ignore[attr-defined] + monkeypatch.setitem(sys.modules, "segment_anything", segment_anything) + + class _SAM: + """Minimal raw SAM model.""" + + model_name = "sam_vit_b.pth" + + def to(self, device: torch.device) -> None: + """Accept temporary CPU inference movement.""" + + del device + + def eval(self) -> None: + """Accept evaluation mode.""" + + masks = SAMModelAutomaticSegmenter().segment_all( + _SAM(), + torch.zeros((1, 2, 2, 3)), + execution_device="cpu", + ) + + assert len(masks) == 1 + assert masks[0].confidence == 0.92 + assert torch.equal( + masks[0].mask, + torch.tensor([[1.0, 0.0], [0.0, 1.0]]), + ) diff --git a/tests/test_sam_loader.py b/tests/test_sam_loader.py index c1bf50a..3343dd7 100644 --- a/tests/test_sam_loader.py +++ b/tests/test_sam_loader.py @@ -212,6 +212,45 @@ def test_sam_loader_loads_sam_hq_from_owned_runtime( assert state.checkpoints == [str(tmp_path / "sams" / "sam_hq_vit_b.pth")] +def test_sam_loader_downloads_and_loads_fast_sam_s( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """FastSAM-s uses the existing SAM download, cache, and model wrapper path.""" + + downloader = RecordingDownloader() + constructed: list[str] = [] + + class FakeFastSAM: + """Minimal FastSAM-compatible model fake.""" + + def __init__(self, checkpoint: str) -> None: + """Record the checkpoint passed by the shared SAM loader.""" + + constructed.append(checkpoint) + + def to(self, device: object) -> None: + """Accept device management.""" + + def eval(self) -> None: + """Accept evaluation mode.""" + + ultralytics = ModuleType("ultralytics") + ultralytics.FastSAM = FakeFastSAM # type: ignore[attr-defined] + monkeypatch.setitem(sys.modules, "ultralytics", ultralytics) + + loaded = SAMLoaderService( + downloader=downloader, # type: ignore[arg-type] + folder_paths_module=FakeFolderPaths(tmp_path), + ).load_model("FastSAM-s (23MB)", auto_download=True) + + expected = str(tmp_path / "sams" / "FastSAM-s.pt") + assert downloader.requests[0].destination_path == Path(expected) + assert constructed == [expected] + assert loaded.model_id == "fast_sam_s" + assert loaded.managed_model is not None + + def test_sam_loader_errors_when_missing_and_download_disabled(tmp_path: Path) -> None: """SAM loader fails clearly when downloads are disabled.""" diff --git a/tests/test_sam_model_loader_node.py b/tests/test_sam_model_loader_node.py index c04334d..6ac3030 100644 --- a/tests/test_sam_model_loader_node.py +++ b/tests/test_sam_model_loader_node.py @@ -31,6 +31,7 @@ def test_sam_model_loader_declares_expected_inputs() -> None: assert set(required) == {"sam_model"} assert "sam_vit_b (375MB)" in required["sam_model"][0] + assert "FastSAM-s (23MB)" in required["sam_model"][0] def test_sam_model_loader_uses_settings_aware_choices() -> None: diff --git a/tests/test_segs_from_sam_output_node.py b/tests/test_segs_from_sam_output_node.py new file mode 100644 index 0000000..b8129c0 --- /dev/null +++ b/tests/test_segs_from_sam_output_node.py @@ -0,0 +1,97 @@ +# 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 SEGS from SAM Output ComfyUI node.""" + +from __future__ import annotations + +import pytest +import torch + +from simple_syrup.nodes.segs_from_sam_output import SEGSFromSAMOutput + + +def test_node_declares_automatic_sam_to_segs_contract() -> None: + """The node exposes standard IMAGE, SAM_MODEL, and SEGS sockets.""" + + inputs = SEGSFromSAMOutput.INPUT_TYPES()["required"] + + assert tuple(inputs) == ( + "image", + "sam_model", + "segmentation_resolution", + "minimum_region_area", + ) + assert inputs["segmentation_resolution"][1]["default"] == 640 + assert inputs["segmentation_resolution"][1]["step"] == 64 + assert SEGSFromSAMOutput.RETURN_TYPES == ("SEGS",) + assert SEGSFromSAMOutput.OUTPUT_IS_LIST == (True,) + + +def test_node_builds_one_segs_output_per_image_batch_item( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Automatic segmentation remains aligned with ComfyUI IMAGE batches.""" + + calls: list[torch.Tensor] = [] + phase_progress = _RecordingPhaseProgress() + + class _Service: + """Record image batches and return a distinct opaque SEGS payload.""" + + def build(self, **kwargs: object) -> object: + """Record one source image and return it as an opaque marker.""" + + image = kwargs["image"] + assert isinstance(image, torch.Tensor) + calls.append(image) + progress = kwargs["phase_progress"] + assert progress is phase_progress + for phase in ( + "preparing_segmentation_image", + "generating_automatic_masks", + "building_segs", + ): + phase_progress.advance(phase) + return (image.shape[1:3], ()) + + monkeypatch.setattr(SEGSFromSAMOutput, "service_class", _Service) + monkeypatch.setattr( + SEGSFromSAMOutput, + "progress_factory", + lambda **_kwargs: phase_progress, + ) + + (segs,) = SEGSFromSAMOutput().generate( + image=torch.zeros((2, 16, 16, 3)), + sam_model=object(), + segmentation_resolution=640, + minimum_region_area=0, + ) + + assert len(calls) == 2 + assert len(segs) == 2 + assert phase_progress.phases == [ + "preparing_segmentation_image", + "generating_automatic_masks", + "building_segs", + "preparing_segmentation_image", + "generating_automatic_masks", + "building_segs", + "completed", + ] + + +class _RecordingPhaseProgress: + """Record node progress without constructing a real ComfyUI progress bar.""" + + def __init__(self) -> None: + """Create an empty phase record.""" + + self.phases: list[str] = [] + + def advance(self, phase: str) -> None: + """Record one phase transition.""" + + self.phases.append(phase) diff --git a/tests/test_segs_from_sam_output_service.py b/tests/test_segs_from_sam_output_service.py new file mode 100644 index 0000000..5be626f --- /dev/null +++ b/tests/test_segs_from_sam_output_service.py @@ -0,0 +1,180 @@ +# 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 automatic SAM output conversion into standard SEGS.""" + +from __future__ import annotations + +from typing import cast + +import pytest +import torch + +from simple_syrup.runtime.sam_automatic_segmenter import AutomaticSAMMask +from simple_syrup.services.segs_from_sam_output_service import SEGSFromSAMOutputService + + +class _RecordingSegmenter: + """Return configured guide-space masks while recording the input shape.""" + + def __init__(self, masks: tuple[AutomaticSAMMask, ...]) -> None: + """Store masks for a deterministic automatic segmentation response.""" + + self._masks = masks + self.image_shapes: list[tuple[int, int]] = [] + + def segment_all( + self, + sam_model: object, + image: torch.Tensor, + execution_device: str = "auto", + ) -> tuple[AutomaticSAMMask, ...]: + """Record the guide image and return configured masks.""" + + del sam_model, execution_device + self.image_shapes.append((int(image.shape[1]), int(image.shape[2]))) + return self._masks + + +class _RecordingPhaseProgress: + """Record automatic-SEGS phase transitions without a ComfyUI dependency.""" + + def __init__(self) -> None: + """Create an empty phase record.""" + + self.phases: list[str] = [] + + def advance(self, phase: str) -> None: + """Record one named phase.""" + + self.phases.append(phase) + + +def test_service_downscales_the_segmentation_guide_without_upscaling_source() -> None: + """Guide resolution caps the long edge while preserving source SEGS geometry.""" + + guide_mask = torch.zeros((32, 64), dtype=torch.float32) + guide_mask[8:16, 16:32] = 1.0 + runtime = _RecordingSegmenter((AutomaticSAMMask(guide_mask, 0.8),)) + service = SEGSFromSAMOutputService(runtime) + + segs = service.build( + image=torch.zeros((1, 128, 256, 3)), + sam_model=object(), + segmentation_resolution=64, + minimum_region_area=0, + ) + + assert runtime.image_shapes == [(32, 64)] + assert segs[0] == (128, 256) + segment = segs[1][0] + assert segment.bbox == (64, 32, 128, 64) + assert segment.crop_region == (64, 32, 128, 64) + assert cast(torch.Tensor, segment.cropped_mask).shape == (32, 64) + assert segment.confidence == 0.8 + assert segment.label == "segment_001" + + +def test_service_reports_meaningful_automatic_segmentation_phases() -> None: + """The service exposes each expensive step to the node-owned progress bar.""" + + runtime = _RecordingSegmenter((AutomaticSAMMask(torch.ones((16, 16)), 1.0),)) + reporter = _RecordingPhaseProgress() + + SEGSFromSAMOutputService(runtime).build( + image=torch.zeros((1, 16, 16, 3)), + sam_model=object(), + segmentation_resolution=64, + minimum_region_area=0, + phase_progress=reporter, + ) + + assert reporter.phases == [ + "preparing_segmentation_image", + "generating_automatic_masks", + "building_segs", + ] + + +def test_service_filters_region_area_after_restoring_source_dimensions() -> None: + """Minimum area uses the original image instead of the guide image size.""" + + guide_mask = torch.zeros((32, 64), dtype=torch.float32) + guide_mask[8:10, 8:10] = 1.0 + runtime = _RecordingSegmenter((AutomaticSAMMask(guide_mask, 0.5, "thing"),)) + service = SEGSFromSAMOutputService(runtime) + + retained = service.build( + image=torch.zeros((1, 128, 256, 3)), + sam_model=object(), + segmentation_resolution=64, + minimum_region_area=63, + ) + filtered = service.build( + image=torch.zeros((1, 128, 256, 3)), + sam_model=object(), + segmentation_resolution=64, + minimum_region_area=65, + ) + + assert retained[1][0].label == "thing" + assert filtered[1] == () + + +def test_service_suppresses_duplicate_masks_and_keeps_highest_confidence() -> None: + """Repeated automatic masks do not create duplicate detailer targets.""" + + mask = torch.ones((16, 16), dtype=torch.float32) + runtime = _RecordingSegmenter( + ( + AutomaticSAMMask(mask, 0.4, "first"), + AutomaticSAMMask(mask, 0.9, "second"), + ) + ) + + segs = SEGSFromSAMOutputService(runtime).build( + image=torch.zeros((1, 16, 16, 3)), + sam_model=object(), + segmentation_resolution=64, + minimum_region_area=0, + ) + + assert len(segs[1]) == 1 + assert segs[1][0].confidence == 0.9 + assert segs[1][0].label == "second" + + +def test_service_expands_retained_mask_crops_to_source_resolution() -> None: + """Source-size SEGS creation never materializes a full-image mask per region.""" + + guide_mask = torch.zeros((64, 128), dtype=torch.float32) + guide_mask[16:32, 32:64] = 1.0 + + segs = SEGSFromSAMOutputService( + _RecordingSegmenter((AutomaticSAMMask(guide_mask, 1.0),)) + ).build( + image=torch.zeros((1, 1024, 2048, 3)), + sam_model=object(), + segmentation_resolution=128, + minimum_region_area=0, + ) + + segment = segs[1][0] + assert segment.crop_region == (512, 256, 1024, 512) + assert cast(torch.Tensor, segment.cropped_mask).shape == (256, 512) + + +@pytest.mark.parametrize("resolution", (0, 65, 127)) +def test_service_rejects_non_64_step_segmentation_resolution(resolution: int) -> None: + """Service validation preserves the node's 64-pixel resolution contract.""" + + service = SEGSFromSAMOutputService(_RecordingSegmenter(())) + + with pytest.raises(ValueError, match="segmentation_resolution"): + service.build( + image=torch.zeros((1, 16, 16, 3)), + sam_model=object(), + segmentation_resolution=resolution, + minimum_region_area=0, + ) diff --git a/tests/test_segs_tiled_diffusion.py b/tests/test_segs_tiled_diffusion.py new file mode 100644 index 0000000..ca35e3f --- /dev/null +++ b/tests/test_segs_tiled_diffusion.py @@ -0,0 +1,189 @@ +# 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 SEGS-guided irregular tiled diffusion planning.""" + +from __future__ import annotations + +import pytest +import torch + +from simple_syrup.domain.segs import BoundingBox, CropRegion, Segment +from simple_syrup.domain.segs_tiled_diffusion import ( + _segment_mask_to_latent, + build_segs_guided_tiled_diffusion_plan, +) + + +def test_guided_plan_splits_oversized_region_with_bounded_rectangular_windows() -> None: + """A large region becomes multiple bounded cores instead of one giant tile.""" + + plan = build_segs_guided_tiled_diffusion_plan( + segs=_full_mask_segs(64, 64), + latent_width=64, + latent_height=64, + tile_width=32, + tile_height=32, + overlap=8, + tile_batch_size=3, + ) + + assert len(plan.tiles) > 1 + assert all(tile.width == 32 and tile.height == 32 for tile in plan.tiles) + assert all(tile.weight_mask is not None for tile in plan.tiles) + coverage = torch.zeros((64, 64), dtype=torch.float32) + for tile in plan.tiles: + assert tile.weight_mask is not None + coverage[tile.y : tile.y + tile.height, tile.x : tile.x + tile.width] += ( + tile.weight_mask + ) + assert bool((coverage > 0).all()) + + +def test_guided_plan_feathers_the_configured_overlap_between_irregular_cores() -> None: + """Non-zero overlap produces fractional shared ownership weights.""" + + plan = build_segs_guided_tiled_diffusion_plan( + segs=_half_mask_segs(64, 64), + latent_width=64, + latent_height=64, + tile_width=32, + tile_height=32, + overlap=8, + tile_batch_size=4, + ) + + weights = [tile.weight_mask for tile in plan.tiles if tile.weight_mask is not None] + assert any(bool(torch.any((weight > 0) & (weight < 1))) for weight in weights) + + +def test_guided_plan_prefers_smaller_overlapping_segs_for_ownership() -> None: + """A nested small SEG remains an ownership core instead of being swallowed.""" + + full = torch.ones((32, 32), dtype=torch.float32) + small = torch.zeros((32, 32), dtype=torch.float32) + small[8:24, 8:24] = 1.0 + segs = ( + (32, 32), + ( + _segment(full, "large", 0.6), + _segment(small, "small", 0.9), + ), + ) + + plan = build_segs_guided_tiled_diffusion_plan( + segs=segs, + latent_width=32, + latent_height=32, + tile_width=32, + tile_height=32, + overlap=0, + tile_batch_size=4, + ) + + assert len(plan.tiles) == 2 + assert all(tile.weight_mask is not None for tile in plan.tiles) + + +def test_guided_plan_merges_small_cores_without_comparing_tensor_values() -> None: + """Small nearby cores merge through their indexes instead of tensor equality.""" + + first = torch.zeros((32, 32), dtype=torch.float32) + second = torch.zeros((32, 32), dtype=torch.float32) + first[4:8, 4:8] = 1.0 + second[8:12, 8:12] = 1.0 + segs = ( + (32, 32), + ( + _segment(first, "first", 1.0), + _segment(second, "second", 1.0), + ), + ) + + plan = build_segs_guided_tiled_diffusion_plan( + segs=segs, + latent_width=32, + latent_height=32, + tile_width=32, + tile_height=32, + overlap=0, + tile_batch_size=4, + ) + + assert len(plan.tiles) == 1 + + +def test_crop_local_mask_projects_directly_to_latent_space() -> None: + """A small crop maps without materializing a full source-resolution mask.""" + + crop = CropRegion(512, 256, 768, 512) + segment = Segment( + cropped_image=None, + cropped_mask=torch.ones((256, 256), dtype=torch.float32), + confidence=1.0, + crop_region=crop, + bbox=BoundingBox(*crop), + label="small_region", + ) + + latent_mask = _segment_mask_to_latent( + segment, + source_height=4096, + source_width=4096, + latent_width=512, + latent_height=512, + ) + + assert int(latent_mask.sum().item()) == 32 * 32 + assert bool(latent_mask[32:64, 64:96].all()) + assert not bool(latent_mask[:32].any()) + assert not bool(latent_mask[:, :64].any()) + + +def test_guided_plan_rejects_mismatched_image_aspect_ratio() -> None: + """SEGS from a different image fail before tiled sampling begins.""" + + with pytest.raises(ValueError, match="aspect ratio"): + build_segs_guided_tiled_diffusion_plan( + segs=_full_mask_segs(16, 32), + latent_width=32, + latent_height=32, + tile_width=16, + tile_height=16, + overlap=4, + tile_batch_size=2, + ) + + +def _full_mask_segs( + height: int, width: int +) -> tuple[tuple[int, int], tuple[Segment, ...]]: + """Return one full-image SEG payload.""" + + return ((height, width), (_segment(torch.ones((height, width)), "region", 1.0),)) + + +def _half_mask_segs( + height: int, width: int +) -> tuple[tuple[int, int], tuple[Segment, ...]]: + """Return one left-half SEG payload with implicit background ownership.""" + + mask = torch.zeros((height, width), dtype=torch.float32) + mask[:, : width // 2] = 1.0 + return ((height, width), (_segment(mask, "left", 1.0),)) + + +def _segment(mask: torch.Tensor, label: str, confidence: float) -> Segment: + """Build one full-image Segment whose crop equals the source dimensions.""" + + height, width = mask.shape + region = CropRegion(0, 0, width, height) + return Segment( + cropped_image=None, + cropped_mask=mask, + confidence=confidence, + crop_region=region, + bbox=BoundingBox(0, 0, width, height), + label=label, + ) diff --git a/tests/test_tiled_diffusion_sampling_service.py b/tests/test_tiled_diffusion_sampling_service.py index fcc002b..5a85164 100644 --- a/tests/test_tiled_diffusion_sampling_service.py +++ b/tests/test_tiled_diffusion_sampling_service.py @@ -12,6 +12,7 @@ import pytest import torch from simple_syrup.domain.conditioning_batch import ConditioningBatch +from simple_syrup.domain.segs import BoundingBox, CropRegion, Segment from simple_syrup.services.tiled_diffusion_sampling_service import ( TiledDiffusionSamplingService, ) @@ -127,7 +128,7 @@ def test_service_forwards_sampling_arguments_unchanged( assert result is output assert calls == { key: value for key, value in kwargs.items() if key != "diffusion_mode" - } + } | {"tiled_plan": None} def test_service_forwards_differential_diffusion_request( @@ -216,6 +217,38 @@ def test_service_selects_conditioning_batch_per_latent_item( assert torch.equal(result["samples"][1], torch.full((4, 4, 4), 2.0)) +def test_service_builds_and_forwards_a_segs_guided_plan( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Connected SEGS replace the regular grid with a semantic tile plan.""" + + calls: list[dict[str, Any]] = [] + + def fake_multidiffusion(**kwargs: Any) -> dict[str, Any]: + """Record the semantic plan and return the item unchanged.""" + + calls.append(kwargs) + return {"samples": kwargs["latent_image"]["samples"]} + + monkeypatch.setattr( + "simple_syrup.services.tiled_diffusion_sampling_service." + "multidiffusion_sampling.sample_multidiffusion", + fake_multidiffusion, + ) + segs = ((4, 4), (_full_segment(4, 4),)) + + TiledDiffusionSamplingService().sample( + **(_sample_kwargs(diffusion_mode="multidiffusion") | {"segs": segs}) + ) + + assert len(calls) == 1 + plan = calls[0]["tiled_plan"] + assert plan is not None + assert plan.latent_width == 4 + assert plan.latent_height == 4 + assert all(tile.weight_mask is not None for tile in plan.tiles) + + def test_invalid_mode_fails_before_runtime_call( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -271,3 +304,17 @@ def _sample_kwargs( "differential_diffusion": False, "allow_full_context_masks": False, } + + +def _full_segment(height: int, width: int) -> Segment: + """Return one full-image SEG compatible with the sample latent dimensions.""" + + region = CropRegion(0, 0, width, height) + return Segment( + cropped_image=None, + cropped_mask=torch.ones((height, width), dtype=torch.float32), + confidence=1.0, + crop_region=region, + bbox=BoundingBox(0, 0, width, height), + label="region", + ) diff --git a/tests/test_tiled_sampling_runtime.py b/tests/test_tiled_sampling_runtime.py index d78a398..36fbe78 100644 --- a/tests/test_tiled_sampling_runtime.py +++ b/tests/test_tiled_sampling_runtime.py @@ -185,6 +185,30 @@ def test_new_spatial_weight_buffer_broadcasts_over_spatial_axes() -> None: ).shape == (1, 1, 1, 4, 8) +def test_semantic_tile_weight_cache_reuses_resident_weights() -> None: + """Semantic tile weights are materialized once for repeated model outputs.""" + + tile = LatentTile( + x=0, + y=0, + width=4, + height=4, + weight_mask=torch.ones((4, 4), dtype=torch.float32), + ) + cache = tiled_sampling.SemanticTileWeightCache((tile,)) + output = torch.zeros((1, 4, 4, 4), dtype=torch.float16) + + first = cache.for_output(output) + second = cache.for_output(output) + model_weight, accumulation_weight = cache.for_tile(first, tile) + + assert first is second + assert model_weight is first.model[0] + assert accumulation_weight is first.accumulation[0] + assert model_weight.dtype == torch.float16 + assert accumulation_weight.dtype == torch.float32 + + def test_contains_unsupported_conditioning_key_finds_nested_values() -> None: """Unsupported regional and control keys are detected recursively."""