feat(segmentation): add SAM region overlay

This commit is contained in:
Artificial Sweetener
2026-08-02 02:43:27 -04:00
parent 29ee772b1d
commit 1972e452dc
6 changed files with 286 additions and 35 deletions
+19 -15
View File
@@ -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:
@@ -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)
@@ -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(
+63
View File
@@ -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,
)
+22 -4
View File
@@ -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",
]
+13 -11
View File
@@ -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)