From 652ae51fc4ce1d952161b0c144daea516a7452e6 Mon Sep 17 00:00:00 2001 From: Artificial Sweetener Date: Mon, 3 Aug 2026 21:55:16 -0400 Subject: [PATCH] feat(sampling): expose evaluated context SEGS --- simple_syrup/domain/context_segs.py | 170 ++++++++++++++++++ .../nodes/ksampler_contextual_diffusion.py | 14 +- simple_syrup/nodes/tooltips.py | 4 + simple_syrup/runtime/latent_geometry.py | 35 ++++ simple_syrup/runtime/tiled_sampling.py | 63 +++++++ .../contextual_diffusion_sampling_service.py | 56 +++++- tests/test_context_segs.py | 86 +++++++++ tests/test_contextual_diffusion_sampling.py | 53 ++++++ ...t_contextual_diffusion_sampling_service.py | 7 +- ...test_ksampler_contextual_diffusion_node.py | 30 +++- tests/test_latent_geometry.py | 64 +++++++ tests/test_tiled_sampling_runtime.py | 72 ++++++++ 12 files changed, 637 insertions(+), 17 deletions(-) create mode 100644 simple_syrup/domain/context_segs.py create mode 100644 simple_syrup/runtime/latent_geometry.py create mode 100644 tests/test_context_segs.py create mode 100644 tests/test_latent_geometry.py diff --git a/simple_syrup/domain/context_segs.py b/simple_syrup/domain/context_segs.py new file mode 100644 index 0000000..1da5bf3 --- /dev/null +++ b/simple_syrup/domain/context_segs.py @@ -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.") diff --git a/simple_syrup/nodes/ksampler_contextual_diffusion.py b/simple_syrup/nodes/ksampler_contextual_diffusion.py index d187af2..e624354 100644 --- a/simple_syrup/nodes/ksampler_contextual_diffusion.py +++ b/simple_syrup/nodes/ksampler_contextual_diffusion.py @@ -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 diff --git a/simple_syrup/nodes/tooltips.py b/simple_syrup/nodes/tooltips.py index b4ac98b..3fd85b1 100644 --- a/simple_syrup/nodes/tooltips.py +++ b/simple_syrup/nodes/tooltips.py @@ -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." diff --git a/simple_syrup/runtime/latent_geometry.py b/simple_syrup/runtime/latent_geometry.py new file mode 100644 index 0000000..7e5f48f --- /dev/null +++ b/simple_syrup/runtime/latent_geometry.py @@ -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) diff --git a/simple_syrup/runtime/tiled_sampling.py b/simple_syrup/runtime/tiled_sampling.py index b6d6a21..0d0a707 100644 --- a/simple_syrup/runtime/tiled_sampling.py +++ b/simple_syrup/runtime/tiled_sampling.py @@ -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], *, diff --git a/simple_syrup/services/contextual_diffusion_sampling_service.py b/simple_syrup/services/contextual_diffusion_sampling_service.py index 7f14550..622f643 100644 --- a/simple_syrup/services/contextual_diffusion_sampling_service.py +++ b/simple_syrup/services/contextual_diffusion_sampling_service.py @@ -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, + ), + ) diff --git a/tests/test_context_segs.py b/tests/test_context_segs.py new file mode 100644 index 0000000..3fdfddb --- /dev/null +++ b/tests/test_context_segs.py @@ -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) diff --git a/tests/test_contextual_diffusion_sampling.py b/tests/test_contextual_diffusion_sampling.py index 92291a0..cf13be3 100644 --- a/tests/test_contextual_diffusion_sampling.py +++ b/tests/test_contextual_diffusion_sampling.py @@ -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.""" diff --git a/tests/test_contextual_diffusion_sampling_service.py b/tests/test_contextual_diffusion_sampling_service.py index a7522ec..8c071cd 100644 --- a/tests/test_contextual_diffusion_sampling_service.py +++ b/tests/test_contextual_diffusion_sampling_service.py @@ -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"] diff --git a/tests/test_ksampler_contextual_diffusion_node.py b/tests/test_ksampler_contextual_diffusion_node.py index 944945b..2d420bd 100644 --- a/tests/test_ksampler_contextual_diffusion_node.py +++ b/tests/test_ksampler_contextual_diffusion_node.py @@ -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}, + )() diff --git a/tests/test_latent_geometry.py b/tests/test_latent_geometry.py new file mode 100644 index 0000000..2fbedfd --- /dev/null +++ b/tests/test_latent_geometry.py @@ -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, + ) diff --git a/tests/test_tiled_sampling_runtime.py b/tests/test_tiled_sampling_runtime.py index dde1ef5..a55a984 100644 --- a/tests/test_tiled_sampling_runtime.py +++ b/tests/test_tiled_sampling_runtime.py @@ -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."""