From ace5fafda90cdea7c0e42cfc88b3ceb1a105ca14 Mon Sep 17 00:00:00 2001 From: Artificial Sweetener Date: Mon, 10 Aug 2026 00:55:03 -0400 Subject: [PATCH] feat(sampling): add regional diffusion sampling --- simple_syrup/domain/contextual_diffusion.py | 24 +- .../domain/regional_tiled_diffusion.py | 106 +++++ simple_syrup/domain/segs_tiled_diffusion.py | 346 +--------------- .../domain/semantic_tiled_diffusion.py | 384 ++++++++++++++++++ simple_syrup/masking/regional_prompt_masks.py | 47 ++- .../nodes/ksampler_contextual_diffusion.py | 34 +- .../nodes/ksampler_tiled_diffusion.py | 32 ++ simple_syrup/nodes/tooltips.py | 12 + .../runtime/contextual_diffusion_sampling.py | 13 +- .../contextual_diffusion_sampling_service.py | 49 ++- .../services/regional_conditioning_service.py | 27 +- .../regional_sampling_preparation_service.py | 112 +++++ ...gional_tiled_diffusion_sampling_service.py | 143 +++++++ .../tiled_diffusion_sampling_service.py | 51 ++- tests/test_contextual_diffusion_sampling.py | 20 + ...t_contextual_diffusion_sampling_service.py | 110 ++++- ...test_ksampler_contextual_diffusion_node.py | 6 + tests/test_ksampler_tiled_diffusion_node.py | 12 + tests/test_patcher_lifecycle_policy.py | 2 +- ...t_regional_sampling_preparation_service.py | 101 +++++ tests/test_regional_tiled_diffusion.py | 103 +++++ .../test_tiled_diffusion_sampling_service.py | 121 ++++++ 22 files changed, 1502 insertions(+), 353 deletions(-) create mode 100644 simple_syrup/domain/regional_tiled_diffusion.py create mode 100644 simple_syrup/domain/semantic_tiled_diffusion.py create mode 100644 simple_syrup/services/regional_sampling_preparation_service.py create mode 100644 simple_syrup/services/regional_tiled_diffusion_sampling_service.py create mode 100644 tests/test_regional_sampling_preparation_service.py create mode 100644 tests/test_regional_tiled_diffusion.py diff --git a/simple_syrup/domain/contextual_diffusion.py b/simple_syrup/domain/contextual_diffusion.py index b92aa3f..a6ba8f9 100644 --- a/simple_syrup/domain/contextual_diffusion.py +++ b/simple_syrup/domain/contextual_diffusion.py @@ -8,6 +8,9 @@ from __future__ import annotations from dataclasses import dataclass +import torch + +from .regional_tiled_diffusion import build_region_constrained_tiled_diffusion_plan from .segs import NativeSegs from .segs_tiled_diffusion import build_segs_guided_tiled_diffusion_plan from .tiled_diffusion import TiledDiffusionPlan, build_tiled_diffusion_plan @@ -72,6 +75,7 @@ def build_contextual_diffusion_plan( latent_height: int, controls: ContextualDiffusionControls, segs: NativeSegs | None, + region_masks: torch.Tensor | None = None, ) -> ContextualDiffusionPlan: """Return a global context plus the regular or SEGS-guided context plan.""" @@ -89,8 +93,9 @@ def build_contextual_diffusion_plan( context_width=global_width, context_height=global_height, ) - tile_plan = ( - build_segs_guided_tiled_diffusion_plan( + if region_masks is not None: + tile_plan = build_region_constrained_tiled_diffusion_plan( + region_masks=region_masks, segs=segs, latent_width=latent_width, latent_height=latent_height, @@ -99,8 +104,18 @@ def build_contextual_diffusion_plan( overlap=controls.latent_context_overlap, tile_batch_size=controls.latent_context_batch_size, ) - if segs is not None - else build_tiled_diffusion_plan( + elif segs is not None: + tile_plan = build_segs_guided_tiled_diffusion_plan( + segs=segs, + latent_width=latent_width, + latent_height=latent_height, + tile_width=controls.latent_context_size, + tile_height=controls.latent_context_size, + overlap=controls.latent_context_overlap, + tile_batch_size=controls.latent_context_batch_size, + ) + else: + tile_plan = build_tiled_diffusion_plan( latent_width=latent_width, latent_height=latent_height, tile_width=controls.latent_context_size, @@ -108,7 +123,6 @@ def build_contextual_diffusion_plan( overlap=controls.latent_context_overlap, tile_batch_size=controls.latent_context_batch_size, ) - ) return ContextualDiffusionPlan( latent_width=latent_width, latent_height=latent_height, diff --git a/simple_syrup/domain/regional_tiled_diffusion.py b/simple_syrup/domain/regional_tiled_diffusion.py new file mode 100644 index 0000000..6fd2457 --- /dev/null +++ b/simple_syrup/domain/regional_tiled_diffusion.py @@ -0,0 +1,106 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Constrain semantic tiled diffusion by authored regional composition masks.""" + +from __future__ import annotations + +import torch + +from .segs import coerce_segs +from .segs_tiled_diffusion import segs_ownership_masks, validate_segs_aspect_ratio +from .semantic_tiled_diffusion import build_semantic_tiled_diffusion_plan +from .tiled_diffusion import TiledDiffusionPlan + +REGIONAL_PLANNING_THRESHOLD = 0.5 + + +def build_region_constrained_tiled_diffusion_plan( + *, + region_masks: torch.Tensor, + segs: object | None, + latent_width: int, + latent_height: int, + tile_width: int, + tile_height: int, + overlap: int, + tile_batch_size: int, +) -> TiledDiffusionPlan: + """Build tiles split wherever regional composition or optional SEGS change.""" + + region_ownership = regional_composition_ownership_masks( + region_masks, + latent_height=latent_height, + latent_width=latent_width, + ) + ownership_masks = region_ownership + if segs is not None: + native_segs = coerce_segs(segs) + validate_segs_aspect_ratio(native_segs, latent_height, latent_width) + semantic_ownership = segs_ownership_masks( + native_segs, + latent_height=latent_height, + latent_width=latent_width, + ) + ownership_masks = _intersect_partitions( + region_ownership, + semantic_ownership, + ) + return build_semantic_tiled_diffusion_plan( + ownership_masks=ownership_masks, + latent_width=latent_width, + latent_height=latent_height, + tile_width=tile_width, + tile_height=tile_height, + overlap=overlap, + tile_batch_size=tile_batch_size, + merge_across_masks=False, + ) + + +def regional_composition_ownership_masks( + region_masks: torch.Tensor, + *, + latent_height: int, + latent_width: int, +) -> tuple[torch.Tensor, ...]: + """Partition the canvas by every distinct active regional-mask combination.""" + + if region_masks.ndim != 3: + raise ValueError("Regional tile planning requires a BHW mask batch.") + if tuple(region_masks.shape[1:]) != (latent_height, latent_width): + raise ValueError( + "Regional tile planning masks must match latent shape " + f"{latent_height}x{latent_width}." + ) + membership = (region_masks.detach().cpu() >= REGIONAL_PLANNING_THRESHOLD).permute( + 1, 2, 0 + ) + flattened = membership.reshape(latent_height * latent_width, -1) + signatures, inverse = torch.unique( + flattened, + dim=0, + sorted=True, + return_inverse=True, + ) + del signatures + labels = inverse.reshape(latent_height, latent_width) + return tuple(labels == index for index in range(int(labels.max().item()) + 1)) + + +def _intersect_partitions( + first: tuple[torch.Tensor, ...], + second: tuple[torch.Tensor, ...], +) -> tuple[torch.Tensor, ...]: + """Return non-empty intersections of two complete ownership partitions.""" + + intersections = tuple( + intersection + for first_mask in first + for second_mask in second + if bool((intersection := torch.logical_and(first_mask, second_mask)).any()) + ) + if not intersections: + raise ValueError("Regional and SEGS ownership produced no tile coverage.") + return intersections diff --git a/simple_syrup/domain/segs_tiled_diffusion.py b/simple_syrup/domain/segs_tiled_diffusion.py index 3a3d694..601bce4 100644 --- a/simple_syrup/domain/segs_tiled_diffusion.py +++ b/simple_syrup/domain/segs_tiled_diffusion.py @@ -6,27 +6,11 @@ from __future__ import annotations -from dataclasses import dataclass - import torch -import torch.nn.functional as functional from .segs import NativeSegs, Segment, coerce_segment_mask, 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 +from .semantic_tiled_diffusion import build_semantic_tiled_diffusion_plan +from .tiled_diffusion import TiledDiffusionPlan def build_segs_guided_tiled_diffusion_plan( @@ -46,64 +30,22 @@ def build_segs_guided_tiled_diffusion_plan( boundary and shares a feathered overlap with neighboring cores. """ - base_plan = build_tiled_diffusion_plan( + native_segs = coerce_segs(segs) + validate_segs_aspect_ratio(native_segs, latent_height, latent_width) + ownership_masks = segs_ownership_masks( + native_segs, + latent_height=latent_height, + latent_width=latent_width, + ) + return build_semantic_tiled_diffusion_plan( + ownership_masks=ownership_masks, 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_segs_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, + merge_across_masks=True, ) @@ -126,13 +68,13 @@ def validate_segs_aspect_ratio( ) -def _build_ownership_cores( +def segs_ownership_masks( segs: NativeSegs, *, latent_height: int, latent_width: int, -) -> tuple[_OwnershipCore, ...]: - """Resolve overlapping SEGS into one deterministic latent ownership partition.""" +) -> tuple[torch.Tensor, ...]: + """Resolve overlapping SEGS into a deterministic latent ownership partition.""" source_height, source_width = segs[0] segment_masks = tuple( @@ -154,20 +96,18 @@ def _build_ownership_cores( ), ) occupied = torch.zeros((latent_height, latent_width), dtype=torch.bool) - cores: list[_OwnershipCore] = [] + ownership_masks: list[torch.Tensor] = [] 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)) + ownership_masks.append(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)), - ) + ownership_masks.append(background) + if ownership_masks: + return tuple(ownership_masks) + return (torch.ones((latent_height, latent_width), dtype=torch.bool),) def segment_mask_to_latent( @@ -259,223 +199,6 @@ def segment_weight_to_latent( 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, @@ -487,30 +210,3 @@ def _latent_sample_range( start = (source_start * latent_limit + source_limit - 1) // source_limit end = (source_end * latent_limit + source_limit - 1) // source_limit return max(0, min(latent_limit, start)), max(0, min(latent_limit, end)) - - -def _core_from_mask(mask: torch.Tensor) -> _OwnershipCore: - """Build one core with bounds and area computed exactly once.""" - - bounds = _mask_bounds(mask) - if bounds is None: - raise ValueError("SEGS-guided tiled diffusion cannot use an empty core.") - return _OwnershipCore( - mask=mask, - bounds=bounds, - area=int(mask.sum().item()), - ) - - -def _mask_bounds(mask: torch.Tensor) -> tuple[int, int, int, int] | None: - """Return left, top, right, bottom bounds for one non-empty boolean mask.""" - - y_coords, x_coords = torch.where(mask) - if y_coords.numel() == 0: - return None - return ( - int(x_coords.min().item()), - int(y_coords.min().item()), - int(x_coords.max().item()) + 1, - int(y_coords.max().item()) + 1, - ) diff --git a/simple_syrup/domain/semantic_tiled_diffusion.py b/simple_syrup/domain/semantic_tiled_diffusion.py new file mode 100644 index 0000000..e25a7b5 --- /dev/null +++ b/simple_syrup/domain/semantic_tiled_diffusion.py @@ -0,0 +1,384 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Build tiled diffusion plans from non-overlapping latent ownership masks.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass + +import torch +import torch.nn.functional as functional + +from .tiled_diffusion import ( + LatentTile, + TiledDiffusionPlan, + batch_latent_tiles, + build_tiled_diffusion_plan, +) + + +@dataclass(frozen=True) +class _OwnershipCore: + """Represent one latent ownership region before sampling-window placement.""" + + mask: torch.Tensor + bounds: tuple[int, int, int, int] + area: int + + +def build_semantic_tiled_diffusion_plan( + *, + ownership_masks: Sequence[torch.Tensor], + latent_width: int, + latent_height: int, + tile_width: int, + tile_height: int, + overlap: int, + tile_batch_size: int, + merge_across_masks: bool, +) -> TiledDiffusionPlan: + """Build bounded windows whose write weights follow ownership masks.""" + + 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, + ) + normalized_masks = _validate_ownership_masks( + ownership_masks, + 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_groups = tuple( + tuple( + split_core + for split_core in _split_core( + _core_from_mask(mask), + max_width=max_core_width, + max_height=max_core_height, + ) + ) + for mask in normalized_masks + ) + if merge_across_masks: + cores = _merge_small_cores( + tuple(core for group in split_groups for core in group), + max_width=max_core_width, + max_height=max_core_height, + ) + else: + cores = tuple( + core + for group in split_groups + for core in _merge_small_cores( + group, + 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 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_ownership_masks( + masks: Sequence[torch.Tensor], + *, + latent_height: int, + latent_width: int, +) -> tuple[torch.Tensor, ...]: + """Return non-empty boolean masks that cover the complete latent canvas.""" + + normalized: list[torch.Tensor] = [] + coverage = torch.zeros((latent_height, latent_width), dtype=torch.bool) + for index, mask in enumerate(masks): + if mask.ndim != 2 or tuple(mask.shape) != (latent_height, latent_width): + raise ValueError( + "Semantic ownership mask " + f"{index} must match latent shape {latent_height}x{latent_width}." + ) + boolean_mask = mask.detach().cpu().bool() + if not bool(boolean_mask.any()): + continue + if bool(torch.logical_and(coverage, boolean_mask).any()): + raise ValueError("Semantic ownership masks must not overlap.") + normalized.append(boolean_mask) + coverage = torch.logical_or(coverage, boolean_mask) + if not normalized: + raise ValueError("Semantic tiled diffusion requires non-empty ownership.") + if not bool(coverage.all()): + raise ValueError("Semantic ownership masks must cover the latent canvas.") + return tuple(normalized) + + +def _split_core( + core: _OwnershipCore, + *, + max_width: int, + max_height: int, +) -> tuple[_OwnershipCore, ...]: + """Recursively divide a core into balanced pieces within 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 semantic tile core.") + return first, second + + +def _merge_small_cores( + cores: tuple[_OwnershipCore, ...], + *, + max_width: int, + max_height: int, +) -> tuple[_OwnershipCore, ...]: + """Greedily combine 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 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("Semantic 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 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 _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("Semantic 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 a non-empty boolean mask.""" + + y_coords, x_coords = torch.where(mask) + if y_coords.numel() == 0: + return None + return ( + int(x_coords.min().item()), + int(y_coords.min().item()), + int(x_coords.max().item()) + 1, + int(y_coords.max().item()) + 1, + ) diff --git a/simple_syrup/masking/regional_prompt_masks.py b/simple_syrup/masking/regional_prompt_masks.py index 5215d70..35fa810 100644 --- a/simple_syrup/masking/regional_prompt_masks.py +++ b/simple_syrup/masking/regional_prompt_masks.py @@ -6,14 +6,31 @@ from __future__ import annotations +from dataclasses import dataclass + import torch +import torch.nn.functional as functional from .detailer_masks import gaussian_feather_mask +@dataclass(frozen=True) +class PreparedRegionalMasks: + """Keep authored geometry separate from feathered conditioning influence.""" + + authored: torch.Tensor + conditioning: torch.Tensor + + def prepare_regional_mask_batch(mask: object, feather: int) -> torch.Tensor: """Return a validated and optionally feathered BHW mask batch.""" + return prepare_regional_masks(mask, feather).conditioning + + +def prepare_regional_masks(mask: object, feather: int) -> PreparedRegionalMasks: + """Return normalized authored masks and their feathered conditioning form.""" + if not isinstance(mask, torch.Tensor): raise TypeError("regional prompting requires a torch MASK tensor.") if feather < 0: @@ -30,9 +47,33 @@ def prepare_regional_mask_batch(mask: object, feather: int) -> torch.Tensor: raise ValueError("regional masks must have non-empty height and width.") normalized = working.clamp(0.0, 1.0) - if feather == 0: - return normalized - return gaussian_feather_mask(normalized, feather) + conditioning = ( + normalized if feather == 0 else gaussian_feather_mask(normalized, feather) + ) + return PreparedRegionalMasks( + authored=normalized, + conditioning=conditioning, + ) + + +def resize_regional_mask_batch( + mask_batch: torch.Tensor, + *, + height: int, + width: int, +) -> torch.Tensor: + """Resize BHW masks to one latent canvas using Comfy-compatible scaling.""" + + if height < 1 or width < 1: + raise ValueError("regional mask target height and width must be positive.") + if tuple(mask_batch.shape[1:]) == (height, width): + return mask_batch + return functional.interpolate( + mask_batch.unsqueeze(1), + size=(height, width), + mode="bilinear", + align_corners=False, + ).squeeze(1) def regional_mask(mask_batch: torch.Tensor, index: int) -> torch.Tensor: diff --git a/simple_syrup/nodes/ksampler_contextual_diffusion.py b/simple_syrup/nodes/ksampler_contextual_diffusion.py index e624354..e68d445 100644 --- a/simple_syrup/nodes/ksampler_contextual_diffusion.py +++ b/simple_syrup/nodes/ksampler_contextual_diffusion.py @@ -8,6 +8,7 @@ from __future__ import annotations from typing import Any, ClassVar, TypeAlias +from ..domain.regional_prompting import MAX_REGIONAL_PROMPT_WEIGHT from ..domain.tiled_diffusion import TILED_DIFFUSION_MODES from ..runtime import sampling_samplers, sampling_schedulers from ..services.contextual_diffusion_sampling_service import ( @@ -182,7 +183,32 @@ class KSamplerContextualDiffusion: "segs": ( "SEGS", {"tooltip": tooltips.CONTEXTUAL_DIFFUSION_SEGS}, - ) + ), + "region_masks": ( + "MASK", + {"tooltip": tooltips.OPTIONAL_REGIONAL_MASKS}, + ), + "regional_prompt_weight": ( + "FLOAT", + { + "default": 0.5, + "min": 0.0, + "max": MAX_REGIONAL_PROMPT_WEIGHT, + "step": 0.01, + "round": 0.01, + "tooltip": tooltips.OPTIONAL_REGIONAL_PROMPT_WEIGHT, + }, + ), + "region_mask_feather": ( + "INT", + { + "default": 0, + "min": 0, + "max": 512, + "step": 1, + "tooltip": tooltips.OPTIONAL_REGION_MASK_FEATHER, + }, + ), }, } @@ -206,6 +232,9 @@ class KSamplerContextualDiffusion: global_steps: int = 1, global_decay: float = 0.5, segs: object | None = None, + region_masks: object | None = None, + regional_prompt_weight: float = 0.5, + region_mask_feather: int = 0, ) -> tuple[Latent, object]: """Delegate contextual diffusion sampling to its application service.""" @@ -228,5 +257,8 @@ class KSamplerContextualDiffusion: global_steps=global_steps, global_decay=global_decay, segs=segs, + region_masks=region_masks, + regional_prompt_weight=regional_prompt_weight, + region_mask_feather=region_mask_feather, ) return result.latent, result.contexts diff --git a/simple_syrup/nodes/ksampler_tiled_diffusion.py b/simple_syrup/nodes/ksampler_tiled_diffusion.py index b8ecab7..e42ab5f 100644 --- a/simple_syrup/nodes/ksampler_tiled_diffusion.py +++ b/simple_syrup/nodes/ksampler_tiled_diffusion.py @@ -8,6 +8,7 @@ from __future__ import annotations from typing import Any, ClassVar +from ..domain.regional_prompting import MAX_REGIONAL_PROMPT_WEIGHT from ..domain.tiled_diffusion import TILED_DIFFUSION_MODES from ..runtime import sampling_samplers, sampling_schedulers from ..services.tiled_diffusion_sampling_service import TiledDiffusionSamplingService @@ -160,6 +161,31 @@ class KSamplerTiledDiffusion: ), }, ), + "region_masks": ( + "MASK", + {"tooltip": tooltips.OPTIONAL_REGIONAL_MASKS}, + ), + "regional_prompt_weight": ( + "FLOAT", + { + "default": 0.5, + "min": 0.0, + "max": MAX_REGIONAL_PROMPT_WEIGHT, + "step": 0.01, + "round": 0.01, + "tooltip": tooltips.OPTIONAL_REGIONAL_PROMPT_WEIGHT, + }, + ), + "region_mask_feather": ( + "INT", + { + "default": 0, + "min": 0, + "max": 512, + "step": 1, + "tooltip": tooltips.OPTIONAL_REGION_MASK_FEATHER, + }, + ), }, } @@ -181,6 +207,9 @@ class KSamplerTiledDiffusion: latent_tile_overlap: int = 16, latent_tile_batch_size: int = 4, segs: object | None = None, + region_masks: object | None = None, + regional_prompt_weight: float = 0.5, + region_mask_feather: int = 0, ) -> tuple[Latent]: """Sample a latent with the selected tiled diffusion method.""" @@ -202,5 +231,8 @@ class KSamplerTiledDiffusion: latent_tile_batch_size=latent_tile_batch_size, preview_context=None, segs=segs, + region_masks=region_masks, + regional_prompt_weight=regional_prompt_weight, + region_mask_feather=region_mask_feather, ) return (output,) diff --git a/simple_syrup/nodes/tooltips.py b/simple_syrup/nodes/tooltips.py index 3fd85b1..f16619a 100644 --- a/simple_syrup/nodes/tooltips.py +++ b/simple_syrup/nodes/tooltips.py @@ -187,6 +187,18 @@ GLOBAL_CONTEXT_DECAY = ( CONTEXTUAL_DIFFUSION_SEGS = ( "Optional regions that replace the regular grid with SEGS-guided contexts." ) +OPTIONAL_REGIONAL_MASKS = ( + "Optional authored masks that activate regional prompting when positive or " + "negative conditioning is a batch; SEGS may further subdivide those regions." +) +OPTIONAL_REGIONAL_PROMPT_WEIGHT = ( + "Balances regional prompts against the global prompt when regional masks and " + "a conditioning batch are connected." +) +OPTIONAL_REGION_MASK_FEATHER = ( + "Softens regional conditioning edges by this many source-mask pixels; tile " + "boundaries continue to follow the unfeathered authored regions." +) DETAIL_IMAGE = ( "Source image containing the regions to improve. Detailed crops are blended " diff --git a/simple_syrup/runtime/contextual_diffusion_sampling.py b/simple_syrup/runtime/contextual_diffusion_sampling.py index 82017c6..3ea781b 100644 --- a/simple_syrup/runtime/contextual_diffusion_sampling.py +++ b/simple_syrup/runtime/contextual_diffusion_sampling.py @@ -52,6 +52,7 @@ def sample_contextual_diffusion( diffusion_mode: str, controls: ContextualDiffusionControls, plan: ContextualDiffusionPlan, + allow_full_context_masks: bool = False, ) -> Latent: """Sample one latent through global context and one tiled prediction plan.""" @@ -65,8 +66,16 @@ def sample_contextual_diffusion( controls.validate() if sampler_name in UNIPC_SAMPLERS: raise ValueError("Contextual Diffusion is not compatible with UniPC samplers.") - reject_unsupported_conditioning(positive, sampler_label=SAMPLER_LABEL) - reject_unsupported_conditioning(negative, sampler_label=SAMPLER_LABEL) + reject_unsupported_conditioning( + positive, + sampler_label=SAMPLER_LABEL, + allow_full_context_masks=allow_full_context_masks, + ) + reject_unsupported_conditioning( + negative, + sampler_label=SAMPLER_LABEL, + allow_full_context_masks=allow_full_context_masks, + ) sampler = sampling_samplers.resolve_sampler(sampler_name) sigmas = sampling_schedulers.calculate_sigmas( diff --git a/simple_syrup/services/contextual_diffusion_sampling_service.py b/simple_syrup/services/contextual_diffusion_sampling_service.py index 622f643..1a45a25 100644 --- a/simple_syrup/services/contextual_diffusion_sampling_service.py +++ b/simple_syrup/services/contextual_diffusion_sampling_service.py @@ -7,7 +7,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Any, TypeAlias +from typing import Any, ClassVar, TypeAlias import torch @@ -25,6 +25,9 @@ from ..domain.segs import NativeSegs, coerce_segs_group from ..domain.tiled_diffusion import validate_tiled_diffusion_mode from ..runtime.contextual_diffusion_sampling import sample_contextual_diffusion from ..runtime.latent_geometry import decoded_image_dimensions +from .regional_sampling_preparation_service import ( + RegionalSamplingPreparationService, +) from .sampling_batch import ( combine_latent_outputs, latent_batch_size, @@ -45,6 +48,10 @@ class ContextualDiffusionSamplingResult: class ContextualDiffusionSamplingService: """Plan and execute composition-preserving contextual diffusion.""" + regional_preparation_service_class: ClassVar[ + type[RegionalSamplingPreparationService] + ] = RegionalSamplingPreparationService + def sample( self, *, @@ -66,6 +73,9 @@ class ContextualDiffusionSamplingService: global_steps: int, global_decay: float, segs: object | None = None, + region_masks: object | None = None, + regional_prompt_weight: float = 0.5, + region_mask_feather: int = 0, ) -> ContextualDiffusionSamplingResult: """Sample a latent with global and bounded detail contexts.""" @@ -79,6 +89,14 @@ class ContextualDiffusionSamplingService: global_decay=global_decay, ) controls.validate() + regional = self.regional_preparation_service_class().prepare( + positive=positive, + negative=negative, + latent_image=latent_image, + region_masks=region_masks, + regional_prompt_weight=regional_prompt_weight, + region_mask_feather=region_mask_feather, + ) batch_size = latent_batch_size(latent_image) segs_group = coerce_segs_group(segs) if segs is not None else () if segs_group and len(segs_group) not in (1, batch_size): @@ -98,8 +116,9 @@ class ContextualDiffusionSamplingService: ) split_batch = ( bool(segs_group) - or isinstance(positive, ConditioningBatch) - or isinstance(negative, ConditioningBatch) + or regional.active + or isinstance(regional.positive, ConditioningBatch) + or isinstance(regional.negative, ConditioningBatch) ) if not split_batch: return self._sample_item( @@ -109,8 +128,8 @@ class ContextualDiffusionSamplingService: cfg=cfg, sampler_name=sampler_name, scheduler=scheduler, - positive=positive, - negative=negative, + positive=regional.positive, + negative=regional.negative, latent_image=latent_image, denoise=denoise, diffusion_mode=diffusion_mode, @@ -118,6 +137,8 @@ class ContextualDiffusionSamplingService: segs=None, image_height=image_height, image_width=image_width, + region_masks=regional.planning_masks, + allow_full_context_masks=regional.active, ) outputs: list[torch.Tensor] = [] @@ -131,14 +152,14 @@ class ContextualDiffusionSamplingService: sampler_name=sampler_name, scheduler=scheduler, positive=( - select_conditioning(positive, index) - if isinstance(positive, ConditioningBatch) - else positive + select_conditioning(regional.positive, index) + if isinstance(regional.positive, ConditioningBatch) + else regional.positive ), negative=( - select_conditioning(negative, index) - if isinstance(negative, ConditioningBatch) - else negative + select_conditioning(regional.negative, index) + if isinstance(regional.negative, ConditioningBatch) + else regional.negative ), latent_image=single_item_latent(latent_image, index), denoise=denoise, @@ -151,6 +172,8 @@ class ContextualDiffusionSamplingService: ), image_height=image_height, image_width=image_width, + region_masks=regional.planning_masks, + allow_full_context_masks=regional.active, ) samples = item_result.latent.get("samples") if not isinstance(samples, torch.Tensor): @@ -182,6 +205,8 @@ class ContextualDiffusionSamplingService: segs: NativeSegs | None, image_height: int, image_width: int, + region_masks: torch.Tensor | None, + allow_full_context_masks: bool, ) -> ContextualDiffusionSamplingResult: """Build one canvas plan and execute it through the runtime adapter.""" @@ -195,6 +220,7 @@ class ContextualDiffusionSamplingService: latent_height=int(samples.shape[-2]), controls=controls, segs=segs, + region_masks=region_masks, ) latent = sample_contextual_diffusion( model=model, @@ -210,6 +236,7 @@ class ContextualDiffusionSamplingService: diffusion_mode=diffusion_mode, controls=controls, plan=plan, + allow_full_context_masks=allow_full_context_masks, ) return ContextualDiffusionSamplingResult( latent=latent, diff --git a/simple_syrup/services/regional_conditioning_service.py b/simple_syrup/services/regional_conditioning_service.py index df5f757..697c470 100644 --- a/simple_syrup/services/regional_conditioning_service.py +++ b/simple_syrup/services/regional_conditioning_service.py @@ -42,15 +42,36 @@ class RegionalConditioningService: validate_regional_prompt_weight(regional_prompt_weight) mask_batch = prepare_regional_mask_batch(masks, region_mask_feather) + return self.assemble_prepared( + positive=positive, + negative=negative, + mask_batch=mask_batch, + regional_prompt_weight=regional_prompt_weight, + region_mask_feather=region_mask_feather, + ) + + def assemble_prepared( + self, + *, + positive: object, + negative: object, + mask_batch: torch.Tensor, + regional_prompt_weight: float, + region_mask_feather: int = 0, + ) -> tuple[Conditioning, Conditioning]: + """Assemble conditioning from a validated, already-feathered mask batch.""" + + validate_regional_prompt_weight(regional_prompt_weight) + prepared_mask_batch = prepare_regional_mask_batch(mask_batch, feather=0) assembled_positive = self._assemble_input( positive, - mask_batch, + prepared_mask_batch, regional_prompt_weight=regional_prompt_weight, input_name="positive", ) assembled_negative = self._assemble_input( negative, - mask_batch, + prepared_mask_batch, regional_prompt_weight=regional_prompt_weight, input_name="negative", ) @@ -58,7 +79,7 @@ class RegionalConditioningService: "Regional conditioning assembled", extra={ "operation": "assemble_regional_conditioning", - "region_count": int(mask_batch.shape[0]), + "region_count": int(prepared_mask_batch.shape[0]), "positive_entry_count": len(assembled_positive), "negative_entry_count": len(assembled_negative), "regional_prompt_weight": regional_prompt_weight, diff --git a/simple_syrup/services/regional_sampling_preparation_service.py b/simple_syrup/services/regional_sampling_preparation_service.py new file mode 100644 index 0000000..f4cb9c7 --- /dev/null +++ b/simple_syrup/services/regional_sampling_preparation_service.py @@ -0,0 +1,112 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Prepare optional regional conditioning before sampler batch interpretation.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, ClassVar, TypeAlias + +import torch + +from ..domain.conditioning_batch import ConditioningBatch +from ..masking.regional_prompt_masks import ( + prepare_regional_masks, + resize_regional_mask_batch, +) +from ..shared.logging import get_logger +from .regional_conditioning_service import RegionalConditioningService + +Latent: TypeAlias = dict[str, Any] +LOGGER = get_logger(__name__) + + +@dataclass(frozen=True) +class RegionalSamplingPreparation: + """Carry resolved conditioning and latent-space planning masks.""" + + positive: object + negative: object + planning_masks: torch.Tensor | None + + @property + def active(self) -> bool: + """Return whether regional sampling owns conditioning interpretation.""" + + return self.planning_masks is not None + + +class RegionalSamplingPreparationService: + """Activate and assemble regional sampling from masks plus a condition batch.""" + + conditioning_service_class: ClassVar[type[RegionalConditioningService]] = ( + RegionalConditioningService + ) + + def prepare( + self, + *, + positive: object, + negative: object, + latent_image: Latent, + region_masks: object | None, + regional_prompt_weight: float, + region_mask_feather: int, + ) -> RegionalSamplingPreparation: + """Return unchanged inputs unless masks explicitly disambiguate a batch.""" + + uses_batch = isinstance(positive, ConditioningBatch) or isinstance( + negative, + ConditioningBatch, + ) + if region_masks is None or not uses_batch: + return RegionalSamplingPreparation( + positive=positive, + negative=negative, + planning_masks=None, + ) + samples = latent_image.get("samples") + if not isinstance(samples, torch.Tensor): + raise TypeError("Regional sampling latent samples must be a torch.Tensor.") + if samples.ndim < 4: + raise ValueError( + "Regional sampling latent samples must include height and width axes." + ) + latent_height = int(samples.shape[-2]) + latent_width = int(samples.shape[-1]) + prepared_masks = prepare_regional_masks(region_masks, region_mask_feather) + planning_masks = resize_regional_mask_batch( + prepared_masks.authored, + height=latent_height, + width=latent_width, + ) + conditioning_masks = resize_regional_mask_batch( + prepared_masks.conditioning, + height=latent_height, + width=latent_width, + ) + assembled_positive, assembled_negative = ( + self.conditioning_service_class().assemble_prepared( + positive=positive, + negative=negative, + mask_batch=conditioning_masks, + regional_prompt_weight=regional_prompt_weight, + region_mask_feather=region_mask_feather, + ) + ) + LOGGER.info( + "Regional sampler inputs prepared", + extra={ + "operation": "prepare_regional_sampling", + "region_count": int(planning_masks.shape[0]), + "latent_height": latent_height, + "latent_width": latent_width, + }, + ) + return RegionalSamplingPreparation( + positive=assembled_positive, + negative=assembled_negative, + planning_masks=planning_masks, + ) diff --git a/simple_syrup/services/regional_tiled_diffusion_sampling_service.py b/simple_syrup/services/regional_tiled_diffusion_sampling_service.py new file mode 100644 index 0000000..083cce3 --- /dev/null +++ b/simple_syrup/services/regional_tiled_diffusion_sampling_service.py @@ -0,0 +1,143 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Execute region-constrained tiled diffusion across a latent batch.""" + +from __future__ import annotations + +from typing import Any, Protocol, TypeAlias + +import torch + +from ..domain.regional_tiled_diffusion import ( + build_region_constrained_tiled_diffusion_plan, +) +from ..domain.segs import coerce_segs_group +from ..domain.tiled_diffusion import TiledDiffusionPlan +from ..runtime.detail_previews import DetailPreviewContext +from .sampling_batch import combine_latent_outputs, single_item_latent + +Latent: TypeAlias = dict[str, Any] + + +class TiledDiffusionItemSampler(Protocol): + """Sample one latent item through a supplied tiled diffusion plan.""" + + def __call__( + 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: + """Return one sampled latent item.""" + + +class RegionalTiledDiffusionSamplingService: + """Own regional plan construction and latent-batch execution.""" + + def sample( + self, + *, + item_sampler: TiledDiffusionItemSampler, + region_masks: torch.Tensor, + segs: object | None, + 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, + ) -> Latent: + """Sample every latent item with one shared regional composition.""" + + 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]) + segs_group = coerce_segs_group(segs) if segs is not None else () + if segs_group and len(segs_group) not in (1, batch_size): + raise ValueError( + "Region-constrained tiled diffusion requires one SEGS payload or " + f"one per latent batch item; received {len(segs_group)} SEGS " + f"payloads for batch size {batch_size}." + ) + + outputs: list[torch.Tensor] = [] + for index in range(batch_size): + item_latent = single_item_latent(latent_image, index) + item_samples = item_latent.get("samples") + if not isinstance(item_samples, torch.Tensor): + raise TypeError( + "Tiled diffusion latent samples must be a torch.Tensor." + ) + plan = build_region_constrained_tiled_diffusion_plan( + region_masks=region_masks, + segs=( + segs_group[0 if len(segs_group) == 1 else index] + if segs_group + else None + ), + latent_width=int(item_samples.shape[-1]), + latent_height=int(item_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 = item_sampler( + diffusion_mode=diffusion_mode, + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=positive, + negative=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=True, + tiled_plan=plan, + ) + output_samples = output.get("samples") + if not isinstance(output_samples, torch.Tensor): + raise TypeError( + "Tiled diffusion output samples must be a torch.Tensor." + ) + outputs.append(output_samples) + return combine_latent_outputs(latent_image, outputs) diff --git a/simple_syrup/services/tiled_diffusion_sampling_service.py b/simple_syrup/services/tiled_diffusion_sampling_service.py index e05e5d6..06e29f4 100644 --- a/simple_syrup/services/tiled_diffusion_sampling_service.py +++ b/simple_syrup/services/tiled_diffusion_sampling_service.py @@ -6,7 +6,7 @@ from __future__ import annotations -from typing import Any +from typing import Any, ClassVar import torch @@ -16,6 +16,12 @@ 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 +from .regional_sampling_preparation_service import ( + RegionalSamplingPreparationService, +) +from .regional_tiled_diffusion_sampling_service import ( + RegionalTiledDiffusionSamplingService, +) from .sampling_batch import combine_latent_outputs, single_item_latent Latent = dict[str, Any] @@ -24,6 +30,13 @@ Latent = dict[str, Any] class TiledDiffusionSamplingService: """Route tiled diffusion sampling requests to the selected runtime.""" + regional_preparation_service_class: ClassVar[ + type[RegionalSamplingPreparationService] + ] = RegionalSamplingPreparationService + regional_sampling_service_class: ClassVar[ + type[RegionalTiledDiffusionSamplingService] + ] = RegionalTiledDiffusionSamplingService + def sample( self, *, @@ -46,10 +59,46 @@ class TiledDiffusionSamplingService: differential_diffusion: bool = False, allow_full_context_masks: bool = False, segs: object | None = None, + region_masks: object | None = None, + regional_prompt_weight: float = 0.5, + region_mask_feather: int = 0, ) -> Latent: """Sample a latent with the selected tiled diffusion method.""" validate_tiled_diffusion_mode(diffusion_mode) + regional = self.regional_preparation_service_class().prepare( + positive=positive, + negative=negative, + latent_image=latent_image, + region_masks=region_masks, + regional_prompt_weight=regional_prompt_weight, + region_mask_feather=region_mask_feather, + ) + if regional.active: + if regional.planning_masks is None: + raise RuntimeError("Active regional sampling requires planning masks.") + return self.regional_sampling_service_class().sample( + item_sampler=self._sample_single, + region_masks=regional.planning_masks, + segs=segs, + diffusion_mode=diffusion_mode, + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=regional.positive, + negative=regional.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, + ) if segs is not None: return self._sample_segs_guided( diffusion_mode=diffusion_mode, diff --git a/tests/test_contextual_diffusion_sampling.py b/tests/test_contextual_diffusion_sampling.py index f366c6b..df7b6f6 100644 --- a/tests/test_contextual_diffusion_sampling.py +++ b/tests/test_contextual_diffusion_sampling.py @@ -257,6 +257,24 @@ def test_runtime_delegates_sampling_to_comfy_with_wrapped_clone( ) monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None) monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", False) + mask_support: list[bool] = [] + + def record_conditioning_policy( + _conditioning: object, + *, + sampler_label: str, + allow_full_context_masks: bool = False, + ) -> None: + """Record the regional mask policy at the runtime boundary.""" + + assert sampler_label == "Contextual Diffusion" + mask_support.append(allow_full_context_masks) + + monkeypatch.setattr( + contextual_diffusion_sampling, + "reject_unsupported_conditioning", + record_conditioning_policy, + ) def fake_sample_custom( sampling_model: _FakeModel, @@ -292,6 +310,7 @@ def test_runtime_delegates_sampling_to_comfy_with_wrapped_clone( diffusion_mode="mixture_of_diffusers", controls=controls, plan=plan, + allow_full_context_masks=True, ) assert calls["model"] is not model @@ -301,6 +320,7 @@ def test_runtime_delegates_sampling_to_comfy_with_wrapped_clone( assert output["samples"] is sampled assert output["kept"] == "metadata" assert "downscale_ratio_spacial" not in output + assert mask_support == [True, True] def _controls( diff --git a/tests/test_contextual_diffusion_sampling_service.py b/tests/test_contextual_diffusion_sampling_service.py index 8c071cd..8a1607e 100644 --- a/tests/test_contextual_diffusion_sampling_service.py +++ b/tests/test_contextual_diffusion_sampling_service.py @@ -11,6 +11,7 @@ from typing import Any 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 import ( contextual_diffusion_sampling_service as service_module, @@ -97,6 +98,104 @@ def test_service_rejects_segs_batch_mismatch_before_runtime( assert not called +def test_service_selects_conditioning_batch_per_latent_with_shared_segs( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Contextual SEGS sampling keeps its per-latent conditioning contract.""" + + calls: list[dict[str, Any]] = [] + + def fake_sample_contextual_diffusion(**kwargs: Any) -> dict[str, Any]: + """Record one item and return its latent unchanged.""" + + calls.append(kwargs) + latent_image = kwargs["latent_image"] + if not isinstance(latent_image, dict): + raise TypeError("Test runtime expected a latent dictionary.") + return latent_image + + monkeypatch.setattr( + service_module, + "sample_contextual_diffusion", + fake_sample_contextual_diffusion, + ) + latent = {"samples": torch.zeros((2, 4, 64, 96))} + + ContextualDiffusionSamplingService().sample( + **( + _sample_kwargs(latent=latent, segs=_segs(512, 768)) + | { + "positive": ConditioningBatch(("positive-0", "positive-1")), + "negative": ConditioningBatch(("negative-0", "negative-1")), + } + ) + ) + + assert [call["positive"] for call in calls] == ["positive-0", "positive-1"] + assert [call["negative"] for call in calls] == ["negative-0", "negative-1"] + + +def test_service_applies_regional_conditioning_to_contextual_runtime( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Contextual local and global predictions share assembled regional inputs.""" + + calls: list[dict[str, Any]] = [] + + def fake_sample_contextual_diffusion(**kwargs: Any) -> dict[str, Any]: + """Record one regional contextual request and return its latent.""" + + calls.append(kwargs) + latent_image = kwargs["latent_image"] + if not isinstance(latent_image, dict): + raise TypeError("Test runtime expected a latent dictionary.") + return latent_image + + monkeypatch.setattr( + service_module, + "sample_contextual_diffusion", + fake_sample_contextual_diffusion, + ) + masks = torch.zeros((2, 64, 96)) + masks[0, :, :48] = 1.0 + masks[1, :, 48:] = 1.0 + + ContextualDiffusionSamplingService().sample( + **( + _sample_kwargs( + latent={ + "samples": torch.zeros((1, 4, 64, 96)), + "downscale_ratio_spacial": 8, + }, + segs=None, + ) + | { + "positive": ConditioningBatch( + ( + _conditioning("global"), + _conditioning("left"), + _conditioning("right"), + ) + ), + "negative": _conditioning("negative"), + "region_masks": masks, + } + ) + ) + + assert len(calls) == 1 + assert [item[0] for item in calls[0]["positive"]] == [ + "global", + "left", + "right", + ] + assert calls[0]["allow_full_context_masks"] is True + assert len(calls[0]["plan"].tile_plan.tiles) >= 2 + assert all( + tile.weight_mask is not None for tile in calls[0]["plan"].tile_plan.tiles + ) + + def test_service_rejects_invalid_controls_before_runtime( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -112,7 +211,10 @@ def test_service_rejects_invalid_controls_before_runtime( ContextualDiffusionSamplingService().sample( **( _sample_kwargs( - latent={"samples": torch.zeros((1, 4, 64, 96))}, + latent={ + "samples": torch.zeros((1, 4, 64, 96)), + "downscale_ratio_spacial": 8, + }, segs=None, ) | {"latent_context_overlap": 64} @@ -181,3 +283,9 @@ def _segs(height: int, width: int) -> object: label="subject", ) return ((height, width), (segment,)) + + +def _conditioning(name: str) -> list[list[object]]: + """Return one structurally valid standard conditioning value.""" + + return [[name, {}]] diff --git a/tests/test_ksampler_contextual_diffusion_node.py b/tests/test_ksampler_contextual_diffusion_node.py index 2d420bd..947f8f3 100644 --- a/tests/test_ksampler_contextual_diffusion_node.py +++ b/tests/test_ksampler_contextual_diffusion_node.py @@ -69,6 +69,9 @@ def test_input_types_expose_concise_klein_oriented_controls( assert required["global_steps"][1]["default"] == 1 assert required["global_decay"][1]["default"] == 0.5 assert declared["optional"]["segs"][0] == "SEGS" + assert declared["optional"]["region_masks"][0] == "MASK" + assert declared["optional"]["regional_prompt_weight"][1]["default"] == 0.5 + assert declared["optional"]["region_mask_feather"][1]["default"] == 0 def test_node_metadata_matches_separate_sampler_contract() -> None: @@ -147,6 +150,9 @@ def test_sample_delegates_every_control_to_service( "global_steps": 2, "global_decay": 0.4, "segs": segs, + "region_masks": None, + "regional_prompt_weight": 0.5, + "region_mask_feather": 0, } ] diff --git a/tests/test_ksampler_tiled_diffusion_node.py b/tests/test_ksampler_tiled_diffusion_node.py index 46d62c9..eb96044 100644 --- a/tests/test_ksampler_tiled_diffusion_node.py +++ b/tests/test_ksampler_tiled_diffusion_node.py @@ -60,6 +60,9 @@ def test_input_types_match_tiled_diffusion_contract( assert required["latent_tile_overlap"][1]["default"] == 16 assert required["latent_tile_batch_size"][1]["default"] == 4 assert optional["segs"][0] == "SEGS" + assert optional["region_masks"][0] == "MASK" + assert optional["regional_prompt_weight"][1]["default"] == 0.5 + assert optional["region_mask_feather"][1]["default"] == 0 def test_node_metadata_matches_contract() -> None: @@ -121,6 +124,9 @@ def test_sample_delegates_to_shared_service( assert call["latent_tile_batch_size"] == 3 assert call["preview_context"] is None assert call["segs"] is None + assert call["region_masks"] is None + assert call["regional_prompt_weight"] == 0.5 + assert call["region_mask_feather"] == 0 def test_invalid_diffusion_mode_fails_before_runtime_sampling() -> None: @@ -175,6 +181,9 @@ class _FakeTiledDiffusionSamplingService: latent_tile_batch_size: int, preview_context: Any | None = None, segs: object | None = None, + region_masks: object | None = None, + regional_prompt_weight: float = 0.5, + region_mask_feather: int = 0, ) -> dict[str, Any]: """Record sampling arguments and return a fixed latent.""" @@ -197,6 +206,9 @@ class _FakeTiledDiffusionSamplingService: "latent_tile_batch_size": latent_tile_batch_size, "preview_context": preview_context, "segs": segs, + "region_masks": region_masks, + "regional_prompt_weight": regional_prompt_weight, + "region_mask_feather": region_mask_feather, } ) return self.output diff --git a/tests/test_patcher_lifecycle_policy.py b/tests/test_patcher_lifecycle_policy.py index d13922b..f29b699 100644 --- a/tests/test_patcher_lifecycle_policy.py +++ b/tests/test_patcher_lifecycle_policy.py @@ -25,7 +25,7 @@ FORBIDDEN_PATCHER_CALLS = frozenset( FORBIDDEN_PATCHER_WRITES = frozenset({"forced_hooks", "use_clip_schedule"}) APPROVED_VALUE_CLONES = Counter( { - ("simple_syrup/domain/segs_tiled_diffusion.py", "mask"): 2, + ("simple_syrup/domain/semantic_tiled_diffusion.py", "mask"): 2, ("simple_syrup/image/crop_composite.py", "image"): 1, ( "simple_syrup/image/resize_service.py", diff --git a/tests/test_regional_sampling_preparation_service.py b/tests/test_regional_sampling_preparation_service.py new file mode 100644 index 0000000..ef8ad94 --- /dev/null +++ b/tests/test_regional_sampling_preparation_service.py @@ -0,0 +1,101 @@ +# 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 optional regional sampler input preparation.""" + +from __future__ import annotations + +import pytest +import torch + +from simple_syrup.domain.conditioning_batch import ConditioningBatch +from simple_syrup.services.regional_sampling_preparation_service import ( + RegionalSamplingPreparationService, +) + + +def test_preparation_is_inactive_without_both_masks_and_a_batch() -> None: + """Partial regional inputs preserve ordinary sampler interpretation.""" + + service = RegionalSamplingPreparationService() + latent = {"samples": torch.zeros((1, 4, 2, 4))} + + without_masks = service.prepare( + positive=ConditioningBatch((_conditioning("global"),)), + negative=_conditioning("negative"), + latent_image=latent, + region_masks=None, + regional_prompt_weight=0.5, + region_mask_feather=0, + ) + without_batch = service.prepare( + positive=_conditioning("global"), + negative=_conditioning("negative"), + latent_image=latent, + region_masks=object(), + regional_prompt_weight=0.5, + region_mask_feather=0, + ) + + assert not without_masks.active + assert isinstance(without_masks.positive, ConditioningBatch) + assert not without_batch.active + assert without_batch.planning_masks is None + + +def test_preparation_assembles_global_first_batch_at_latent_size() -> None: + """Active regional sampling shares one resized mask batch with its planner.""" + + masks = torch.zeros((2, 4, 8)) + masks[0, :, :4] = 1.0 + masks[1, :, 4:] = 1.0 + + prepared = RegionalSamplingPreparationService().prepare( + positive=ConditioningBatch( + ( + _conditioning("global"), + _conditioning("left"), + _conditioning("right"), + ) + ), + negative=_conditioning("negative"), + latent_image={"samples": torch.zeros((1, 4, 2, 4))}, + region_masks=masks, + regional_prompt_weight=0.75, + region_mask_feather=0, + ) + + assert prepared.active + assert prepared.planning_masks is not None + assert prepared.planning_masks.shape == (2, 2, 4) + assert isinstance(prepared.positive, list) + assert [item[0] for item in prepared.positive] == ["global", "left", "right"] + assert prepared.positive[1][1]["mask"].shape == (1, 2, 4) + assert prepared.positive[1][1]["mask_strength"] == 0.75 + + +def test_preparation_reuses_prompt_by_region_mismatch_policy() -> None: + """Excess regional prompts fail through the authoritative pairing service.""" + + with pytest.raises(ValueError, match="2 regional entries but only 1 authored"): + RegionalSamplingPreparationService().prepare( + positive=ConditioningBatch( + ( + _conditioning("global"), + _conditioning("first"), + _conditioning("excess"), + ) + ), + negative=_conditioning("negative"), + latent_image={"samples": torch.zeros((1, 4, 2, 4))}, + region_masks=torch.ones((1, 4, 8)), + regional_prompt_weight=0.5, + region_mask_feather=0, + ) + + +def _conditioning(name: str) -> list[list[object]]: + """Return one structurally valid standard conditioning value.""" + + return [[name, {}]] diff --git a/tests/test_regional_tiled_diffusion.py b/tests/test_regional_tiled_diffusion.py new file mode 100644 index 0000000..14b0f9f --- /dev/null +++ b/tests/test_regional_tiled_diffusion.py @@ -0,0 +1,103 @@ +# 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 region-constrained semantic tile planning.""" + +from __future__ import annotations + +import torch + +from simple_syrup.domain.regional_tiled_diffusion import ( + build_region_constrained_tiled_diffusion_plan, + regional_composition_ownership_masks, +) +from simple_syrup.domain.segs import BoundingBox, CropRegion, Segment + + +def test_composition_partitions_preserve_overlapping_mask_membership() -> None: + """Every distinct global/regional conditioning combination owns one zone.""" + + masks = torch.zeros((2, 4, 8)) + masks[0, :, :6] = 1.0 + masks[1, :, 2:] = 1.0 + + ownership = regional_composition_ownership_masks( + masks, + latent_height=4, + latent_width=8, + ) + + assert len(ownership) == 3 + coverage = torch.stack([mask.to(dtype=torch.int32) for mask in ownership]).sum( + dim=0 + ) + assert torch.equal(coverage, torch.ones((4, 8), dtype=torch.int32)) + + +def test_region_plan_never_merges_across_composition_boundaries() -> None: + """Broad authored regions remain distinct even when one window could hold both.""" + + masks = torch.zeros((2, 4, 8)) + masks[0, :, :4] = 1.0 + masks[1, :, 4:] = 1.0 + + plan = build_region_constrained_tiled_diffusion_plan( + region_masks=masks, + segs=None, + latent_width=8, + latent_height=4, + tile_width=8, + tile_height=4, + overlap=0, + tile_batch_size=4, + ) + + assert len(plan.tiles) == 2 + coverage = torch.zeros((4, 8), 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 torch.equal(coverage, torch.ones((4, 8))) + + +def test_segs_subdivide_each_regional_composition_zone() -> None: + """SEGS geometry is intersected with region ownership instead of replacing it.""" + + masks = torch.zeros((2, 8, 8)) + masks[0, :, :4] = 1.0 + masks[1, :, 4:] = 1.0 + seg_mask = torch.zeros((8, 8)) + seg_mask[:4, :] = 1.0 + segs = ((8, 8), (_segment(seg_mask),)) + + plan = build_region_constrained_tiled_diffusion_plan( + region_masks=masks, + segs=segs, + latent_width=8, + latent_height=8, + tile_width=8, + tile_height=8, + overlap=0, + tile_batch_size=4, + ) + + assert len(plan.tiles) == 4 + assert all(tile.weight_mask is not None for tile in plan.tiles) + + +def _segment(mask: torch.Tensor) -> Segment: + """Return one full-canvas SEG for intersection tests.""" + + height, width = mask.shape + crop = CropRegion(0, 0, width, height) + return Segment( + cropped_image=None, + cropped_mask=mask, + confidence=1.0, + crop_region=crop, + bbox=BoundingBox(*crop), + label="top", + ) diff --git a/tests/test_tiled_diffusion_sampling_service.py b/tests/test_tiled_diffusion_sampling_service.py index 5a85164..3ae50e4 100644 --- a/tests/test_tiled_diffusion_sampling_service.py +++ b/tests/test_tiled_diffusion_sampling_service.py @@ -249,6 +249,121 @@ def test_service_builds_and_forwards_a_segs_guided_plan( assert all(tile.weight_mask is not None for tile in plan.tiles) +def test_service_selects_conditioning_batch_per_latent_with_shared_segs( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """SEGS guidance preserves the established per-latent batch interpretation.""" + + calls: list[dict[str, Any]] = [] + + def fake_multidiffusion(**kwargs: Any) -> dict[str, Any]: + """Record one latent item and return it 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") + | { + "positive": ConditioningBatch(("positive-0", "positive-1")), + "negative": ConditioningBatch(("negative-0", "negative-1")), + "latent_image": {"samples": torch.zeros((2, 4, 4, 4))}, + "segs": segs, + } + ) + ) + + assert [call["positive"] for call in calls] == ["positive-0", "positive-1"] + assert [call["negative"] for call in calls] == ["negative-0", "negative-1"] + assert all(call["tiled_plan"] is not None for call in calls) + + +def test_service_activates_regional_conditioning_and_constrained_tiles( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Masks plus a conditioning batch activate shared regional sampling.""" + + calls: list[dict[str, Any]] = [] + + def fake_multidiffusion(**kwargs: Any) -> dict[str, Any]: + """Record assembled conditioning and its region-constrained plan.""" + + calls.append(kwargs) + return {"samples": kwargs["latent_image"]["samples"]} + + monkeypatch.setattr( + "simple_syrup.services.tiled_diffusion_sampling_service." + "multidiffusion_sampling.sample_multidiffusion", + fake_multidiffusion, + ) + masks = torch.zeros((2, 4, 4)) + masks[0, :, :2] = 1.0 + masks[1, :, 2:] = 1.0 + + TiledDiffusionSamplingService().sample( + **( + _sample_kwargs(diffusion_mode="multidiffusion") + | { + "positive": ConditioningBatch( + ( + _conditioning("global"), + _conditioning("left"), + _conditioning("right"), + ) + ), + "negative": _conditioning("negative"), + "region_masks": masks, + } + ) + ) + + assert len(calls) == 1 + assert [item[0] for item in calls[0]["positive"]] == [ + "global", + "left", + "right", + ] + assert calls[0]["allow_full_context_masks"] is True + assert len(calls[0]["tiled_plan"].tiles) == 2 + + +def test_region_masks_without_conditioning_batch_leave_sampling_unchanged( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Masks alone neither validate nor alter the established ordinary path.""" + + calls: list[dict[str, Any]] = [] + + def fake_multidiffusion(**kwargs: Any) -> dict[str, Any]: + """Record the unchanged runtime request.""" + + calls.append(kwargs) + return {"samples": kwargs["latent_image"]["samples"]} + + monkeypatch.setattr( + "simple_syrup.services.tiled_diffusion_sampling_service." + "multidiffusion_sampling.sample_multidiffusion", + fake_multidiffusion, + ) + + TiledDiffusionSamplingService().sample( + **(_sample_kwargs(diffusion_mode="multidiffusion") | {"region_masks": object()}) + ) + + assert len(calls) == 1 + assert calls[0]["positive"] == "positive" + assert calls[0]["tiled_plan"] is None + assert calls[0]["allow_full_context_masks"] is False + + def test_invalid_mode_fails_before_runtime_call( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -318,3 +433,9 @@ def _full_segment(height: int, width: int) -> Segment: bbox=BoundingBox(0, 0, width, height), label="region", ) + + +def _conditioning(name: str) -> list[list[object]]: + """Return one structurally valid standard conditioning value.""" + + return [[name, {}]]