feat(sampling): expose evaluated context SEGS

This commit is contained in:
Artificial Sweetener
2026-08-03 21:55:16 -04:00
parent 823fe209d8
commit 652ae51fc4
12 changed files with 637 additions and 17 deletions
+170
View File
@@ -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
+4
View File
@@ -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."
+35
View File
@@ -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)
+63
View File
@@ -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,
),
)
+86
View File
@@ -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},
)()
+64
View File
@@ -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,
)
+72
View File
@@ -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."""