feat(segmentation): add SAM-guided tiled diffusion
This commit is contained in:
@@ -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,
|
||||
)
|
||||
@@ -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]:
|
||||
|
||||
@@ -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,)
|
||||
|
||||
@@ -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__
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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))
|
||||
@@ -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}."
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -44,6 +44,7 @@ BASE_NODE_IDS = [
|
||||
"SimpleSyrup.PromptSEGSWithSAM",
|
||||
"SimpleSyrup.ResizeImageToTarget",
|
||||
"SimpleSyrup.SAMModelLoader",
|
||||
"SimpleSyrup.SEGSFromSAMOutput",
|
||||
"SimpleSyrup.ScaleFactor",
|
||||
"SimpleSyrup.Seed",
|
||||
"SimpleSyrup.SimpleLoadAnima",
|
||||
|
||||
@@ -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]]),
|
||||
)
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user