feat(sampling): add regional diffusion sampling
This commit is contained in:
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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,)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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, {}]]
|
||||
@@ -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, {}]]
|
||||
|
||||
Reference in New Issue
Block a user