feat(sampling): expose evaluated context SEGS
This commit is contained in:
@@ -0,0 +1,170 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Project evaluated latent context windows into lazily materialized SEGS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from collections.abc import Iterable, Iterator, Sequence
|
||||
from threading import Lock
|
||||
from typing import TypeAlias, overload
|
||||
|
||||
import torch
|
||||
|
||||
from .segs import BoundingBox, CropRegion, Segment, SegsHeader
|
||||
from .tiled_diffusion import TiledDiffusionPlan
|
||||
|
||||
ContextSegs: TypeAlias = tuple[SegsHeader, "ContextSegmentSequence"]
|
||||
|
||||
|
||||
class ContextSegmentSequence(Sequence[Segment]):
|
||||
"""Delay large rectangular mask allocation until a SEGS consumer reads it."""
|
||||
|
||||
def __init__(self, windows: Iterable[CropRegion]) -> None:
|
||||
"""Store deterministic context windows without allocating their masks."""
|
||||
|
||||
self._windows = tuple(windows)
|
||||
self._materialized: tuple[Segment, ...] | None = None
|
||||
self._lock = Lock()
|
||||
|
||||
@property
|
||||
def windows(self) -> tuple[CropRegion, ...]:
|
||||
"""Return immutable projected windows without forcing mask allocation."""
|
||||
|
||||
return self._windows
|
||||
|
||||
@property
|
||||
def is_materialized(self) -> bool:
|
||||
"""Report whether a downstream consumer has requested concrete segments."""
|
||||
|
||||
return self._materialized is not None
|
||||
|
||||
def __len__(self) -> int:
|
||||
"""Return the number of contexts without materializing masks."""
|
||||
|
||||
return len(self._windows)
|
||||
|
||||
@overload
|
||||
def __getitem__(self, index: int) -> Segment: ...
|
||||
|
||||
@overload
|
||||
def __getitem__(self, index: slice) -> tuple[Segment, ...]: ...
|
||||
|
||||
def __getitem__(self, index: int | slice) -> Segment | tuple[Segment, ...]:
|
||||
"""Materialize masks and return one context or a context slice."""
|
||||
|
||||
return self._segments()[index]
|
||||
|
||||
def __iter__(self) -> Iterator[Segment]:
|
||||
"""Materialize masks once and iterate contexts in evaluation order."""
|
||||
|
||||
return iter(self._segments())
|
||||
|
||||
def _segments(self) -> tuple[Segment, ...]:
|
||||
"""Create full rectangular masks once on first downstream access."""
|
||||
|
||||
if self._materialized is not None:
|
||||
return self._materialized
|
||||
with self._lock:
|
||||
if self._materialized is None:
|
||||
self._materialized = tuple(
|
||||
_segment_from_window(window, index)
|
||||
for index, window in enumerate(self._windows, start=1)
|
||||
)
|
||||
return self._materialized
|
||||
|
||||
|
||||
def context_segs_from_tile_plan(
|
||||
plan: TiledDiffusionPlan,
|
||||
*,
|
||||
image_height: int,
|
||||
image_width: int,
|
||||
) -> ContextSegs:
|
||||
"""Return lazy rectangular SEGS for every non-global evaluated tile."""
|
||||
|
||||
_validate_image_dimensions(image_height, image_width)
|
||||
windows = tuple(
|
||||
_project_tile(
|
||||
x=tile.x,
|
||||
y=tile.y,
|
||||
width=tile.width,
|
||||
height=tile.height,
|
||||
latent_width=plan.latent_width,
|
||||
latent_height=plan.latent_height,
|
||||
image_width=image_width,
|
||||
image_height=image_height,
|
||||
)
|
||||
for tile in plan.tiles
|
||||
)
|
||||
return (image_height, image_width), ContextSegmentSequence(windows)
|
||||
|
||||
|
||||
def merge_context_segs(values: Iterable[ContextSegs]) -> ContextSegs:
|
||||
"""Combine batched context windows without duplicates or mask allocation."""
|
||||
|
||||
items = tuple(values)
|
||||
if not items:
|
||||
raise ValueError("Context SEGS requires at least one context plan.")
|
||||
header = items[0][0]
|
||||
if any(item[0] != header for item in items[1:]):
|
||||
raise ValueError(
|
||||
"Context SEGS requires batched image dimensions to match exactly."
|
||||
)
|
||||
windows = list(items[0][1].windows)
|
||||
seen = set(windows)
|
||||
for _header, segments in items[1:]:
|
||||
for window in segments.windows:
|
||||
if window not in seen:
|
||||
windows.append(window)
|
||||
seen.add(window)
|
||||
return header, ContextSegmentSequence(windows)
|
||||
|
||||
|
||||
def _project_tile(
|
||||
*,
|
||||
x: int,
|
||||
y: int,
|
||||
width: int,
|
||||
height: int,
|
||||
latent_width: int,
|
||||
latent_height: int,
|
||||
image_width: int,
|
||||
image_height: int,
|
||||
) -> CropRegion:
|
||||
"""Project one latent rectangle outward into integer image coordinates."""
|
||||
|
||||
left = math.floor(x * image_width / latent_width)
|
||||
top = math.floor(y * image_height / latent_height)
|
||||
right = math.ceil((x + width) * image_width / latent_width)
|
||||
bottom = math.ceil((y + height) * image_height / latent_height)
|
||||
return CropRegion(
|
||||
max(0, min(image_width - 1, left)),
|
||||
max(0, min(image_height - 1, top)),
|
||||
max(1, min(image_width, right)),
|
||||
max(1, min(image_height, bottom)),
|
||||
)
|
||||
|
||||
|
||||
def _segment_from_window(window: CropRegion, index: int) -> Segment:
|
||||
"""Materialize one Impact-compatible full rectangular context mask."""
|
||||
|
||||
return Segment(
|
||||
cropped_image=None,
|
||||
cropped_mask=torch.ones(
|
||||
(window.height, window.width),
|
||||
dtype=torch.uint8,
|
||||
),
|
||||
confidence=1.0,
|
||||
crop_region=window,
|
||||
bbox=BoundingBox(*window),
|
||||
label=f"context_{index:03d}",
|
||||
)
|
||||
|
||||
|
||||
def _validate_image_dimensions(height: int, width: int) -> None:
|
||||
"""Reject image geometry that cannot host projected context rectangles."""
|
||||
|
||||
if height < 1 or width < 1:
|
||||
raise ValueError("Context SEGS image dimensions must be positive.")
|
||||
@@ -22,8 +22,12 @@ MAX_LATENT_CONTEXT_SIZE = 512
|
||||
class KSamplerContextualDiffusion:
|
||||
"""Edit large latents through coordinated global and detailed contexts."""
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
OUTPUT_TOOLTIPS = (tooltips.DENOISED_LATENT_OUTPUT,)
|
||||
RETURN_TYPES = ("LATENT", "SEGS")
|
||||
RETURN_NAMES = ("latent", "contexts_segs")
|
||||
OUTPUT_TOOLTIPS = (
|
||||
tooltips.DENOISED_LATENT_OUTPUT,
|
||||
tooltips.CONTEXTUAL_DIFFUSION_CONTEXTS_OUTPUT,
|
||||
)
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "SimpleSyrup/Sampling"
|
||||
DESCRIPTION = (
|
||||
@@ -202,10 +206,10 @@ class KSamplerContextualDiffusion:
|
||||
global_steps: int = 1,
|
||||
global_decay: float = 0.5,
|
||||
segs: object | None = None,
|
||||
) -> tuple[Latent]:
|
||||
) -> tuple[Latent, object]:
|
||||
"""Delegate contextual diffusion sampling to its application service."""
|
||||
|
||||
output = self.service_class().sample(
|
||||
result = self.service_class().sample(
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
@@ -225,4 +229,4 @@ class KSamplerContextualDiffusion:
|
||||
global_decay=global_decay,
|
||||
segs=segs,
|
||||
)
|
||||
return (output,)
|
||||
return result.latent, result.contexts
|
||||
|
||||
@@ -137,6 +137,10 @@ DENOISE_STRENGTH = (
|
||||
"larger changes."
|
||||
)
|
||||
DENOISED_LATENT_OUTPUT = "Denoised latent for VAE decode or more latent processing."
|
||||
CONTEXTUAL_DIFFUSION_CONTEXTS_OUTPUT = (
|
||||
"Rectangular non-global contexts actually evaluated during sampling. Their "
|
||||
"SEGS masks are created only when this output is connected."
|
||||
)
|
||||
TILED_DIFFUSION_MODE = (
|
||||
"Tile overlap blend. MultiDiffusion averages predictions; Mixture of Diffusers "
|
||||
"gives tile centers more influence."
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Resolve decoded image geometry from ComfyUI latent runtime metadata."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
|
||||
def decoded_image_dimensions(
|
||||
*,
|
||||
model: Any,
|
||||
latent_image: dict[str, Any],
|
||||
latent_height: int,
|
||||
latent_width: int,
|
||||
) -> tuple[int, int]:
|
||||
"""Return decoded height and width using ComfyUI's spatial latent ratio."""
|
||||
|
||||
ratio = latent_image.get("downscale_ratio_spacial")
|
||||
if ratio is None:
|
||||
get_model_object = getattr(model, "get_model_object", None)
|
||||
if not callable(get_model_object):
|
||||
raise ValueError(
|
||||
"Context SEGS requires SEGS input dimensions or a model exposing "
|
||||
"its latent format."
|
||||
)
|
||||
latent_format = get_model_object("latent_format")
|
||||
ratio = getattr(latent_format, "spacial_downscale_ratio", None)
|
||||
if isinstance(ratio, bool) or not isinstance(ratio, (int, float)) or ratio <= 0:
|
||||
raise ValueError(
|
||||
"Context SEGS requires a positive latent spatial downscale ratio."
|
||||
)
|
||||
return round(latent_height * ratio), round(latent_width * ratio)
|
||||
@@ -31,6 +31,7 @@ ModelFunctionWrapper: TypeAlias = Callable[[ApplyModel, dict[str, Any]], torch.T
|
||||
TileEvaluator: TypeAlias = Callable[[dict[str, Any]], torch.Tensor]
|
||||
|
||||
UNSUPPORTED_CONDITIONING_KEYS = frozenset({"area", "control", "gligen"})
|
||||
CONTEXT_INVARIANT_CONDITIONING_KEYS = frozenset({"ref_latents"})
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -463,6 +464,14 @@ def spatial_context_conditioning(
|
||||
tiled_timestep=context_timestep,
|
||||
)
|
||||
continue
|
||||
if key in CONTEXT_INVARIANT_CONDITIONING_KEYS:
|
||||
transformed[key] = repeat_context_invariant_value(
|
||||
value,
|
||||
context_count=len(contexts),
|
||||
input_batch_size=input_batch_size,
|
||||
conditioning_key=key,
|
||||
)
|
||||
continue
|
||||
transformed[key] = spatial_context_value(
|
||||
value,
|
||||
contexts=contexts,
|
||||
@@ -587,6 +596,14 @@ def tile_conditioning(
|
||||
tiled_timestep=tiled_timestep,
|
||||
)
|
||||
continue
|
||||
if key in CONTEXT_INVARIANT_CONDITIONING_KEYS:
|
||||
tiled[key] = repeat_context_invariant_value(
|
||||
value,
|
||||
context_count=len(tiles),
|
||||
input_batch_size=input_batch_size,
|
||||
conditioning_key=key,
|
||||
)
|
||||
continue
|
||||
tiled[key] = tile_value(
|
||||
value,
|
||||
tiles=tiles,
|
||||
@@ -597,6 +614,52 @@ def tile_conditioning(
|
||||
return tiled
|
||||
|
||||
|
||||
def repeat_context_invariant_value(
|
||||
value: Any,
|
||||
*,
|
||||
context_count: int,
|
||||
input_batch_size: int,
|
||||
conditioning_key: str,
|
||||
) -> Any:
|
||||
"""Repeat non-spatial conditioning without cropping its tensor contents."""
|
||||
|
||||
if isinstance(value, torch.Tensor):
|
||||
if value.ndim < 1:
|
||||
raise ValueError(
|
||||
f"{conditioning_key} tensors must include a batch dimension."
|
||||
)
|
||||
if value.shape[0] == input_batch_size:
|
||||
return torch.cat([value] * context_count, dim=0)
|
||||
if value.shape[0] == 1:
|
||||
repeats = [input_batch_size * context_count] + [1] * (value.ndim - 1)
|
||||
return value.repeat(repeats)
|
||||
raise ValueError(
|
||||
f"{conditioning_key} tensor batch size must be 1 or match the model "
|
||||
f"input batch size {input_batch_size}; received {value.shape[0]}."
|
||||
)
|
||||
if isinstance(value, list):
|
||||
return [
|
||||
repeat_context_invariant_value(
|
||||
item,
|
||||
context_count=context_count,
|
||||
input_batch_size=input_batch_size,
|
||||
conditioning_key=conditioning_key,
|
||||
)
|
||||
for item in value
|
||||
]
|
||||
if isinstance(value, tuple):
|
||||
return tuple(
|
||||
repeat_context_invariant_value(
|
||||
item,
|
||||
context_count=context_count,
|
||||
input_batch_size=input_batch_size,
|
||||
conditioning_key=conditioning_key,
|
||||
)
|
||||
for item in value
|
||||
)
|
||||
return value
|
||||
|
||||
|
||||
def tile_transformer_options(
|
||||
options: dict[str, Any],
|
||||
*,
|
||||
|
||||
@@ -6,11 +6,17 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, TypeAlias
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.conditioning_batch import ConditioningBatch, select_conditioning
|
||||
from ..domain.context_segs import (
|
||||
ContextSegs,
|
||||
context_segs_from_tile_plan,
|
||||
merge_context_segs,
|
||||
)
|
||||
from ..domain.contextual_diffusion import (
|
||||
ContextualDiffusionControls,
|
||||
build_contextual_diffusion_plan,
|
||||
@@ -18,6 +24,7 @@ from ..domain.contextual_diffusion import (
|
||||
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 .sampling_batch import (
|
||||
combine_latent_outputs,
|
||||
latent_batch_size,
|
||||
@@ -27,6 +34,14 @@ from .sampling_batch import (
|
||||
Latent: TypeAlias = dict[str, Any]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ContextualDiffusionSamplingResult:
|
||||
"""Return the sampled latent and lazy non-global context SEGS."""
|
||||
|
||||
latent: Latent
|
||||
contexts: ContextSegs
|
||||
|
||||
|
||||
class ContextualDiffusionSamplingService:
|
||||
"""Plan and execute composition-preserving contextual diffusion."""
|
||||
|
||||
@@ -51,7 +66,7 @@ class ContextualDiffusionSamplingService:
|
||||
global_steps: int,
|
||||
global_decay: float,
|
||||
segs: object | None = None,
|
||||
) -> Latent:
|
||||
) -> ContextualDiffusionSamplingResult:
|
||||
"""Sample a latent with global and bounded detail contexts."""
|
||||
|
||||
validate_tiled_diffusion_mode(diffusion_mode)
|
||||
@@ -71,6 +86,16 @@ class ContextualDiffusionSamplingService:
|
||||
"Contextual Diffusion requires one SEGS payload or one per latent "
|
||||
f"batch item; received {len(segs_group)} for batch size {batch_size}."
|
||||
)
|
||||
image_height, image_width = (
|
||||
segs_group[0][0]
|
||||
if segs_group
|
||||
else decoded_image_dimensions(
|
||||
model=model,
|
||||
latent_image=latent_image,
|
||||
latent_height=int(latent_image["samples"].shape[-2]),
|
||||
latent_width=int(latent_image["samples"].shape[-1]),
|
||||
)
|
||||
)
|
||||
split_batch = (
|
||||
bool(segs_group)
|
||||
or isinstance(positive, ConditioningBatch)
|
||||
@@ -91,11 +116,14 @@ class ContextualDiffusionSamplingService:
|
||||
diffusion_mode=diffusion_mode,
|
||||
controls=controls,
|
||||
segs=None,
|
||||
image_height=image_height,
|
||||
image_width=image_width,
|
||||
)
|
||||
|
||||
outputs: list[torch.Tensor] = []
|
||||
contexts: list[ContextSegs] = []
|
||||
for index in range(batch_size):
|
||||
item_output = self._sample_item(
|
||||
item_result = self._sample_item(
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
@@ -121,14 +149,20 @@ class ContextualDiffusionSamplingService:
|
||||
if segs_group
|
||||
else None
|
||||
),
|
||||
image_height=image_height,
|
||||
image_width=image_width,
|
||||
)
|
||||
samples = item_output.get("samples")
|
||||
samples = item_result.latent.get("samples")
|
||||
if not isinstance(samples, torch.Tensor):
|
||||
raise TypeError(
|
||||
"Contextual Diffusion output samples must be a torch.Tensor."
|
||||
)
|
||||
outputs.append(samples)
|
||||
return combine_latent_outputs(latent_image, outputs)
|
||||
contexts.append(item_result.contexts)
|
||||
return ContextualDiffusionSamplingResult(
|
||||
latent=combine_latent_outputs(latent_image, outputs),
|
||||
contexts=merge_context_segs(contexts),
|
||||
)
|
||||
|
||||
def _sample_item(
|
||||
self,
|
||||
@@ -146,7 +180,9 @@ class ContextualDiffusionSamplingService:
|
||||
diffusion_mode: str,
|
||||
controls: ContextualDiffusionControls,
|
||||
segs: NativeSegs | None,
|
||||
) -> Latent:
|
||||
image_height: int,
|
||||
image_width: int,
|
||||
) -> ContextualDiffusionSamplingResult:
|
||||
"""Build one canvas plan and execute it through the runtime adapter."""
|
||||
|
||||
samples = latent_image.get("samples")
|
||||
@@ -160,7 +196,7 @@ class ContextualDiffusionSamplingService:
|
||||
controls=controls,
|
||||
segs=segs,
|
||||
)
|
||||
return sample_contextual_diffusion(
|
||||
latent = sample_contextual_diffusion(
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
@@ -175,3 +211,11 @@ class ContextualDiffusionSamplingService:
|
||||
controls=controls,
|
||||
plan=plan,
|
||||
)
|
||||
return ContextualDiffusionSamplingResult(
|
||||
latent=latent,
|
||||
contexts=context_segs_from_tile_plan(
|
||||
plan.tile_plan,
|
||||
image_height=image_height,
|
||||
image_width=image_width,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
# 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 projecting authoritative sampling contexts into SEGS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.context_segs import context_segs_from_tile_plan
|
||||
from simple_syrup.domain.segs import BoundingBox, coerce_segs
|
||||
from simple_syrup.domain.tiled_diffusion import build_tiled_diffusion_plan
|
||||
from simple_syrup.services.simple_preview_segs_service import SimplePreviewSEGSService
|
||||
|
||||
|
||||
def test_tile_plan_projects_exact_non_global_context_rectangles() -> None:
|
||||
"""Every evaluated tile becomes one full rectangular SEG in image space."""
|
||||
|
||||
plan = build_tiled_diffusion_plan(
|
||||
latent_width=64,
|
||||
latent_height=32,
|
||||
tile_width=32,
|
||||
tile_height=24,
|
||||
overlap=8,
|
||||
tile_batch_size=2,
|
||||
)
|
||||
|
||||
header, contexts = context_segs_from_tile_plan(
|
||||
plan,
|
||||
image_height=256,
|
||||
image_width=512,
|
||||
)
|
||||
|
||||
assert header == (256, 512)
|
||||
assert not contexts.is_materialized
|
||||
assert len(contexts) == len(plan.tiles)
|
||||
assert not contexts.is_materialized
|
||||
assert [context.label for context in contexts] == [
|
||||
f"context_{index:03d}" for index in range(1, len(contexts) + 1)
|
||||
]
|
||||
assert all(context.crop_region != (0, 0, 512, 256) for context in contexts)
|
||||
for tile, context in zip(plan.tiles, contexts, strict=True):
|
||||
assert context.crop_region == (
|
||||
tile.x * 8,
|
||||
tile.y * 8,
|
||||
(tile.x + tile.width) * 8,
|
||||
(tile.y + tile.height) * 8,
|
||||
)
|
||||
assert context.bbox == BoundingBox(*context.crop_region)
|
||||
mask = context.cropped_mask
|
||||
assert isinstance(mask, torch.Tensor)
|
||||
assert mask.dtype == torch.uint8
|
||||
assert tuple(mask.shape) == (
|
||||
context.crop_region.height,
|
||||
context.crop_region.width,
|
||||
)
|
||||
assert bool(mask.all())
|
||||
assert contexts.is_materialized
|
||||
|
||||
|
||||
def test_lazy_contexts_feed_simple_preview_segs_directly() -> None:
|
||||
"""Simple Preview SEGS materializes and displays a connected context output."""
|
||||
|
||||
plan = build_tiled_diffusion_plan(
|
||||
latent_width=16,
|
||||
latent_height=8,
|
||||
tile_width=8,
|
||||
tile_height=8,
|
||||
overlap=2,
|
||||
tile_batch_size=2,
|
||||
)
|
||||
contexts_segs = context_segs_from_tile_plan(
|
||||
plan,
|
||||
image_height=64,
|
||||
image_width=128,
|
||||
)
|
||||
|
||||
document = SimplePreviewSEGSService().build(
|
||||
image=torch.zeros((1, 64, 128, 3)),
|
||||
segs=coerce_segs(contexts_segs),
|
||||
)
|
||||
|
||||
assert contexts_segs[1].is_materialized
|
||||
assert len(document.regions) == len(plan.tiles)
|
||||
assert all(region.label.startswith("context_") for region in document.regions)
|
||||
@@ -119,6 +119,59 @@ def test_global_prediction_decays_then_stops_after_configured_steps() -> None:
|
||||
assert torch.allclose(output, torch.ones_like(output))
|
||||
|
||||
|
||||
def test_every_prediction_receives_each_complete_reference_image() -> None:
|
||||
"""Local and global predictions preserve ordered multi-image conditioning."""
|
||||
|
||||
controls = _controls()
|
||||
plan = build_contextual_diffusion_plan(
|
||||
latent_width=32,
|
||||
latent_height=16,
|
||||
controls=controls,
|
||||
segs=None,
|
||||
)
|
||||
wrapper = ContextualDiffusionModelWrapper(
|
||||
plan=plan,
|
||||
controls=controls,
|
||||
sigmas=torch.tensor([1.0, 0.0]),
|
||||
existing_wrapper=None,
|
||||
)
|
||||
image_1 = torch.arange(1 * 4 * 16 * 32, dtype=torch.float32).reshape((1, 4, 16, 32))
|
||||
image_2 = torch.full((1, 4, 10, 12), 2.0)
|
||||
received: list[list[torch.Tensor]] = []
|
||||
|
||||
def apply_model(
|
||||
x: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
**conditioning: object,
|
||||
) -> torch.Tensor:
|
||||
"""Capture reference batches presented to each bounded prediction."""
|
||||
|
||||
del timestep
|
||||
references = conditioning["ref_latents"]
|
||||
assert isinstance(references, list)
|
||||
assert all(isinstance(reference, torch.Tensor) for reference in references)
|
||||
received.append(references)
|
||||
return torch.zeros_like(x)
|
||||
|
||||
wrapper(
|
||||
apply_model,
|
||||
{
|
||||
"input": torch.zeros((1, 1, 16, 32)),
|
||||
"timestep": torch.tensor([1.0]),
|
||||
"c": {
|
||||
"ref_latents": [image_1, image_2],
|
||||
"ref_latents_method": "index",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
assert len(received) == 2
|
||||
assert torch.equal(received[0][0], torch.cat([image_1, image_1], dim=0))
|
||||
assert torch.equal(received[0][1], torch.cat([image_2, image_2], dim=0))
|
||||
assert torch.equal(received[1][0], image_1)
|
||||
assert torch.equal(received[1][1], image_2)
|
||||
|
||||
|
||||
def test_native_sized_canvas_delegates_to_one_original_model_call() -> None:
|
||||
"""A canvas already inside the model view limit behaves like normal sampling."""
|
||||
|
||||
|
||||
@@ -47,7 +47,12 @@ def test_service_builds_global_context_and_segs_guided_tile_plan(
|
||||
**_sample_kwargs(latent=latent, segs=_segs(512, 768))
|
||||
)
|
||||
|
||||
assert torch.equal(result["samples"], latent["samples"])
|
||||
assert torch.equal(result.latent["samples"], latent["samples"])
|
||||
assert result.contexts[0] == (512, 768)
|
||||
assert len(result.contexts[1]) == len(calls[0]["plan"].tile_plan.tiles)
|
||||
assert not result.contexts[1].is_materialized
|
||||
assert all(segment.label.startswith("context_") for segment in result.contexts[1])
|
||||
assert result.contexts[1].is_materialized
|
||||
assert len(calls) == 1
|
||||
assert calls[0]["diffusion_mode"] == "mixture_of_diffusers"
|
||||
plan = calls[0]["plan"]
|
||||
|
||||
@@ -11,9 +11,13 @@ from typing import Any
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.segs import NativeSegs
|
||||
from simple_syrup.nodes.ksampler_contextual_diffusion import (
|
||||
KSamplerContextualDiffusion,
|
||||
)
|
||||
from simple_syrup.nodes_v3.legacy_node_wrappers import (
|
||||
KSamplerContextualDiffusionV3,
|
||||
)
|
||||
from simple_syrup.runtime import sampling_samplers, sampling_schedulers
|
||||
|
||||
|
||||
@@ -68,14 +72,24 @@ def test_input_types_expose_concise_klein_oriented_controls(
|
||||
|
||||
|
||||
def test_node_metadata_matches_separate_sampler_contract() -> None:
|
||||
"""Contextual Diffusion remains a distinct sampler with latent output."""
|
||||
"""Contextual Diffusion exposes its latent and actual local contexts."""
|
||||
|
||||
assert KSamplerContextualDiffusion.RETURN_TYPES == ("LATENT",)
|
||||
assert KSamplerContextualDiffusion.RETURN_TYPES == ("LATENT", "SEGS")
|
||||
assert KSamplerContextualDiffusion.RETURN_NAMES == ("latent", "contexts_segs")
|
||||
assert len(KSamplerContextualDiffusion.OUTPUT_TOOLTIPS) == 2
|
||||
assert KSamplerContextualDiffusion.FUNCTION == "sample"
|
||||
assert KSamplerContextualDiffusion.CATEGORY == "SimpleSyrup/Sampling"
|
||||
assert "composition" in KSamplerContextualDiffusion.DESCRIPTION
|
||||
|
||||
|
||||
def test_v3_schema_names_both_contextual_diffusion_outputs() -> None:
|
||||
"""The exported v3 schema exposes the workflow-facing context socket."""
|
||||
|
||||
schema = KSamplerContextualDiffusionV3.define_schema()
|
||||
|
||||
assert [output.id for output in schema.outputs] == ["latent", "contexts_segs"]
|
||||
|
||||
|
||||
def test_sample_delegates_every_control_to_service(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
@@ -90,7 +104,7 @@ def test_sample_delegates_every_control_to_service(
|
||||
latent = {"samples": torch.zeros((1, 4, 32, 48))}
|
||||
segs = object()
|
||||
|
||||
(result,) = KSamplerContextualDiffusion().sample(
|
||||
result, contexts = KSamplerContextualDiffusion().sample(
|
||||
model="model",
|
||||
seed=12,
|
||||
steps=8,
|
||||
@@ -112,6 +126,7 @@ def test_sample_delegates_every_control_to_service(
|
||||
)
|
||||
|
||||
assert result is fake_service.output
|
||||
assert contexts is fake_service.contexts
|
||||
assert fake_service.calls == [
|
||||
{
|
||||
"model": "model",
|
||||
@@ -143,10 +158,15 @@ class _FakeContextualDiffusionService:
|
||||
"""Create a stable output and empty call history."""
|
||||
|
||||
self.output: dict[str, Any] = {"samples": torch.ones((1, 4, 32, 48))}
|
||||
self.contexts: NativeSegs = ((256, 384), ())
|
||||
self.calls: list[dict[str, Any]] = []
|
||||
|
||||
def sample(self, **kwargs: Any) -> dict[str, Any]:
|
||||
def sample(self, **kwargs: Any) -> Any:
|
||||
"""Record one call and return the stable latent."""
|
||||
|
||||
self.calls.append(kwargs)
|
||||
return self.output
|
||||
return type(
|
||||
"Result",
|
||||
(),
|
||||
{"latent": self.output, "contexts": self.contexts},
|
||||
)()
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
# 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 decoded image geometry derived from Comfy latent metadata."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import pytest
|
||||
|
||||
from simple_syrup.runtime.latent_geometry import decoded_image_dimensions
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _LatentFormat:
|
||||
"""Expose ComfyUI's intentionally misspelled spatial ratio attribute."""
|
||||
|
||||
spacial_downscale_ratio: int
|
||||
|
||||
|
||||
class _Model:
|
||||
"""Provide one latent format through the ModelPatcher boundary."""
|
||||
|
||||
def get_model_object(self, name: str) -> object:
|
||||
"""Return the requested latent format fixture."""
|
||||
|
||||
assert name == "latent_format"
|
||||
return _LatentFormat(8)
|
||||
|
||||
|
||||
def test_decoded_dimensions_prefer_explicit_latent_ratio() -> None:
|
||||
"""Per-latent metadata overrides the model format for generated geometry."""
|
||||
|
||||
assert decoded_image_dimensions(
|
||||
model=_Model(),
|
||||
latent_image={"downscale_ratio_spacial": 16},
|
||||
latent_height=32,
|
||||
latent_width=48,
|
||||
) == (512, 768)
|
||||
|
||||
|
||||
def test_decoded_dimensions_fall_back_to_model_latent_format() -> None:
|
||||
"""Ordinary latents use the model's authoritative spatial scale."""
|
||||
|
||||
assert decoded_image_dimensions(
|
||||
model=_Model(),
|
||||
latent_image={},
|
||||
latent_height=32,
|
||||
latent_width=48,
|
||||
) == (256, 384)
|
||||
|
||||
|
||||
def test_decoded_dimensions_reject_missing_geometry_source() -> None:
|
||||
"""A model-neutral context output fails clearly when scale is unknowable."""
|
||||
|
||||
with pytest.raises(ValueError, match="latent format"):
|
||||
decoded_image_dimensions(
|
||||
model=object(),
|
||||
latent_image={},
|
||||
latent_height=32,
|
||||
latent_width=48,
|
||||
)
|
||||
@@ -171,6 +171,78 @@ def test_tile_transformer_options_repeats_model_metadata() -> None:
|
||||
assert torch.equal(tiled["sample_sigmas"], torch.tensor([1.0, 0.0]))
|
||||
|
||||
|
||||
def test_tile_conditioning_repeats_complete_reference_latents() -> None:
|
||||
"""Independent reference images remain complete for every tiled model call."""
|
||||
|
||||
canvas_sized_reference = torch.arange(
|
||||
1 * 4 * 4 * 8,
|
||||
dtype=torch.float32,
|
||||
).reshape((1, 4, 4, 8))
|
||||
smaller_reference = torch.full((1, 4, 3, 5), 7.0)
|
||||
tiles = (LatentTile(0, 0, 4, 4), LatentTile(4, 0, 4, 4))
|
||||
|
||||
transformed = tiled_sampling.tile_conditioning(
|
||||
conditioning={
|
||||
"ref_latents": [canvas_sized_reference, smaller_reference],
|
||||
"ref_latents_method": "index",
|
||||
},
|
||||
tiles=tiles,
|
||||
input_batch_size=1,
|
||||
latent_height=4,
|
||||
latent_width=8,
|
||||
tiled_timestep=torch.tensor([1.0, 1.0]),
|
||||
)
|
||||
|
||||
references = transformed["ref_latents"]
|
||||
assert isinstance(references, list)
|
||||
assert torch.equal(
|
||||
references[0],
|
||||
torch.cat([canvas_sized_reference, canvas_sized_reference], dim=0),
|
||||
)
|
||||
assert torch.equal(
|
||||
references[1],
|
||||
torch.cat([smaller_reference, smaller_reference], dim=0),
|
||||
)
|
||||
assert transformed["ref_latents_method"] == "index"
|
||||
|
||||
|
||||
def test_spatial_context_conditioning_repeats_complete_reference_latents() -> None:
|
||||
"""Global and semantic contexts share complete independent reference images."""
|
||||
|
||||
reference = torch.arange(1 * 4 * 8 * 12, dtype=torch.float32).reshape((1, 4, 8, 12))
|
||||
contexts = (
|
||||
SpatialContext(0, 0, 12, 8, 6, 4),
|
||||
SpatialContext(2, 1, 8, 6, 6, 4),
|
||||
)
|
||||
|
||||
transformed = tiled_sampling.spatial_context_conditioning(
|
||||
conditioning={"ref_latents": [reference]},
|
||||
contexts=contexts,
|
||||
input_batch_size=1,
|
||||
latent_height=8,
|
||||
latent_width=12,
|
||||
context_timestep=torch.tensor([1.0, 1.0]),
|
||||
)
|
||||
|
||||
references = transformed["ref_latents"]
|
||||
assert isinstance(references, list)
|
||||
assert torch.equal(references[0], torch.cat([reference, reference], dim=0))
|
||||
|
||||
|
||||
def test_reference_latents_reject_ambiguous_batch_alignment() -> None:
|
||||
"""Reference batches must align explicitly with each model input batch."""
|
||||
|
||||
with pytest.raises(ValueError, match="ref_latents tensor batch size"):
|
||||
tiled_sampling.tile_conditioning(
|
||||
conditioning={"ref_latents": [torch.zeros((3, 4, 8, 12))]},
|
||||
tiles=(LatentTile(0, 0, 6, 8), LatentTile(6, 0, 6, 8)),
|
||||
input_batch_size=2,
|
||||
latent_height=8,
|
||||
latent_width=12,
|
||||
tiled_timestep=torch.ones((4,)),
|
||||
)
|
||||
|
||||
|
||||
def test_spatial_context_args_resize_latent_and_canvas_conditioning() -> None:
|
||||
"""A global context resizes spatial tensors and repeats aligned metadata."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user