From 1972e452dc894d769018ceb2ab4e7858c4c54345 Mon Sep 17 00:00:00 2001 From: Artificial Sweetener Date: Sun, 2 Aug 2026 01:15:47 -0400 Subject: [PATCH] feat(segmentation): add SAM region overlay --- simple_syrup/nodes/segs_from_sam_output.py | 34 +++-- .../runtime/sam_region_overlay_renderer.py | 142 ++++++++++++++++++ .../services/segs_from_sam_output_service.py | 32 +++- tests/test_sam_region_overlay_renderer.py | 63 ++++++++ tests/test_segs_from_sam_output_node.py | 26 +++- tests/test_segs_from_sam_output_service.py | 24 +-- 6 files changed, 286 insertions(+), 35 deletions(-) create mode 100644 simple_syrup/runtime/sam_region_overlay_renderer.py create mode 100644 tests/test_sam_region_overlay_renderer.py diff --git a/simple_syrup/nodes/segs_from_sam_output.py b/simple_syrup/nodes/segs_from_sam_output.py index c6d1e80..2ec3e19 100644 --- a/simple_syrup/nodes/segs_from_sam_output.py +++ b/simple_syrup/nodes/segs_from_sam_output.py @@ -9,6 +9,8 @@ from __future__ import annotations from collections.abc import Callable from typing import Any, ClassVar +import torch + 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 @@ -22,11 +24,12 @@ class SEGSFromSAMOutput: create_comfy_phase_progress ) - RETURN_TYPES = ("SEGS",) - RETURN_NAMES = ("segs",) - OUTPUT_IS_LIST = (True,) + RETURN_TYPES = ("SEGS", "IMAGE") + RETURN_NAMES = ("segs", "overlay") + OUTPUT_IS_LIST = (True, False) OUTPUT_TOOLTIPS = ( "Automatic image regions as SEGS for detailing, masking, or tiled diffusion.", + "Source images with retained SAM regions shown as translucent colors.", ) FUNCTION = "generate" CATEGORY = "SimpleSyrup/Detection" @@ -82,33 +85,34 @@ class SEGSFromSAMOutput: 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.""" + ) -> tuple[list[object], torch.Tensor]: + """Return aligned automatic SEGS and a SAM-style overlay image batch.""" 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, + total_phases=int(image_batch.shape[0]) * 4 + 1, ) outputs: list[object] = [] + overlays: list[torch.Tensor] = [] 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, - ) + result = service.build( + image=single_image, + sam_model=sam_model, + segmentation_resolution=segmentation_resolution, + minimum_region_area=minimum_region_area, + phase_progress=phase_progress, ) + outputs.append(result.segs) + overlays.append(result.overlay) except Exception: phase_progress.advance("failed") raise phase_progress.advance("completed") - return (outputs,) + return outputs, torch.cat(overlays, dim=0) def _sam_model_subject(sam_model: object) -> str: diff --git a/simple_syrup/runtime/sam_region_overlay_renderer.py b/simple_syrup/runtime/sam_region_overlay_renderer.py new file mode 100644 index 0000000..c1bebd4 --- /dev/null +++ b/simple_syrup/runtime/sam_region_overlay_renderer.py @@ -0,0 +1,142 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Render deterministic SAM-style colored overlays from retained SEGS.""" + +from __future__ import annotations + +import torch +import torch.nn.functional as functional + +from ..domain.segs import NativeSegs, Segment, coerce_segs + +_REGION_COLORS: tuple[tuple[float, float, float], ...] = ( + (0.95, 0.26, 0.21), + (0.13, 0.59, 0.95), + (0.30, 0.69, 0.31), + (1.00, 0.76, 0.03), + (0.61, 0.15, 0.69), + (1.00, 0.34, 0.13), + (0.00, 0.74, 0.83), + (0.91, 0.12, 0.39), + (0.55, 0.76, 0.29), + (0.40, 0.23, 0.72), + (1.00, 0.60, 0.00), + (0.00, 0.59, 0.53), +) + + +class SAMRegionOverlayRenderer: + """Color retained SAM regions without changing their source SEGS geometry.""" + + def render(self, *, image: torch.Tensor, segs: NativeSegs) -> torch.Tensor: + """Return a source-sized IMAGE with translucent masks and stronger edges.""" + + _validate_image(image) + (segs_height, segs_width), segments = coerce_segs(segs) + image_height = int(image.shape[1]) + image_width = int(image.shape[2]) + if (segs_height, segs_width) != (image_height, image_width): + raise ValueError( + "SAM region overlay requires SEGS dimensions to match the image." + ) + if not segments: + return image.detach().clone() + + color_channels = min(3, int(image.shape[-1])) + device = image.device + color_sum = torch.zeros( + (image_height, image_width, color_channels), + device=device, + dtype=torch.float32, + ) + coverage = torch.zeros( + (image_height, image_width, 1), + device=device, + dtype=torch.float32, + ) + boundaries = torch.zeros_like(coverage) + boundary_thickness = max(1, round(max(image_height, image_width) / 1024)) + + for index, segment in enumerate(segments): + mask = _validated_local_mask(segment, device=device) + region = segment.crop_region + color = torch.tensor( + _REGION_COLORS[index % len(_REGION_COLORS)][:color_channels], + device=device, + dtype=torch.float32, + ) + mask_channels = mask.unsqueeze(-1) + region_slice = ( + slice(region.top, region.bottom), + slice(region.left, region.right), + ) + color_sum[region_slice] += mask_channels * color + coverage[region_slice] += mask_channels + boundaries[region_slice] = torch.maximum( + boundaries[region_slice], + _mask_boundary(mask, thickness=boundary_thickness).unsqueeze(-1), + ) + + covered = coverage > 0 + mean_color = color_sum / coverage.clamp_min(1.0) + alpha = torch.where( + boundaries > 0, + torch.full_like(coverage, 0.82), + torch.full_like(coverage, 0.46), + ) + alpha = torch.where(covered, alpha, torch.zeros_like(alpha)) + output = image.detach().clone() + source_color = output[0, :, :, :color_channels].float() + output[0, :, :, :color_channels] = ( + source_color * (1.0 - alpha) + mean_color * alpha + ).to(dtype=output.dtype) + return output.clamp(0.0, 1.0) + + +def _validate_image(image: torch.Tensor) -> None: + """Reject tensors that cannot represent one ComfyUI image.""" + + if image.ndim != 4 or int(image.shape[0]) != 1: + raise ValueError("SAM region overlay requires one BHWC image.") + if int(image.shape[-1]) < 1: + raise ValueError("SAM region overlay requires at least one image channel.") + + +def _validated_local_mask( + segment: Segment, + *, + device: torch.device, +) -> torch.Tensor: + """Return one binary crop-local mask after validating its SEG geometry.""" + + mask = segment.cropped_mask + if not isinstance(mask, torch.Tensor): + raise TypeError("SAM region overlay requires tensor SEG masks.") + working = mask.detach().to(device=device, dtype=torch.float32) + if working.ndim == 3 and int(working.shape[0]) == 1: + working = working.squeeze(0) + if working.ndim != 2: + raise ValueError("SAM region overlay requires HW or 1HW SEG masks.") + expected_shape = (segment.crop_region.height, segment.crop_region.width) + if tuple(working.shape) != expected_shape: + raise ValueError("SAM region overlay mask dimensions must match its SEG crop.") + return (working >= 0.5).float() + + +def _mask_boundary(mask: torch.Tensor, *, thickness: int) -> torch.Tensor: + """Return the interior boundary band of one binary mask.""" + + kernel_size = thickness * 2 + 1 + padded = functional.pad( + mask.unsqueeze(0).unsqueeze(0), + (thickness, thickness, thickness, thickness), + value=0.0, + ) + eroded = -functional.max_pool2d( + -padded, + kernel_size=kernel_size, + stride=1, + ) + return (mask - eroded.squeeze(0).squeeze(0)).clamp(0.0, 1.0) diff --git a/simple_syrup/services/segs_from_sam_output_service.py b/simple_syrup/services/segs_from_sam_output_service.py index 403f8d4..20aab62 100644 --- a/simple_syrup/services/segs_from_sam_output_service.py +++ b/simple_syrup/services/segs_from_sam_output_service.py @@ -21,6 +21,7 @@ from ..runtime.sam_automatic_segmenter import ( SAMAutomaticSegmenter, SAMModelAutomaticSegmenter, ) +from ..runtime.sam_region_overlay_renderer import SAMRegionOverlayRenderer from ..shared.logging import get_logger LOGGER = get_logger(__name__) @@ -34,6 +35,14 @@ class SAMAutoSegsSettings: minimum_region_area: int +@dataclass(frozen=True) +class SEGSFromSAMOutputResult: + """Return retained SEGS together with their source-image visualization.""" + + segs: NativeSegs + overlay: torch.Tensor + + @dataclass(frozen=True) class _GuideMaskCandidate: """Keep a retained SAM mask in compact guide-image coordinates.""" @@ -47,10 +56,15 @@ class _GuideMaskCandidate: class SEGSFromSAMOutputService: """Build reusable SEGS from a SAM model's unprompted masks.""" - def __init__(self, segmenter: SAMAutomaticSegmenter | None = None) -> None: + def __init__( + self, + segmenter: SAMAutomaticSegmenter | None = None, + overlay_renderer: SAMRegionOverlayRenderer | None = None, + ) -> None: """Create the service with an injectable automatic segmentation runtime.""" self._segmenter = segmenter or SAMModelAutomaticSegmenter() + self._overlay_renderer = overlay_renderer or SAMRegionOverlayRenderer() def build( self, @@ -60,8 +74,8 @@ class SEGSFromSAMOutputService: segmentation_resolution: int, minimum_region_area: int, phase_progress: PhaseProgressReporter | None = None, - ) -> NativeSegs: - """Return source-sized SEGS from unprompted SAM masks.""" + ) -> SEGSFromSAMOutputResult: + """Return source-sized SEGS and their SAM-style colored overlay.""" operation_started_at = perf_counter() reporter = phase_progress or NullPhaseProgressReporter() @@ -98,6 +112,10 @@ class SEGSFromSAMOutputService: ) for index, candidate in enumerate(candidates, start=1) ) + segs: NativeSegs = (image_height, image_width), segments + segs_built_at = perf_counter() + reporter.advance("rendering_overlay") + overlay = self._overlay_renderer.render(image=source_image, segs=segs) LOGGER.info( "Built SEGS from SAM output", extra={ @@ -113,12 +131,16 @@ class SEGSFromSAMOutputService: 2, ), "segs_construction_ms": round( - (perf_counter() - masks_generated_at) * 1000.0, + (segs_built_at - masks_generated_at) * 1000.0, + 2, + ), + "overlay_rendering_ms": round( + (perf_counter() - segs_built_at) * 1000.0, 2, ), }, ) - return (image_height, image_width), segments + return SEGSFromSAMOutputResult(segs=segs, overlay=overlay) def _validate_settings( diff --git a/tests/test_sam_region_overlay_renderer.py b/tests/test_sam_region_overlay_renderer.py new file mode 100644 index 0000000..53762df --- /dev/null +++ b/tests/test_sam_region_overlay_renderer.py @@ -0,0 +1,63 @@ +# 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 deterministic SAM-style region overlay rendering.""" + +from __future__ import annotations + +import torch + +from simple_syrup.domain.segs import BoundingBox, CropRegion, Segment +from simple_syrup.runtime.sam_region_overlay_renderer import SAMRegionOverlayRenderer + + +def test_renderer_colors_retained_regions_and_preserves_uncovered_pixels() -> None: + """The overlay colors each SEG while leaving the source visible elsewhere.""" + + image = torch.full((1, 8, 10, 3), 0.25) + first = _segment(CropRegion(1, 1, 6, 6), torch.ones((5, 5)), "first") + second_mask = torch.zeros((5, 5)) + second_mask[1:4, 1:4] = 1.0 + second = _segment(CropRegion(4, 2, 9, 7), second_mask, "second") + + overlay = SAMRegionOverlayRenderer().render( + image=image, + segs=((8, 10), (first, second)), + ) + + assert overlay.shape == image.shape + assert overlay.dtype == image.dtype + assert torch.equal(overlay[:, 0, 0], image[:, 0, 0]) + assert not torch.equal(overlay[:, 2, 2], image[:, 2, 2]) + assert not torch.equal(overlay[:, 4, 5], overlay[:, 2, 2]) + assert torch.equal(image, torch.full_like(image, 0.25)) + + +def test_renderer_is_deterministic_and_empty_segs_return_source_copy() -> None: + """Stable colors support comparisons and empty detections remain readable.""" + + image = torch.rand((1, 6, 6, 3), generator=torch.Generator().manual_seed(4)) + segment = _segment(CropRegion(1, 1, 5, 5), torch.ones((4, 4)), "subject") + renderer = SAMRegionOverlayRenderer() + + first = renderer.render(image=image, segs=((6, 6), (segment,))) + second = renderer.render(image=image, segs=((6, 6), (segment,))) + empty = renderer.render(image=image, segs=((6, 6), ())) + + assert torch.equal(first, second) + assert torch.equal(empty, image) + assert empty.data_ptr() != image.data_ptr() + + +def _segment(region: CropRegion, mask: torch.Tensor, label: str) -> Segment: + """Create one source-aligned segment for renderer tests.""" + + return Segment( + cropped_image=None, + cropped_mask=mask, + confidence=1.0, + crop_region=region, + bbox=BoundingBox(*region), + label=label, + ) diff --git a/tests/test_segs_from_sam_output_node.py b/tests/test_segs_from_sam_output_node.py index b8129c0..8a5066b 100644 --- a/tests/test_segs_from_sam_output_node.py +++ b/tests/test_segs_from_sam_output_node.py @@ -9,7 +9,11 @@ from __future__ import annotations import pytest import torch +from simple_syrup.domain.segs import NativeSegs from simple_syrup.nodes.segs_from_sam_output import SEGSFromSAMOutput +from simple_syrup.services.segs_from_sam_output_service import ( + SEGSFromSAMOutputResult, +) def test_node_declares_automatic_sam_to_segs_contract() -> None: @@ -25,8 +29,9 @@ def test_node_declares_automatic_sam_to_segs_contract() -> None: ) 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,) + assert SEGSFromSAMOutput.RETURN_TYPES == ("SEGS", "IMAGE") + assert SEGSFromSAMOutput.RETURN_NAMES == ("segs", "overlay") + assert SEGSFromSAMOutput.OUTPUT_IS_LIST == (True, False) def test_node_builds_one_segs_output_per_image_batch_item( @@ -52,9 +57,17 @@ def test_node_builds_one_segs_output_per_image_batch_item( "preparing_segmentation_image", "generating_automatic_masks", "building_segs", + "rendering_overlay", ): phase_progress.advance(phase) - return (image.shape[1:3], ()) + empty_segs: NativeSegs = ( + (int(image.shape[1]), int(image.shape[2])), + (), + ) + return SEGSFromSAMOutputResult( + segs=empty_segs, + overlay=image + len(calls) / 10.0, + ) monkeypatch.setattr(SEGSFromSAMOutput, "service_class", _Service) monkeypatch.setattr( @@ -63,7 +76,7 @@ def test_node_builds_one_segs_output_per_image_batch_item( lambda **_kwargs: phase_progress, ) - (segs,) = SEGSFromSAMOutput().generate( + segs, overlay = SEGSFromSAMOutput().generate( image=torch.zeros((2, 16, 16, 3)), sam_model=object(), segmentation_resolution=640, @@ -72,13 +85,18 @@ def test_node_builds_one_segs_output_per_image_batch_item( assert len(calls) == 2 assert len(segs) == 2 + assert overlay.shape == (2, 16, 16, 3) + assert torch.allclose(overlay[0], torch.full((16, 16, 3), 0.1)) + assert torch.allclose(overlay[1], torch.full((16, 16, 3), 0.2)) assert phase_progress.phases == [ "preparing_segmentation_image", "generating_automatic_masks", "building_segs", + "rendering_overlay", "preparing_segmentation_image", "generating_automatic_masks", "building_segs", + "rendering_overlay", "completed", ] diff --git a/tests/test_segs_from_sam_output_service.py b/tests/test_segs_from_sam_output_service.py index 5be626f..c88837f 100644 --- a/tests/test_segs_from_sam_output_service.py +++ b/tests/test_segs_from_sam_output_service.py @@ -59,13 +59,14 @@ def test_service_downscales_the_segmentation_guide_without_upscaling_source() -> runtime = _RecordingSegmenter((AutomaticSAMMask(guide_mask, 0.8),)) service = SEGSFromSAMOutputService(runtime) - segs = service.build( + result = service.build( image=torch.zeros((1, 128, 256, 3)), sam_model=object(), segmentation_resolution=64, minimum_region_area=0, ) + segs = result.segs assert runtime.image_shapes == [(32, 64)] assert segs[0] == (128, 256) segment = segs[1][0] @@ -94,6 +95,7 @@ def test_service_reports_meaningful_automatic_segmentation_phases() -> None: "preparing_segmentation_image", "generating_automatic_masks", "building_segs", + "rendering_overlay", ] @@ -105,21 +107,21 @@ def test_service_filters_region_area_after_restoring_source_dimensions() -> None runtime = _RecordingSegmenter((AutomaticSAMMask(guide_mask, 0.5, "thing"),)) service = SEGSFromSAMOutputService(runtime) - retained = service.build( + retained_result = service.build( image=torch.zeros((1, 128, 256, 3)), sam_model=object(), segmentation_resolution=64, minimum_region_area=63, ) - filtered = service.build( + filtered_result = 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] == () + assert retained_result.segs[1][0].label == "thing" + assert filtered_result.segs[1] == () def test_service_suppresses_duplicate_masks_and_keeps_highest_confidence() -> None: @@ -133,16 +135,16 @@ def test_service_suppresses_duplicate_masks_and_keeps_highest_confidence() -> No ) ) - segs = SEGSFromSAMOutputService(runtime).build( + result = 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" + assert len(result.segs[1]) == 1 + assert result.segs[1][0].confidence == 0.9 + assert result.segs[1][0].label == "second" def test_service_expands_retained_mask_crops_to_source_resolution() -> None: @@ -151,7 +153,7 @@ def test_service_expands_retained_mask_crops_to_source_resolution() -> None: guide_mask = torch.zeros((64, 128), dtype=torch.float32) guide_mask[16:32, 32:64] = 1.0 - segs = SEGSFromSAMOutputService( + result = SEGSFromSAMOutputService( _RecordingSegmenter((AutomaticSAMMask(guide_mask, 1.0),)) ).build( image=torch.zeros((1, 1024, 2048, 3)), @@ -160,7 +162,7 @@ def test_service_expands_retained_mask_crops_to_source_resolution() -> None: minimum_region_area=0, ) - segment = segs[1][0] + segment = result.segs[1][0] assert segment.crop_region == (512, 256, 1024, 512) assert cast(torch.Tensor, segment.cropped_mask).shape == (256, 512)