feat(segmentation): add SAM-guided tiled diffusion

This commit is contained in:
Artificial Sweetener
2026-07-31 22:08:15 -04:00
parent 546ee5db01
commit 36081d7771
24 changed files with 2331 additions and 18 deletions
+505
View File
@@ -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,
)
+4 -3
View File
@@ -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]:
+14 -1
View File
@@ -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,)
+120
View File
@@ -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__
+2
View File
@@ -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,
@@ -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",
@@ -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."""
+19
View File
@@ -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, ...] = (
@@ -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,
@@ -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))
+22 -4
View File
@@ -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}."
+67
View File
@@ -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,
@@ -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()
)
@@ -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,
*,
@@ -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
+1
View File
@@ -44,6 +44,7 @@ BASE_NODE_IDS = [
"SimpleSyrup.PromptSEGSWithSAM",
"SimpleSyrup.ResizeImageToTarget",
"SimpleSyrup.SAMModelLoader",
"SimpleSyrup.SEGSFromSAMOutput",
"SimpleSyrup.ScaleFactor",
"SimpleSyrup.Seed",
"SimpleSyrup.SimpleLoadAnima",
+131
View File
@@ -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]]),
)
+39
View File
@@ -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."""
+1
View File
@@ -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:
+97
View File
@@ -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)
+180
View File
@@ -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,
)
+189
View File
@@ -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,
)
+48 -1
View File
@@ -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",
)
+24
View File
@@ -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."""