feat(sampling): add regional diffusion sampling

This commit is contained in:
Artificial Sweetener
2026-08-10 00:55:03 -04:00
parent f031c28589
commit ace5fafda9
22 changed files with 1502 additions and 353 deletions
+19 -5
View File
@@ -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,
@@ -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
+21 -325
View File
@@ -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,
)
@@ -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,
)
+44 -3
View File
@@ -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:
@@ -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
@@ -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,)
+12
View File
@@ -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 "
@@ -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(
@@ -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,
@@ -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,
@@ -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,
)
@@ -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)
@@ -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,
@@ -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(
@@ -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, {}]]
@@ -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,
}
]
@@ -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
+1 -1
View File
@@ -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",
@@ -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, {}]]
+103
View File
@@ -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",
)
@@ -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, {}]]