From b58539097a4b5349fec4b7cfc0122da0bc4ad40b Mon Sep 17 00:00:00 2001 From: Artificial Sweetener Date: Sun, 2 Aug 2026 01:15:16 -0400 Subject: [PATCH] feat(sampling): add contextual diffusion sampler --- simple_syrup/domain/contextual_diffusion.py | 146 +++++++ simple_syrup/domain/segs_tiled_diffusion.py | 36 +- .../nodes/ksampler_contextual_diffusion.py | 228 +++++++++++ simple_syrup/nodes/ksampler_extras.py | 2 +- .../nodes/ksampler_tiled_diffusion.py | 6 +- simple_syrup/nodes/tooltips.py | 31 ++ simple_syrup/nodes_v3/__init__.py | 2 + simple_syrup/nodes_v3/legacy_node_wrappers.py | 10 + .../runtime/contextual_diffusion_sampling.py | 362 ++++++++++++++++++ simple_syrup/runtime/detail_sampling.py | 16 +- .../runtime/mixture_of_diffusers_sampling.py | 87 +---- .../runtime/multidiffusion_sampling.py | 64 +--- .../regional_multidiffusion_sampling.py | 16 +- simple_syrup/runtime/sampling_schedulers.py | 90 ++++- simple_syrup/runtime/tiled_sampling.py | 299 ++++++++++++++- .../contextual_diffusion_sampling_service.py | 177 +++++++++ .../services/ksampler_sampling_service.py | 15 +- simple_syrup/services/sampling_batch.py | 56 +++ .../tiled_diffusion_sampling_service.py | 47 +-- tests/test_contextual_diffusion_domain.py | 98 +++++ tests/test_contextual_diffusion_sampling.py | 289 ++++++++++++++ ...t_contextual_diffusion_sampling_service.py | 178 +++++++++ ...test_ksampler_contextual_diffusion_node.py | 152 ++++++++ tests/test_ksampler_extras_node.py | 8 + tests/test_mixture_of_diffusers_sampling.py | 3 + tests/test_multidiffusion_sampling.py | 3 + tests/test_persisted_widget_order_contract.py | 21 + tests/test_registration.py | 1 + tests/test_sampling_schedulers.py | 74 +++- tests/test_segs_tiled_diffusion.py | 29 +- tests/test_tiled_sampling_runtime.py | 83 ++++ 31 files changed, 2420 insertions(+), 209 deletions(-) create mode 100644 simple_syrup/domain/contextual_diffusion.py create mode 100644 simple_syrup/nodes/ksampler_contextual_diffusion.py create mode 100644 simple_syrup/runtime/contextual_diffusion_sampling.py create mode 100644 simple_syrup/services/contextual_diffusion_sampling_service.py create mode 100644 simple_syrup/services/sampling_batch.py create mode 100644 tests/test_contextual_diffusion_domain.py create mode 100644 tests/test_contextual_diffusion_sampling.py create mode 100644 tests/test_contextual_diffusion_sampling_service.py create mode 100644 tests/test_ksampler_contextual_diffusion_node.py diff --git a/simple_syrup/domain/contextual_diffusion.py b/simple_syrup/domain/contextual_diffusion.py new file mode 100644 index 0000000..b92aa3f --- /dev/null +++ b/simple_syrup/domain/contextual_diffusion.py @@ -0,0 +1,146 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Plan one bounded global context and one authoritative tiled context set.""" + +from __future__ import annotations + +from dataclasses import dataclass + +from .segs import NativeSegs +from .segs_tiled_diffusion import build_segs_guided_tiled_diffusion_plan +from .tiled_diffusion import TiledDiffusionPlan, build_tiled_diffusion_plan + + +@dataclass(frozen=True) +class ContextualDiffusionControls: + """Validate workflow controls for contextual diffusion sampling.""" + + latent_context_size: int + latent_context_overlap: int + latent_context_batch_size: int + global_weight: float + global_steps: int + global_decay: float + + def validate(self) -> None: + """Reject controls that cannot produce a stable bounded context plan.""" + + if self.latent_context_size < 16: + raise ValueError("latent_context_size must be at least 16 latent pixels.") + if not 0 <= self.latent_context_overlap < self.latent_context_size: + raise ValueError( + "latent_context_overlap must be non-negative and smaller than " + "latent_context_size." + ) + if self.latent_context_batch_size < 1: + raise ValueError("latent_context_batch_size must be at least 1.") + if not 0.0 <= self.global_weight <= 2.0: + raise ValueError("global_weight must be between 0 and 2.") + if self.global_steps < 0: + raise ValueError("global_steps must be non-negative.") + if not 0.0 <= self.global_decay <= 1.0: + raise ValueError("global_decay must be between 0 and 1.") + + +@dataclass(frozen=True) +class SpatialContext: + """Describe one source rectangle evaluated at a bounded model context shape.""" + + x: int + y: int + width: int + height: int + context_width: int + context_height: int + + +@dataclass(frozen=True) +class ContextualDiffusionPlan: + """Own the global context and sole tiled plan for one latent canvas.""" + + latent_width: int + latent_height: int + global_context: SpatialContext + tile_plan: TiledDiffusionPlan + + +def build_contextual_diffusion_plan( + *, + latent_width: int, + latent_height: int, + controls: ContextualDiffusionControls, + segs: NativeSegs | None, +) -> ContextualDiffusionPlan: + """Return a global context plus the regular or SEGS-guided context plan.""" + + controls.validate() + global_width, global_height = fit_context_shape( + latent_width, + latent_height, + controls.latent_context_size, + ) + global_context = SpatialContext( + x=0, + y=0, + width=latent_width, + height=latent_height, + context_width=global_width, + context_height=global_height, + ) + 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, + ) + if segs is not None + else build_tiled_diffusion_plan( + 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, + ) + ) + return ContextualDiffusionPlan( + latent_width=latent_width, + latent_height=latent_height, + global_context=global_context, + tile_plan=tile_plan, + ) + + +def fit_context_shape(width: int, height: int, max_size: int) -> tuple[int, int]: + """Fit a rectangle inside one maximum latent dimension without boxing it.""" + + if width < 1 or height < 1: + raise ValueError("Context source dimensions must be positive.") + if max_size < 1: + raise ValueError("Context maximum size must be positive.") + if max(width, height) <= max_size: + return width, height + scale = max_size / max(width, height) + fitted_width = max(2, round(width * scale)) + fitted_height = max(2, round(height * scale)) + return _even_at_most(fitted_width, max_size), _even_at_most( + fitted_height, + max_size, + ) + + +def _even_at_most(value: int, maximum: int) -> int: + """Return a positive even model-context dimension within its maximum.""" + + bounded = min(maximum, max(2, value)) + if bounded % 2 == 0: + return bounded + if bounded == maximum: + return max(2, bounded - 1) + return bounded + 1 diff --git a/simple_syrup/domain/segs_tiled_diffusion.py b/simple_syrup/domain/segs_tiled_diffusion.py index 89e5a4a..d3f528e 100644 --- a/simple_syrup/domain/segs_tiled_diffusion.py +++ b/simple_syrup/domain/segs_tiled_diffusion.py @@ -55,7 +55,7 @@ def build_segs_guided_tiled_diffusion_plan( tile_batch_size=tile_batch_size, ) native_segs = coerce_segs(segs) - _validate_aspect_ratio(native_segs, latent_height, latent_width) + validate_segs_aspect_ratio(native_segs, latent_height, latent_width) ownership = _build_ownership_cores( native_segs, latent_height=latent_height, @@ -107,7 +107,7 @@ def build_segs_guided_tiled_diffusion_plan( ) -def _validate_aspect_ratio( +def validate_segs_aspect_ratio( segs: NativeSegs, latent_height: int, latent_width: int, @@ -136,7 +136,7 @@ def _build_ownership_cores( source_height, source_width = segs[0] segment_masks = tuple( - _segment_mask_to_latent( + segment_mask_to_latent( segment, source_height=source_height, source_width=source_width, @@ -170,7 +170,7 @@ def _build_ownership_cores( ) -def _segment_mask_to_latent( +def segment_mask_to_latent( segment: Segment, *, source_height: int, @@ -180,6 +180,28 @@ def _segment_mask_to_latent( ) -> torch.Tensor: """Restore one crop-local SEG mask and map it to a latent-space mask.""" + return ( + segment_weight_to_latent( + segment, + source_height=source_height, + source_width=source_width, + latent_height=latent_height, + latent_width=latent_width, + ) + >= 0.5 + ) + + +def segment_weight_to_latent( + segment: Segment, + *, + source_height: int, + source_width: int, + latent_height: int, + latent_width: int, +) -> torch.Tensor: + """Project one crop-local SEG mask into latent space without binarizing it.""" + crop = segment.crop_region if ( crop.left < 0 @@ -215,7 +237,7 @@ def _segment_mask_to_latent( source_width, latent_width, ) - latent_mask = torch.zeros((latent_height, latent_width), dtype=torch.bool) + latent_mask = torch.zeros((latent_height, latent_width), dtype=torch.float32) if latent_bottom <= latent_top or latent_right <= latent_left: return latent_mask sampled_rows = ( @@ -242,9 +264,7 @@ def _segment_mask_to_latent( ) .index_select(1, sampled_columns) ) - latent_mask[latent_top:latent_bottom, latent_left:latent_right] = ( - sampled_mask >= 0.5 - ) + latent_mask[latent_top:latent_bottom, latent_left:latent_right] = sampled_mask return latent_mask diff --git a/simple_syrup/nodes/ksampler_contextual_diffusion.py b/simple_syrup/nodes/ksampler_contextual_diffusion.py new file mode 100644 index 0000000..d187af2 --- /dev/null +++ b/simple_syrup/nodes/ksampler_contextual_diffusion.py @@ -0,0 +1,228 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""ComfyUI node declaration for contextual diffusion sampling.""" + +from __future__ import annotations + +from typing import Any, ClassVar, TypeAlias + +from ..domain.tiled_diffusion import TILED_DIFFUSION_MODES +from ..runtime import sampling_samplers, sampling_schedulers +from ..services.contextual_diffusion_sampling_service import ( + ContextualDiffusionSamplingService, +) +from . import tooltips + +Latent: TypeAlias = dict[str, Any] +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,) + FUNCTION = "sample" + CATEGORY = "SimpleSyrup/Sampling" + DESCRIPTION = ( + "Preserves composition while applying appearance and subject-detail edits " + "to large latents through global context and optional SEGS-guided tiles." + ) + SEARCH_ALIASES = [ + "ksampler", + "contextual diffusion", + "contextual tiled diffusion", + "high resolution edit", + "sam tiled diffusion", + ] + + service_class: ClassVar[type[ContextualDiffusionSamplingService]] = ( + ContextualDiffusionSamplingService + ) + + @classmethod + def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]: + """Declare KSampler inputs and bounded contextual controls.""" + + return { + "required": { + "model": ("MODEL", {"tooltip": tooltips.SAMPLING_MODEL}), + "seed": ( + "INT", + { + "default": 0, + "min": 0, + "max": 0xFFFFFFFFFFFFFFFF, + "control_after_generate": True, + "tooltip": tooltips.SAMPLING_SEED, + }, + ), + "steps": ( + "INT", + { + "default": 4, + "min": 1, + "max": 10000, + "tooltip": tooltips.SAMPLING_STEPS, + }, + ), + "cfg": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 100.0, + "step": 0.1, + "round": 0.01, + "tooltip": tooltips.SAMPLING_CFG, + }, + ), + "sampler_name": ( + sampling_samplers.available_samplers(), + {"tooltip": tooltips.SAMPLER_NAME}, + ), + "scheduler": ( + sampling_schedulers.available_schedulers(), + {"tooltip": tooltips.SCHEDULER}, + ), + "positive": ( + "CONDITIONING,CONDITIONING_BATCH", + {"tooltip": tooltips.POSITIVE_CONDITIONING}, + ), + "negative": ( + "CONDITIONING,CONDITIONING_BATCH", + {"tooltip": tooltips.NEGATIVE_CONDITIONING}, + ), + "latent_image": ("LATENT", {"tooltip": tooltips.LATENT_IMAGE}), + "denoise": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": tooltips.DENOISE_STRENGTH, + }, + ), + "diffusion_mode": ( + list(TILED_DIFFUSION_MODES), + { + "default": "multidiffusion", + "tooltip": tooltips.TILED_DIFFUSION_MODE, + }, + ), + "latent_context_size": ( + "INT", + { + "default": 96, + "min": 16, + "max": MAX_LATENT_CONTEXT_SIZE, + "step": 16, + "tooltip": tooltips.LATENT_CONTEXT_SIZE, + }, + ), + "latent_context_overlap": ( + "INT", + { + "default": 32, + "min": 0, + "max": 256, + "step": 4, + "tooltip": tooltips.LATENT_CONTEXT_OVERLAP, + }, + ), + "latent_context_batch_size": ( + "INT", + { + "default": 4, + "min": 1, + "max": 8, + "step": 1, + "tooltip": tooltips.LATENT_CONTEXT_BATCH_SIZE, + }, + ), + "global_weight": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 2.0, + "step": 0.05, + "tooltip": tooltips.GLOBAL_CONTEXT_WEIGHT, + }, + ), + "global_steps": ( + "INT", + { + "default": 1, + "min": 0, + "max": 10000, + "step": 1, + "tooltip": tooltips.GLOBAL_CONTEXT_STEPS, + }, + ), + "global_decay": ( + "FLOAT", + { + "default": 0.5, + "min": 0.0, + "max": 1.0, + "step": 0.05, + "tooltip": tooltips.GLOBAL_CONTEXT_DECAY, + }, + ), + }, + "optional": { + "segs": ( + "SEGS", + {"tooltip": tooltips.CONTEXTUAL_DIFFUSION_SEGS}, + ) + }, + } + + def sample( + self, + model: Any, + seed: int, + steps: int, + cfg: float, + sampler_name: str, + scheduler: str, + positive: Any, + negative: Any, + latent_image: Latent, + denoise: float = 1.0, + diffusion_mode: str = "multidiffusion", + latent_context_size: int = 96, + latent_context_overlap: int = 32, + latent_context_batch_size: int = 4, + global_weight: float = 1.0, + global_steps: int = 1, + global_decay: float = 0.5, + segs: object | None = None, + ) -> tuple[Latent]: + """Delegate contextual diffusion sampling to its application service.""" + + output = self.service_class().sample( + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=positive, + negative=negative, + latent_image=latent_image, + denoise=denoise, + diffusion_mode=diffusion_mode, + latent_context_size=latent_context_size, + latent_context_overlap=latent_context_overlap, + latent_context_batch_size=latent_context_batch_size, + global_weight=global_weight, + global_steps=global_steps, + global_decay=global_decay, + segs=segs, + ) + return (output,) diff --git a/simple_syrup/nodes/ksampler_extras.py b/simple_syrup/nodes/ksampler_extras.py index b78d8c5..0adee61 100644 --- a/simple_syrup/nodes/ksampler_extras.py +++ b/simple_syrup/nodes/ksampler_extras.py @@ -16,7 +16,7 @@ Latent = dict[str, Any] class KSamplerExtras: - """Expose KSampler-style sampling with AYS and GITS scheduler options.""" + """Expose KSampler-style sampling with extended scheduler options.""" RETURN_TYPES = ("LATENT",) OUTPUT_TOOLTIPS = (tooltips.DENOISED_LATENT_OUTPUT,) diff --git a/simple_syrup/nodes/ksampler_tiled_diffusion.py b/simple_syrup/nodes/ksampler_tiled_diffusion.py index 4454e3a..b8ecab7 100644 --- a/simple_syrup/nodes/ksampler_tiled_diffusion.py +++ b/simple_syrup/nodes/ksampler_tiled_diffusion.py @@ -106,11 +106,7 @@ class KSamplerTiledDiffusion: list(TILED_DIFFUSION_MODES), { "default": "multidiffusion", - "tooltip": ( - "Tiled sampling blend method. MultiDiffusion is steady; " - "Mixture of Diffusers can blend tile predictions more " - "softly." - ), + "tooltip": tooltips.TILED_DIFFUSION_MODE, }, ), "latent_tile_width": ( diff --git a/simple_syrup/nodes/tooltips.py b/simple_syrup/nodes/tooltips.py index 3d970f3..b4ac98b 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." +TILED_DIFFUSION_MODE = ( + "Tile overlap blend. MultiDiffusion averages predictions; Mixture of Diffusers " + "gives tile centers more influence." +) LATENT_TILE_WIDTH = ( "Width of each latent tile. Larger tiles see more context but use more memory." @@ -152,6 +156,33 @@ LATENT_TILE_BATCH_SIZE = ( "Number of latent tiles sampled together. Higher values can be faster but use " "more memory." ) +LATENT_CONTEXT_SIZE = ( + "Maximum side of each model context in latent pixels. Larger contexts preserve " + "more relationships but use more memory." +) +LATENT_CONTEXT_OVERLAP = ( + "Overlap between local latent contexts in latent pixels. Larger overlaps reduce " + "seams but increase sampling work." +) +LATENT_CONTEXT_BATCH_SIZE = ( + "Number of equal-sized latent contexts sampled together. Higher values can be " + "faster but use more memory." +) +GLOBAL_CONTEXT_WEIGHT = ( + "Strength of whole-image low-frequency guidance. 1 makes the global context " + "authoritative; lower values allow more tile interpretation." +) +GLOBAL_CONTEXT_STEPS = ( + "Number of initial denoising steps that use the global context. Fewer steps leave " + "more late sampling for local detail." +) +GLOBAL_CONTEXT_DECAY = ( + "Multiplier applied to whole-image strength after each global step. Lower " + "values hand control to local contexts faster." +) +CONTEXTUAL_DIFFUSION_SEGS = ( + "Optional regions that replace the regular grid with SEGS-guided contexts." +) DETAIL_IMAGE = ( "Source image containing the regions to improve. Detailed crops are blended " diff --git a/simple_syrup/nodes_v3/__init__.py b/simple_syrup/nodes_v3/__init__.py index 3bbfa89..b282d50 100644 --- a/simple_syrup/nodes_v3/__init__.py +++ b/simple_syrup/nodes_v3/__init__.py @@ -28,6 +28,7 @@ def get_nodes() -> list[type[object]]: EncodePromptBatchV3, GroundedSAMModelInfoV3, GroundingDINOModelLoaderV3, + KSamplerContextualDiffusionV3, KSamplerExtrasV3, KSamplerTiledDiffusionV3, LatentDiagnosticsV3, @@ -73,6 +74,7 @@ def get_nodes() -> list[type[object]]: KSamplerExtrasV3, KSamplerPromptByRegionV3, KSamplerPromptByTiledRegionV3, + KSamplerContextualDiffusionV3, KSamplerTiledDiffusionV3, LatentDiagnosticsV3, LayerStyleSAMModelsAdapterV3, diff --git a/simple_syrup/nodes_v3/legacy_node_wrappers.py b/simple_syrup/nodes_v3/legacy_node_wrappers.py index 1c29167..4e075ad 100644 --- a/simple_syrup/nodes_v3/legacy_node_wrappers.py +++ b/simple_syrup/nodes_v3/legacy_node_wrappers.py @@ -24,6 +24,7 @@ from ..nodes.encode_prompt_batch import EncodePromptBatch from ..nodes.grounded_sam_model_info import GroundedSAMModelInfo from ..nodes.grounding_dino_model_loader import GroundingDINOModelLoader from ..nodes.image_resize_to_target import ResizeImageToTarget +from ..nodes.ksampler_contextual_diffusion import KSamplerContextualDiffusion from ..nodes.ksampler_extras import KSamplerExtras from ..nodes.ksampler_tiled_diffusion import KSamplerTiledDiffusion from ..nodes.latent_diagnostics import LatentDiagnostics @@ -157,6 +158,14 @@ class KSamplerTiledDiffusionV3(LegacyNodeV3Adapter): DISPLAY_NAME = "KSampler (Tiled Diffusion)" +class KSamplerContextualDiffusionV3(LegacyNodeV3Adapter): + """Expose KSampler Contextual Diffusion through Comfy v3 only.""" + + LEGACY_NODE_CLASS = KSamplerContextualDiffusion + NODE_ID = "SimpleSyrup.KSamplerContextualDiffusion" + DISPLAY_NAME = "KSampler (Contextual Diffusion)" + + class LayerStyleSAMModelsAdapterV3(LegacyNodeV3Adapter): """Expose LayerStyle SAM Models Adapter through Comfy v3 only.""" @@ -505,6 +514,7 @@ __all__ = [ "GroundedSAMModelInfoV3", "GroundingDINOModelLoaderV3", "KSamplerExtrasV3", + "KSamplerContextualDiffusionV3", "KSamplerTiledDiffusionV3", "LatentDiagnosticsV3", "LayerStyleSAMModelsAdapterV3", diff --git a/simple_syrup/runtime/contextual_diffusion_sampling.py b/simple_syrup/runtime/contextual_diffusion_sampling.py new file mode 100644 index 0000000..53401ce --- /dev/null +++ b/simple_syrup/runtime/contextual_diffusion_sampling.py @@ -0,0 +1,362 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""ComfyUI runtime adapter for contextual diffusion latent sampling.""" + +from __future__ import annotations + +from importlib import import_module +from types import ModuleType +from typing import Any, cast + +import torch + +from ..domain.contextual_diffusion import ( + ContextualDiffusionControls, + ContextualDiffusionPlan, +) +from ..shared.logging import get_logger +from . import sampling_samplers, sampling_schedulers +from .tiled_sampling import ( + ApplyModel, + Latent, + ModelFunctionWrapper, + TilePredictionAccumulator, + make_spatial_context_model_args, + reject_unsupported_conditioning, + resize_spatial_tensor, + validate_latent_samples, + validate_sampling_controls, + validate_tensor_shape, +) + +LOGGER = get_logger(__name__) +SAMPLER_LABEL = "Contextual Diffusion" +UNIPC_SAMPLERS = frozenset({"uni_pc", "uni_pc_bh2"}) + + +def sample_contextual_diffusion( + *, + model: Any, + seed: int, + steps: int, + cfg: float, + sampler_name: str, + scheduler: str, + positive: Any, + negative: Any, + latent_image: Latent, + denoise: float, + diffusion_mode: str, + controls: ContextualDiffusionControls, + plan: ContextualDiffusionPlan, +) -> Latent: + """Sample one latent through global context and one tiled prediction plan.""" + + validate_sampling_controls( + steps=steps, + denoise=denoise, + latent_tile_width=controls.latent_context_size, + latent_tile_height=controls.latent_context_size, + latent_tile_batch_size=controls.latent_context_batch_size, + ) + 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) + + sampler = sampling_samplers.resolve_sampler(sampler_name) + sigmas = sampling_schedulers.calculate_sigmas( + model=model, + scheduler_name=scheduler, + sampler_name=sampler_name, + steps=steps, + denoise=denoise, + view=sampling_schedulers.SchedulerView( + latent_width=controls.latent_context_size, + latent_height=controls.latent_context_size, + ), + ).to(model.load_device) + latent_samples = validate_latent_samples(latent_image, sampler_label=SAMPLER_LABEL) + comfy_sample = _comfy_sample() + comfy_utils = _comfy_utils() + latent_samples = comfy_sample.fix_empty_latent_channels( + model, + latent_samples, + latent_image.get("downscale_ratio_spacial", None), + ) + validate_tensor_shape(latent_samples, sampler_label=SAMPLER_LABEL) + if latent_samples.shape[-2:] != (plan.latent_height, plan.latent_width): + raise ValueError( + "Contextual Diffusion plan dimensions must match the sampled latent shape." + ) + + sampling_model = clone_model_with_contextual_diffusion( + model, + plan=plan, + controls=controls, + sigmas=sigmas, + diffusion_mode=diffusion_mode, + ) + batch_inds = latent_image.get("batch_index") + noise = comfy_sample.prepare_noise(latent_samples, seed, batch_inds) + callback = _latent_preview().prepare_callback(sampling_model, steps) + samples = comfy_sample.sample_custom( + sampling_model, + noise, + cfg, + sampler, + sigmas, + positive, + negative, + latent_samples, + noise_mask=latent_image.get("noise_mask"), + callback=callback, + disable_pbar=not comfy_utils.PROGRESS_BAR_ENABLED, + seed=seed, + ) + + LOGGER.info( + "KSampler Contextual Diffusion pass completed", + extra={ + "operation": "ksampler_contextual_diffusion", + "sampler": sampler_name, + "scheduler": scheduler, + "steps": steps, + "denoise": denoise, + "diffusion_mode": diffusion_mode, + "latent_width": plan.latent_width, + "latent_height": plan.latent_height, + "context_size": controls.latent_context_size, + "overlap": controls.latent_context_overlap, + "tile_count": len(plan.tile_plan.tiles), + "segs_guided": any( + tile.weight_mask is not None for tile in plan.tile_plan.tiles + ), + "model_calls_per_prediction": ( + 1 + if len(plan.tile_plan.tiles) == 1 + else len(plan.tile_plan.batches) + int(controls.global_weight > 0) + ), + }, + ) + output = latent_image.copy() + output.pop("downscale_ratio_spacial", None) + output["samples"] = samples + return output + + +def clone_model_with_contextual_diffusion( + model: Any, + *, + plan: ContextualDiffusionPlan, + controls: ContextualDiffusionControls, + sigmas: torch.Tensor, + diffusion_mode: str, +) -> Any: + """Clone a model and install one pre-CFG contextual prediction wrapper.""" + + cloned_model = model.clone() + old_wrapper = cloned_model.model_options.get("model_function_wrapper") + if old_wrapper is not None and not callable(old_wrapper): + raise ValueError("Existing model_function_wrapper is not callable.") + wrapper = ContextualDiffusionModelWrapper( + plan=plan, + controls=controls, + sigmas=sigmas, + diffusion_mode=diffusion_mode, + existing_wrapper=cast(ModelFunctionWrapper | None, old_wrapper), + ) + cloned_model.set_model_unet_function_wrapper(wrapper) + return cloned_model + + +class ContextualDiffusionModelWrapper: + """Fuse one authoritative tiled prediction with scheduled global structure.""" + + def __init__( + self, + *, + plan: ContextualDiffusionPlan, + controls: ContextualDiffusionControls, + sigmas: torch.Tensor, + diffusion_mode: str = "multidiffusion", + existing_wrapper: ModelFunctionWrapper | None, + ) -> None: + """Create a wrapper for one immutable latent and context plan.""" + + self._plan = plan + self._controls = controls + self._global_schedule = GlobalContextSchedule( + sigmas=sigmas, + active_steps=controls.global_steps, + decay=controls.global_decay, + ) + self._existing_wrapper = existing_wrapper + self._diffusion_mode = diffusion_mode + self._tile_predictions = TilePredictionAccumulator( + plan.tile_plan, + diffusion_mode=diffusion_mode, + ) + + @property + def diffusion_mode(self) -> str: + """Return the selected local tile blending policy.""" + + return self._diffusion_mode + + def __call__( + self, + apply_model: ApplyModel, + args: dict[str, Any], + ) -> torch.Tensor: + """Evaluate bounded contexts and return one canvas-sized prediction.""" + + x = args.get("input") + if not isinstance(x, torch.Tensor): + raise ValueError("Contextual Diffusion model input must be a tensor.") + validate_tensor_shape(x, sampler_label=SAMPLER_LABEL) + if x.shape[-2:] != (self._plan.latent_height, self._plan.latent_width): + return self._call_original(apply_model, args) + if len(self._plan.tile_plan.tiles) == 1: + return self._call_original(apply_model, args) + + tile_prediction = self._predict_tiles(apply_model, args, x) + global_scale = self._global_schedule.scale_for(args.get("timestep")) + return ( + self._apply_global_authority( + apply_model, + args, + tile_prediction, + weight=self._controls.global_weight * global_scale, + ) + if global_scale > 0 and self._controls.global_weight > 0 + else tile_prediction + ) + + def _predict_tiles( + self, + apply_model: ApplyModel, + args: dict[str, Any], + x: torch.Tensor, + ) -> torch.Tensor: + """Blend the authoritative regular or SEGS-guided tiled prediction.""" + + return self._tile_predictions.predict( + args=args, + x=x, + evaluate=lambda tiled_args: self._call_original(apply_model, tiled_args), + ) + + def _apply_global_authority( + self, + apply_model: ApplyModel, + args: dict[str, Any], + local_prediction: torch.Tensor, + *, + weight: float, + ) -> torch.Tensor: + """Replace tile-scale scene interpretation with the whole-image prediction.""" + + global_args = make_spatial_context_model_args( + args=args, + contexts=(self._plan.global_context,), + input_batch_size=int(local_prediction.shape[0]), + latent_height=self._plan.latent_height, + latent_width=self._plan.latent_width, + ) + global_prediction = self._call_original(apply_model, global_args) + context = self._plan.global_context + global_canvas = resize_spatial_tensor( + global_prediction, + height=self._plan.latent_height, + width=self._plan.latent_width, + mode="bilinear", + ) + local_low = resize_spatial_tensor( + local_prediction, + height=context.context_height, + width=context.context_width, + mode="area", + ) + local_low = resize_spatial_tensor( + local_low, + height=self._plan.latent_height, + width=self._plan.latent_width, + mode="bilinear", + ) + return local_prediction + weight * (global_canvas - local_low) + + def _call_original( + self, + apply_model: ApplyModel, + args: dict[str, Any], + ) -> torch.Tensor: + """Call the preserved model wrapper or raw apply_model.""" + + if self._existing_wrapper is not None: + return self._existing_wrapper(apply_model, args) + conditioning = args.get("c", {}) + if not isinstance(conditioning, dict): + raise ValueError("Contextual Diffusion conditioning must be a dict.") + return apply_model(args["input"], args["timestep"], **conditioning) + + +class GlobalContextSchedule: + """Limit whole-image authority to an initial denoising-step fraction.""" + + def __init__( + self, + *, + sigmas: torch.Tensor, + active_steps: int, + decay: float, + ) -> None: + """Capture model-evaluation sigmas and the active initial step count.""" + + step_sigmas = sigmas.detach().to(device="cpu", dtype=torch.float64).flatten() + if step_sigmas.numel() < 2: + raise ValueError( + "Contextual Diffusion requires at least one denoising step." + ) + self._step_sigmas = step_sigmas[:-1] + self._active_steps = min( + len(self._step_sigmas), + max(0, active_steps), + ) + self._decay = decay + + def scale_for(self, timestep: object) -> float: + """Return the decayed global scale for the nearest scheduled step.""" + + if self._active_steps == 0: + return 0.0 + if not isinstance(timestep, torch.Tensor) or timestep.numel() == 0: + raise ValueError( + "Contextual Diffusion timestep must be a non-empty tensor." + ) + sigma = timestep.detach().flatten()[0].to(device="cpu", dtype=torch.float64) + step_index = int(torch.argmin(torch.abs(self._step_sigmas - sigma)).item()) + if step_index >= self._active_steps: + return 0.0 + return self._decay**step_index + + +def _comfy_sample() -> ModuleType: + """Import ComfyUI sample helpers lazily.""" + + return import_module("comfy.sample") + + +def _comfy_utils() -> ModuleType: + """Import ComfyUI progress state lazily.""" + + return import_module("comfy.utils") + + +def _latent_preview() -> ModuleType: + """Import ComfyUI preview helpers lazily.""" + + return import_module("latent_preview") diff --git a/simple_syrup/runtime/detail_sampling.py b/simple_syrup/runtime/detail_sampling.py index 0044cc7..f19e326 100644 --- a/simple_syrup/runtime/detail_sampling.py +++ b/simple_syrup/runtime/detail_sampling.py @@ -60,14 +60,6 @@ class DetailSampler: """Sample a latent with SimpleSyrup's sampler and scheduler helpers.""" sampler = sampling_samplers.resolve_sampler(sampler_name) - sigmas = sampling_schedulers.calculate_sigmas( - model=model, - scheduler_name=scheduler, - sampler_name=sampler_name, - steps=steps, - denoise=denoise, - ).to(model.load_device) - latent_samples = cast(torch.Tensor, latent_image["samples"]) comfy_sample = _comfy_sample() comfy_utils = _comfy_utils() @@ -77,6 +69,14 @@ class DetailSampler: latent_samples, latent_image.get("downscale_ratio_spacial", None), ) + sigmas = sampling_schedulers.calculate_sigmas( + model=model, + scheduler_name=scheduler, + sampler_name=sampler_name, + steps=steps, + denoise=denoise, + view=sampling_schedulers.SchedulerView.from_tensor(latent_samples), + ).to(model.load_device) batch_inds = ( latent_image["batch_index"] if "batch_index" in latent_image else None ) diff --git a/simple_syrup/runtime/mixture_of_diffusers_sampling.py b/simple_syrup/runtime/mixture_of_diffusers_sampling.py index 249a9a4..917b7ac 100644 --- a/simple_syrup/runtime/mixture_of_diffusers_sampling.py +++ b/simple_syrup/runtime/mixture_of_diffusers_sampling.py @@ -10,7 +10,6 @@ from __future__ import annotations -from collections.abc import Sequence from importlib import import_module from types import ModuleType from typing import Any, cast @@ -18,10 +17,8 @@ from typing import Any, cast import torch from ..domain.tiled_diffusion import ( - LatentTile, TiledDiffusionPlan, build_tiled_diffusion_plan, - gaussian_tile_weights, ) from ..shared.logging import get_logger from . import sampling_samplers, sampling_schedulers @@ -31,11 +28,8 @@ from .tiled_sampling import ( ApplyModel, Latent, ModelFunctionWrapper, - SemanticTileWeightCache, - make_tiled_model_args, - new_spatial_weight_buffer, + TilePredictionAccumulator, reject_unsupported_conditioning, - spatial_tile_slicer, validate_latent_samples, validate_sampling_controls, validate_tensor_shape, @@ -93,6 +87,10 @@ def sample_mixture_of_diffusers( sampler_name=sampler_name, steps=steps, denoise=denoise, + view=sampling_schedulers.SchedulerView( + latent_width=latent_tile_width, + latent_height=latent_tile_height, + ), ).to(model.load_device) latent_samples = validate_latent_samples( @@ -217,8 +215,10 @@ class MixtureOfDiffusersModelWrapper: self._plan = plan self._existing_wrapper = existing_wrapper - self._tile_weights_2d: torch.Tensor | None = None - self._semantic_tile_weights = SemanticTileWeightCache(plan.tiles) + self._tile_predictions = TilePredictionAccumulator( + plan, + diffusion_mode="mixture_of_diffusers", + ) def __call__( self, @@ -248,35 +248,11 @@ class MixtureOfDiffusersModelWrapper: "ControlNet in the first implementation." ) - output_buffer = torch.zeros_like(x) - weight_buffer = new_spatial_weight_buffer(x, self._plan) - input_batch_size = int(x.shape[0]) - - for batch in self._plan.batches: - tiled_args = self._make_tiled_args( - args=args, - tiles=batch, - input_batch_size=input_batch_size, - ) - tile_output = self._call_original(apply_model, tiled_args) - weights = self._weights_for(tile_output) - accumulation_weights = weights.to(dtype=weight_buffer.dtype) - semantic_weights = self._semantic_tile_weights.for_output(tile_output) - for index, tile in enumerate(batch): - tile_slice = spatial_tile_slicer(tile, x.ndim) - start = index * input_batch_size - end = start + input_batch_size - model_weight, accumulation_weight = ( - self._semantic_tile_weights.for_tile( - semantic_weights, - tile, - ) - ) - tile_weight = weights * model_weight - output_buffer[tile_slice] += tile_output[start:end] * tile_weight - weight_buffer[tile_slice] += accumulation_weights * accumulation_weight - - return output_buffer / weight_buffer.to(dtype=output_buffer.dtype) + return self._tile_predictions.predict( + args=args, + x=x, + evaluate=lambda tiled_args: self._call_original(apply_model, tiled_args), + ) def _call_original( self, @@ -292,41 +268,6 @@ class MixtureOfDiffusersModelWrapper: raise ValueError("Mixture of Diffusers conditioning must be a dict.") return apply_model(args["input"], args["timestep"], **conditioning) - def _make_tiled_args( - self, - *, - args: dict[str, Any], - tiles: Sequence[LatentTile], - input_batch_size: int, - ) -> dict[str, Any]: - """Create apply-model args for one tile batch.""" - - return make_tiled_model_args( - args=args, - tiles=tiles, - input_batch_size=input_batch_size, - latent_height=self._plan.latent_height, - latent_width=self._plan.latent_width, - ) - - def _weights_for(self, x: torch.Tensor) -> torch.Tensor: - """Return cached Gaussian tile weights for the active device and dtype.""" - - if ( - self._tile_weights_2d is None - or self._tile_weights_2d.device != x.device - or self._tile_weights_2d.dtype != x.dtype - ): - self._tile_weights_2d = gaussian_tile_weights( - self._plan.tile_width, - self._plan.tile_height, - device=x.device, - dtype=x.dtype, - ) - return self._tile_weights_2d.reshape( - (1,) * (x.ndim - 2) + (self._plan.tile_height, self._plan.tile_width) - ) - def _validate_supplied_plan( plan: TiledDiffusionPlan, diff --git a/simple_syrup/runtime/multidiffusion_sampling.py b/simple_syrup/runtime/multidiffusion_sampling.py index 5db98e7..5181b19 100644 --- a/simple_syrup/runtime/multidiffusion_sampling.py +++ b/simple_syrup/runtime/multidiffusion_sampling.py @@ -10,7 +10,6 @@ from __future__ import annotations -from collections.abc import Sequence from importlib import import_module from types import ModuleType from typing import Any, cast @@ -18,7 +17,6 @@ from typing import Any, cast import torch from ..domain.tiled_diffusion import ( - LatentTile, TiledDiffusionPlan, build_tiled_diffusion_plan, ) @@ -30,11 +28,8 @@ from .tiled_sampling import ( ApplyModel, Latent, ModelFunctionWrapper, - SemanticTileWeightCache, - make_tiled_model_args, - new_spatial_weight_buffer, + TilePredictionAccumulator, reject_unsupported_conditioning, - spatial_tile_slicer, validate_latent_samples, validate_sampling_controls, validate_tensor_shape, @@ -94,6 +89,10 @@ def sample_multidiffusion( sampler_name=sampler_name, steps=steps, denoise=denoise, + view=sampling_schedulers.SchedulerView( + latent_width=latent_tile_width, + latent_height=latent_tile_height, + ), ).to(model.load_device) latent_samples = validate_latent_samples( @@ -219,7 +218,10 @@ class MultiDiffusionModelWrapper: self._plan = plan self._existing_wrapper = existing_wrapper - self._semantic_tile_weights = SemanticTileWeightCache(plan.tiles) + self._tile_predictions = TilePredictionAccumulator( + plan, + diffusion_mode="multidiffusion", + ) def __call__( self, @@ -249,32 +251,11 @@ class MultiDiffusionModelWrapper: "ControlNet in the first implementation." ) - output_buffer = torch.zeros_like(x) - weight_buffer = new_spatial_weight_buffer(x, self._plan) - input_batch_size = int(x.shape[0]) - - for batch in self._plan.batches: - tiled_args = self._make_tiled_args( - args=args, - tiles=batch, - input_batch_size=input_batch_size, - ) - tile_output = self._call_original(apply_model, tiled_args) - tile_weights = self._semantic_tile_weights.for_output(tile_output) - for index, tile in enumerate(batch): - tile_slice = spatial_tile_slicer(tile, x.ndim) - start = index * input_batch_size - end = start + input_batch_size - model_weight, accumulation_weight = ( - self._semantic_tile_weights.for_tile( - tile_weights, - tile, - ) - ) - output_buffer[tile_slice] += tile_output[start:end] * model_weight - weight_buffer[tile_slice] += accumulation_weight - - return output_buffer / weight_buffer.to(dtype=output_buffer.dtype) + return self._tile_predictions.predict( + args=args, + x=x, + evaluate=lambda tiled_args: self._call_original(apply_model, tiled_args), + ) def _call_original( self, @@ -290,23 +271,6 @@ class MultiDiffusionModelWrapper: raise ValueError("MultiDiffusion conditioning must be a dict.") return apply_model(args["input"], args["timestep"], **conditioning) - def _make_tiled_args( - self, - *, - args: dict[str, Any], - tiles: Sequence[LatentTile], - input_batch_size: int, - ) -> dict[str, Any]: - """Create apply-model args for one tile batch.""" - - return make_tiled_model_args( - args=args, - tiles=tiles, - input_batch_size=input_batch_size, - latent_height=self._plan.latent_height, - latent_width=self._plan.latent_width, - ) - def _reject_unipc_sampler(sampler_name: str) -> None: """Reject UniPC samplers because MultiDiffusion is incompatible with them.""" diff --git a/simple_syrup/runtime/regional_multidiffusion_sampling.py b/simple_syrup/runtime/regional_multidiffusion_sampling.py index fdd0b8c..e06fa91 100644 --- a/simple_syrup/runtime/regional_multidiffusion_sampling.py +++ b/simple_syrup/runtime/regional_multidiffusion_sampling.py @@ -81,14 +81,6 @@ def sample_regional_multidiffusion( reject_unsupported_conditioning(region.positive, sampler_label=SAMPLER_LABEL) sampler = sampling_samplers.resolve_sampler(sampler_name) - sigmas = sampling_schedulers.calculate_sigmas( - model=model, - scheduler_name=scheduler, - sampler_name=sampler_name, - steps=steps, - denoise=denoise, - ).to(model.load_device) - latent_samples = validate_latent_samples( latent_image, sampler_label=SAMPLER_LABEL, @@ -102,6 +94,14 @@ def sample_regional_multidiffusion( latent_image.get("downscale_ratio_spacial", None), ) validate_tensor_shape(latent_samples, sampler_label=SAMPLER_LABEL) + sigmas = sampling_schedulers.calculate_sigmas( + model=model, + scheduler_name=scheduler, + sampler_name=sampler_name, + steps=steps, + denoise=denoise, + view=sampling_schedulers.SchedulerView.from_tensor(latent_samples), + ).to(model.load_device) latent_height = int(latent_samples.shape[-2]) latent_width = int(latent_samples.shape[-1]) sampling_model, summary = clone_model_with_regional_multidiffusion( diff --git a/simple_syrup/runtime/sampling_schedulers.py b/simple_syrup/runtime/sampling_schedulers.py index 2fe7539..56218ab 100644 --- a/simple_syrup/runtime/sampling_schedulers.py +++ b/simple_syrup/runtime/sampling_schedulers.py @@ -10,7 +10,8 @@ from __future__ import annotations import math -from collections.abc import Sequence +from collections.abc import Callable, Sequence +from dataclasses import dataclass from importlib import import_module from types import ModuleType from typing import Protocol, cast @@ -21,7 +22,14 @@ from ..shared.logging import get_logger LOGGER = get_logger(__name__) -EXTRA_SCHEDULERS = ("AYS SD1", "AYS SDXL", "GITS", "beta57", "automatic_a1111") +EXTRA_SCHEDULERS = ( + "AYS SD1", + "AYS SDXL", + "GITS", + "beta57", + "automatic_a1111", + "Flux2", +) GITS_DEFAULT_COEFF = 1.20 BETA57_ALPHA = 0.5 BETA57_BETA = 0.7 @@ -310,6 +318,33 @@ class SamplingModel(Protocol): """Return a named ComfyUI model object.""" +@dataclass(frozen=True) +class SchedulerView: + """Describe the spatial latent view evaluated by one model prediction.""" + + latent_width: int + latent_height: int + + def __post_init__(self) -> None: + """Reject dimensions that cannot define a spatial schedule.""" + + if self.latent_width <= 0 or self.latent_height <= 0: + raise ValueError("Scheduler model-view dimensions must be positive.") + + @classmethod + def from_tensor(cls, samples: torch.Tensor) -> SchedulerView: + """Create a scheduler view from the tensor's final spatial dimensions.""" + + if samples.ndim < 2: + raise ValueError( + "Scheduler model-view samples must have spatial dimensions." + ) + return cls( + latent_width=int(samples.shape[-1]), + latent_height=int(samples.shape[-2]), + ) + + def available_schedulers() -> tuple[str, ...]: """Return core ComfyUI schedulers plus locally resolved extra schedulers.""" @@ -324,6 +359,8 @@ def calculate_sigmas( sampler_name: str, steps: int, denoise: float, + *, + view: SchedulerView | None = None, ) -> torch.Tensor: """Calculate sigmas for a core or SimpleSyrup-owned scheduler.""" @@ -351,6 +388,7 @@ def calculate_sigmas( sampler_name=sampler_name, steps=steps, denoise=denoise, + view=view, ) return _calculate_core_sigmas( @@ -414,6 +452,7 @@ def _calculate_extra_sigmas( sampler_name: str, steps: int, denoise: float, + view: SchedulerView | None, ) -> torch.Tensor: """Calculate sigmas for locally resolved extra scheduler policies.""" @@ -425,7 +464,12 @@ def _calculate_extra_sigmas( calculation_steps = ( schedule_steps + 1 if discard_penultimate_sigma else schedule_steps ) - sigmas = _calculate_extra_schedule(model, scheduler_name, calculation_steps) + sigmas = _calculate_extra_schedule( + model, + scheduler_name, + calculation_steps, + view=view, + ) if discard_penultimate_sigma: sigmas = torch.cat([sigmas[:-2], sigmas[-1:]]) @@ -456,6 +500,8 @@ def _calculate_extra_schedule( model: SamplingModel, scheduler_name: str, steps: int, + *, + view: SchedulerView | None, ) -> torch.Tensor: """Calculate a full local extra scheduler output.""" @@ -469,9 +515,47 @@ def _calculate_extra_schedule( return _calculate_beta57_schedule(model, steps) if scheduler_name == "automatic_a1111": return _calculate_automatic_a1111_schedule(model, steps) + if scheduler_name == "Flux2": + return _calculate_flux2_schedule(model, steps, view=view) raise ValueError(f"Unsupported extra scheduler '{scheduler_name}'.") +def _calculate_flux2_schedule( + model: SamplingModel, + steps: int, + *, + view: SchedulerView | None, +) -> torch.Tensor: + """Calculate ComfyUI's Flux2 schedule for the effective model view.""" + + if view is None: + raise ValueError("Flux2 scheduler requires a model view resolution.") + latent_format = model.get_model_object("latent_format") + spatial_ratio = getattr(latent_format, "spacial_downscale_ratio", None) + if ( + isinstance(spatial_ratio, bool) + or not isinstance(spatial_ratio, (int, float)) + or spatial_ratio <= 0 + ): + raise ValueError( + "Flux2 scheduler requires the model latent format to expose a " + "positive spacial_downscale_ratio." + ) + + pixel_width = view.latent_width * float(spatial_ratio) + pixel_height = view.latent_height * float(spatial_ratio) + image_sequence_length = round(pixel_width * pixel_height / (16 * 16)) + flux_nodes = import_module("comfy_extras.nodes_flux") + get_schedule = cast( + Callable[[int, int], Sequence[float] | torch.Tensor], + flux_nodes.get_schedule, + ) + return torch.as_tensor( + get_schedule(steps, image_sequence_length), + dtype=torch.float32, + ).detach() + + def _calculate_ays_schedule(model_type: str, steps: int) -> torch.Tensor: """Calculate full AYS sigmas for the requested step count.""" diff --git a/simple_syrup/runtime/tiled_sampling.py b/simple_syrup/runtime/tiled_sampling.py index fde1bbf..b6d6a21 100644 --- a/simple_syrup/runtime/tiled_sampling.py +++ b/simple_syrup/runtime/tiled_sampling.py @@ -15,12 +15,20 @@ from dataclasses import dataclass from typing import Any, TypeAlias import torch +import torch.nn.functional as functional -from ..domain.tiled_diffusion import LatentTile, TiledDiffusionPlan +from ..domain.contextual_diffusion import SpatialContext +from ..domain.tiled_diffusion import ( + LatentTile, + TiledDiffusionPlan, + gaussian_tile_weights, + validate_tiled_diffusion_mode, +) Latent: TypeAlias = dict[str, Any] ApplyModel: TypeAlias = Callable[..., torch.Tensor] ModelFunctionWrapper: TypeAlias = Callable[[ApplyModel, dict[str, Any]], torch.Tensor] +TileEvaluator: TypeAlias = Callable[[dict[str, Any]], torch.Tensor] UNSUPPORTED_CONDITIONING_KEYS = frozenset({"area", "control", "gligen"}) @@ -91,6 +99,110 @@ class SemanticTileWeightCache: return weights.model[index], weights.accumulation[index] +class TileBlendWeightCache: + """Resolve authoritative overlap weights for either tiled diffusion policy.""" + + def __init__(self, plan: TiledDiffusionPlan, diffusion_mode: str) -> None: + """Create weight caches for one immutable tiled prediction plan.""" + + validate_tiled_diffusion_mode(diffusion_mode) + self._diffusion_mode = diffusion_mode + self._semantic_weights = SemanticTileWeightCache(plan.tiles) + self._gaussian_weights: dict[ + tuple[torch.device, torch.dtype, int, int, int], torch.Tensor + ] = {} + + def for_tile( + self, + output: torch.Tensor, + tile: LatentTile, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Return model and accumulation weights for one predicted tile.""" + + semantic_weights = self._semantic_weights.for_output(output) + model_weight, accumulation_weight = self._semantic_weights.for_tile( + semantic_weights, + tile, + ) + if self._diffusion_mode == "multidiffusion": + return model_weight, accumulation_weight + + gaussian_weight = self._gaussian_for(output, tile) + return ( + model_weight * gaussian_weight, + accumulation_weight * gaussian_weight.to(dtype=torch.float32), + ) + + def _gaussian_for( + self, + output: torch.Tensor, + tile: LatentTile, + ) -> torch.Tensor: + """Return cached Mixture of Diffusers weights for one tile shape.""" + + cache_key = ( + output.device, + output.dtype, + output.ndim, + tile.width, + tile.height, + ) + cached = self._gaussian_weights.get(cache_key) + if cached is not None: + return cached + weights = gaussian_tile_weights( + tile.width, + tile.height, + device=output.device, + dtype=output.dtype, + ).reshape((1,) * (output.ndim - 2) + (tile.height, tile.width)) + self._gaussian_weights[cache_key] = weights + return weights + + +class TilePredictionAccumulator: + """Evaluate tiled model views and combine them with one selected policy.""" + + def __init__(self, plan: TiledDiffusionPlan, *, diffusion_mode: str) -> None: + """Bind an immutable plan to its overlap weighting policy.""" + + self._plan = plan + self._blend_weights = TileBlendWeightCache(plan, diffusion_mode) + + def predict( + self, + *, + args: dict[str, Any], + x: torch.Tensor, + evaluate: TileEvaluator, + ) -> torch.Tensor: + """Return one canvas prediction accumulated from bounded model views.""" + + output_buffer = torch.zeros_like(x) + weight_buffer = new_spatial_weight_buffer(x, self._plan) + input_batch_size = int(x.shape[0]) + for batch in self._plan.batches: + tiled_args = make_tiled_model_args( + args=args, + tiles=batch, + input_batch_size=input_batch_size, + latent_height=self._plan.latent_height, + latent_width=self._plan.latent_width, + ) + tile_output = evaluate(tiled_args) + for index, tile in enumerate(batch): + tile_slice = spatial_tile_slicer(tile, x.ndim) + start = index * input_batch_size + end = start + input_batch_size + model_weight, accumulation_weight = self._blend_weights.for_tile( + tile_output, + tile, + ) + output_buffer[tile_slice] += tile_output[start:end] * model_weight + weight_buffer[tile_slice] += accumulation_weight + return output_buffer / weight_buffer.to(dtype=output_buffer.dtype) + + def validate_sampling_controls( *, steps: int, @@ -270,6 +382,191 @@ def make_tiled_model_args( return tiled_args +def make_spatial_context_model_args( + *, + args: dict[str, Any], + contexts: Sequence[SpatialContext], + input_batch_size: int, + latent_height: int, + latent_width: int, +) -> dict[str, Any]: + """Create apply-model arguments for equally shaped spatial contexts.""" + + if not contexts: + raise ValueError("Spatial model arguments require at least one context.") + target_shape = (contexts[0].context_height, contexts[0].context_width) + if any( + (context.context_height, context.context_width) != target_shape + for context in contexts + ): + raise ValueError("Batched spatial contexts must use one model context shape.") + x = args["input"] + timestep = args["timestep"] + conditioning = args.get("c", {}) + if not isinstance(x, torch.Tensor): + raise ValueError("contextual sampler model input must be a tensor.") + if not isinstance(timestep, torch.Tensor): + raise ValueError("contextual sampler timestep must be a tensor.") + if not isinstance(conditioning, dict): + raise ValueError("contextual sampler conditioning must be a dict.") + + context_x = torch.cat( + [ + resize_spatial_tensor( + x[spatial_context_slicer(context, x.ndim)], + height=context.context_height, + width=context.context_width, + mode="nearest-exact", + ) + for context in contexts + ], + dim=0, + ) + context_timestep = torch.cat([timestep] * len(contexts), dim=0) + context_conditioning = spatial_context_conditioning( + conditioning=conditioning, + contexts=contexts, + input_batch_size=input_batch_size, + latent_height=latent_height, + latent_width=latent_width, + context_timestep=context_timestep, + ) + context_args = args.copy() + context_args["input"] = context_x + context_args["timestep"] = context_timestep + context_args["c"] = context_conditioning + if "cond_or_uncond" in args: + context_args["cond_or_uncond"] = repeat_sequence( + args["cond_or_uncond"], + len(contexts), + ) + return context_args + + +def spatial_context_conditioning( + *, + conditioning: dict[str, Any], + contexts: Sequence[SpatialContext], + input_batch_size: int, + latent_height: int, + latent_width: int, + context_timestep: torch.Tensor, +) -> dict[str, Any]: + """Resize spatial conditioning alongside arbitrary latent contexts.""" + + transformed: dict[str, Any] = {} + for key, value in conditioning.items(): + if key == "transformer_options" and isinstance(value, dict): + transformed[key] = tile_transformer_options( + value, + tile_count=len(contexts), + tiled_timestep=context_timestep, + ) + continue + transformed[key] = spatial_context_value( + value, + contexts=contexts, + input_batch_size=input_batch_size, + latent_height=latent_height, + latent_width=latent_width, + ) + return transformed + + +def spatial_context_value( + value: Any, + *, + contexts: Sequence[SpatialContext], + input_batch_size: int, + latent_height: int, + latent_width: int, +) -> Any: + """Transform tensors nested inside one spatial-context conditioning value.""" + + if isinstance(value, torch.Tensor): + if value.ndim >= 4 and value.shape[-2:] == ( + latent_height, + latent_width, + ): + return torch.cat( + [ + resize_spatial_tensor( + value[spatial_context_slicer(context, value.ndim)], + height=context.context_height, + width=context.context_width, + mode="nearest-exact", + ) + for context in contexts + ], + dim=0, + ) + if value.ndim >= 1 and value.shape[0] == input_batch_size: + return torch.cat([value] * len(contexts), dim=0) + if value.ndim >= 1 and value.shape[0] == 1: + repeats = [input_batch_size * len(contexts)] + [1] * (value.ndim - 1) + return value.repeat(repeats) + return value + if isinstance(value, list): + return [ + spatial_context_value( + item, + contexts=contexts, + input_batch_size=input_batch_size, + latent_height=latent_height, + latent_width=latent_width, + ) + for item in value + ] + if isinstance(value, tuple): + return tuple( + spatial_context_value( + item, + contexts=contexts, + input_batch_size=input_batch_size, + latent_height=latent_height, + latent_width=latent_width, + ) + for item in value + ) + return value + + +def spatial_context_slicer( + context: SpatialContext, + tensor_ndim: int, +) -> tuple[slice, ...]: + """Return a slicer for one arbitrary spatial context rectangle.""" + + return ( + (slice(None),) * (tensor_ndim - 2) + + (slice(context.y, context.y + context.height),) + + (slice(context.x, context.x + context.width),) + ) + + +def resize_spatial_tensor( + tensor: torch.Tensor, + *, + height: int, + width: int, + mode: str, +) -> torch.Tensor: + """Resize only the final two axes of a 4D or singleton-depth 5D tensor.""" + + if tensor.shape[-2:] == (height, width): + return tensor + leading_shape = tensor.shape[:-2] + flattened = tensor.reshape(-1, 1, tensor.shape[-2], tensor.shape[-1]) + align_corners = False if mode in {"bilinear", "bicubic"} else None + resized = functional.interpolate( + flattened, + size=(height, width), + mode=mode, + align_corners=align_corners, + ) + return resized.reshape(*leading_shape, height, width) + + def tile_conditioning( *, conditioning: dict[str, Any], diff --git a/simple_syrup/services/contextual_diffusion_sampling_service.py b/simple_syrup/services/contextual_diffusion_sampling_service.py new file mode 100644 index 0000000..7f14550 --- /dev/null +++ b/simple_syrup/services/contextual_diffusion_sampling_service.py @@ -0,0 +1,177 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Application service for contextual diffusion latent sampling.""" + +from __future__ import annotations + +from typing import Any, TypeAlias + +import torch + +from ..domain.conditioning_batch import ConditioningBatch, select_conditioning +from ..domain.contextual_diffusion import ( + ContextualDiffusionControls, + build_contextual_diffusion_plan, +) +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 .sampling_batch import ( + combine_latent_outputs, + latent_batch_size, + single_item_latent, +) + +Latent: TypeAlias = dict[str, Any] + + +class ContextualDiffusionSamplingService: + """Plan and execute composition-preserving contextual diffusion.""" + + def sample( + self, + *, + model: Any, + seed: int, + steps: int, + cfg: float, + sampler_name: str, + scheduler: str, + positive: Any, + negative: Any, + latent_image: Latent, + denoise: float, + diffusion_mode: str, + latent_context_size: int, + latent_context_overlap: int, + latent_context_batch_size: int, + global_weight: float, + global_steps: int, + global_decay: float, + segs: object | None = None, + ) -> Latent: + """Sample a latent with global and bounded detail contexts.""" + + validate_tiled_diffusion_mode(diffusion_mode) + controls = ContextualDiffusionControls( + latent_context_size=latent_context_size, + latent_context_overlap=latent_context_overlap, + latent_context_batch_size=latent_context_batch_size, + global_weight=global_weight, + global_steps=global_steps, + global_decay=global_decay, + ) + controls.validate() + 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): + raise ValueError( + "Contextual Diffusion requires one SEGS payload or one per latent " + f"batch item; received {len(segs_group)} for batch size {batch_size}." + ) + split_batch = ( + bool(segs_group) + or isinstance(positive, ConditioningBatch) + or isinstance(negative, ConditioningBatch) + ) + if not split_batch: + return self._sample_item( + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=positive, + negative=negative, + latent_image=latent_image, + denoise=denoise, + diffusion_mode=diffusion_mode, + controls=controls, + segs=None, + ) + + outputs: list[torch.Tensor] = [] + for index in range(batch_size): + item_output = self._sample_item( + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=( + select_conditioning(positive, index) + if isinstance(positive, ConditioningBatch) + else positive + ), + negative=( + select_conditioning(negative, index) + if isinstance(negative, ConditioningBatch) + else negative + ), + latent_image=single_item_latent(latent_image, index), + denoise=denoise, + diffusion_mode=diffusion_mode, + controls=controls, + segs=( + segs_group[0 if len(segs_group) == 1 else index] + if segs_group + else None + ), + ) + samples = item_output.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) + + def _sample_item( + self, + *, + model: Any, + seed: int, + steps: int, + cfg: float, + sampler_name: str, + scheduler: str, + positive: Any, + negative: Any, + latent_image: Latent, + denoise: float, + diffusion_mode: str, + controls: ContextualDiffusionControls, + segs: NativeSegs | None, + ) -> Latent: + """Build one canvas plan and execute it through the runtime adapter.""" + + samples = latent_image.get("samples") + if not isinstance(samples, torch.Tensor): + raise TypeError( + "Contextual Diffusion latent samples must be a torch.Tensor." + ) + plan = build_contextual_diffusion_plan( + latent_width=int(samples.shape[-1]), + latent_height=int(samples.shape[-2]), + controls=controls, + segs=segs, + ) + return sample_contextual_diffusion( + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=positive, + negative=negative, + latent_image=latent_image, + denoise=denoise, + diffusion_mode=diffusion_mode, + controls=controls, + plan=plan, + ) diff --git a/simple_syrup/services/ksampler_sampling_service.py b/simple_syrup/services/ksampler_sampling_service.py index f3fb921..86066b3 100644 --- a/simple_syrup/services/ksampler_sampling_service.py +++ b/simple_syrup/services/ksampler_sampling_service.py @@ -39,13 +39,6 @@ class KSamplerSamplingService: """Sample a latent with configured SimpleSyrup sampler extensions.""" sampler = sampling_samplers.resolve_sampler(sampler_name) - sigmas = sampling_schedulers.calculate_sigmas( - model=model, - scheduler_name=scheduler, - sampler_name=sampler_name, - steps=steps, - denoise=denoise, - ).to(model.load_device) latent_samples = latent_image["samples"] if not isinstance(latent_samples, torch.Tensor): raise TypeError("KSampler latent samples must be a torch.Tensor.") @@ -61,6 +54,14 @@ class KSamplerSamplingService: raise TypeError( "KSampler normalized latent samples must be a torch.Tensor." ) + sigmas = sampling_schedulers.calculate_sigmas( + model=model, + scheduler_name=scheduler, + sampler_name=sampler_name, + steps=steps, + denoise=denoise, + view=sampling_schedulers.SchedulerView.from_tensor(latent_samples), + ).to(model.load_device) noise = comfy_sample.prepare_noise( latent_samples, seed, diff --git a/simple_syrup/services/sampling_batch.py b/simple_syrup/services/sampling_batch.py new file mode 100644 index 0000000..0fbfbb8 --- /dev/null +++ b/simple_syrup/services/sampling_batch.py @@ -0,0 +1,56 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Share latent-batch adaptation across sampling application services.""" + +from __future__ import annotations + +from typing import Any, TypeAlias + +import torch + +Latent: TypeAlias = dict[str, Any] + + +def latent_batch_size(latent_image: Latent) -> int: + """Return the validated number of latent batch items.""" + + samples = latent_image.get("samples") + if not isinstance(samples, torch.Tensor): + raise TypeError("Sampler latent samples must be a torch.Tensor.") + if samples.ndim < 1: + raise ValueError("Sampler latent samples must include a batch axis.") + return int(samples.shape[0]) + + +def single_item_latent(latent_image: Latent, index: int) -> Latent: + """Return one latent batch item with aligned batch metadata and mask.""" + + samples = latent_image.get("samples") + if not isinstance(samples, torch.Tensor): + raise TypeError("Sampler latent samples must be a torch.Tensor.") + item = latent_image.copy() + item["samples"] = samples[index : index + 1] + if "batch_index" in item: + item["batch_index"] = [item["batch_index"][index]] + noise_mask = item.get("noise_mask") + if isinstance(noise_mask, torch.Tensor) and noise_mask.shape[0] == int( + samples.shape[0] + ): + item["noise_mask"] = noise_mask[index : index + 1] + return item + + +def combine_latent_outputs( + latent_image: Latent, + outputs: list[torch.Tensor], +) -> Latent: + """Return one latent dictionary containing ordered sampled batch outputs.""" + + if not outputs: + raise ValueError("Sampler batch execution produced no outputs.") + result = latent_image.copy() + result.pop("downscale_ratio_spacial", None) + result["samples"] = torch.cat(outputs, dim=0) + return result diff --git a/simple_syrup/services/tiled_diffusion_sampling_service.py b/simple_syrup/services/tiled_diffusion_sampling_service.py index a0fe716..e05e5d6 100644 --- a/simple_syrup/services/tiled_diffusion_sampling_service.py +++ b/simple_syrup/services/tiled_diffusion_sampling_service.py @@ -16,6 +16,7 @@ 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 .sampling_batch import combine_latent_outputs, single_item_latent Latent = dict[str, Any] @@ -219,7 +220,7 @@ class TiledDiffusionSamplingService: outputs: list[torch.Tensor] = [] for index in range(batch_size): - item_latent = self._single_item_latent(latent_image, index) + item_latent = single_item_latent(latent_image, index) samples = item_latent["samples"] if not isinstance(samples, torch.Tensor): raise TypeError( @@ -266,10 +267,7 @@ class TiledDiffusionSamplingService: "Tiled diffusion output samples must be a torch.Tensor." ) outputs.append(output_samples) - result = latent_image.copy() - result.pop("downscale_ratio_spacial", None) - result["samples"] = torch.cat(outputs, dim=0) - return result + return combine_latent_outputs(latent_image, outputs) def _sample_conditioning_batch( self, @@ -301,7 +299,7 @@ class TiledDiffusionSamplingService: outputs: list[torch.Tensor] = [] for index in range(int(latent_samples.shape[0])): - item_latent = self._single_item_latent(latent_image, index) + item_latent = single_item_latent(latent_image, index) output = self.sample( diffusion_mode=diffusion_mode, model=model, @@ -329,42 +327,7 @@ class TiledDiffusionSamplingService: ) outputs.append(output_samples) - result = latent_image.copy() - result.pop("downscale_ratio_spacial", None) - result["samples"] = torch.cat(outputs, dim=0) - return result - - def _single_item_latent(self, latent_image: Latent, index: int) -> Latent: - """Return a latent dictionary for one batch item.""" - - latent_samples = latent_image["samples"] - if not isinstance(latent_samples, torch.Tensor): - raise TypeError("Tiled diffusion latent samples must be a torch.Tensor.") - item = latent_image.copy() - item["samples"] = latent_samples[index : index + 1] - if "batch_index" in item: - item["batch_index"] = [item["batch_index"][index]] - if "noise_mask" in item: - item["noise_mask"] = self._slice_noise_mask( - item["noise_mask"], - index, - latent_samples, - ) - return item - - def _slice_noise_mask( - self, - noise_mask: Any, - index: int, - latent_samples: torch.Tensor, - ) -> Any: - """Return the noise mask slice matching one latent batch item.""" - - if isinstance(noise_mask, torch.Tensor) and noise_mask.shape[0] == int( - latent_samples.shape[0], - ): - return noise_mask[index : index + 1] - return noise_mask + return combine_latent_outputs(latent_image, outputs) def _uses_conditioning_batch(self, positive: Any, negative: Any) -> bool: """Return whether tiled sampling needs per-item conditioning selection.""" diff --git a/tests/test_contextual_diffusion_domain.py b/tests/test_contextual_diffusion_domain.py new file mode 100644 index 0000000..16eeafa --- /dev/null +++ b/tests/test_contextual_diffusion_domain.py @@ -0,0 +1,98 @@ +# 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 deterministic contextual diffusion planning.""" + +from __future__ import annotations + +import pytest +import torch + +from simple_syrup.domain.contextual_diffusion import ( + ContextualDiffusionControls, + build_contextual_diffusion_plan, + fit_context_shape, +) +from simple_syrup.domain.segs import BoundingBox, CropRegion, Segment + + +def test_global_context_fits_maximum_dimension_without_boxing() -> None: + """The whole canvas keeps its aspect ratio and uses no artificial padding.""" + + assert fit_context_shape(384, 256, 128) == (128, 86) + assert fit_context_shape(64, 96, 128) == (64, 96) + + +def test_connected_segs_replace_regular_grid_with_guided_tile_plan() -> None: + """Contextual Diffusion delegates its tile path to SEGS-guided planning.""" + + large = torch.ones((64, 64), dtype=torch.float32) + small = torch.zeros((64, 64), dtype=torch.float32) + small[24:40, 24:40] = 1.0 + plan = build_contextual_diffusion_plan( + latent_width=64, + latent_height=64, + controls=_controls(), + segs=((64, 64), (_segment(large, 0.8), _segment(small, 0.9))), + ) + + assert len(plan.tile_plan.tiles) > 1 + assert all(tile.weight_mask is not None for tile in plan.tile_plan.tiles) + + +def test_missing_segs_uses_regular_tiled_diffusion_plan() -> None: + """The optional SEGS input preserves ordinary tiled diffusion as fallback.""" + + plan = build_contextual_diffusion_plan( + latent_width=64, + latent_height=64, + controls=_controls(), + segs=None, + ) + + assert len(plan.tile_plan.tiles) > 1 + assert all(tile.weight_mask is None for tile in plan.tile_plan.tiles) + + +def test_invalid_overlap_fails_before_planning() -> None: + """An overlap that cannot advance a context is rejected explicitly.""" + + with pytest.raises(ValueError, match="latent_context_overlap"): + build_contextual_diffusion_plan( + latent_width=64, + latent_height=64, + controls=_controls(latent_context_overlap=32), + segs=None, + ) + + +def _controls( + *, + latent_context_overlap: int = 8, +) -> ContextualDiffusionControls: + """Return compact valid controls for planner tests.""" + + return ContextualDiffusionControls( + latent_context_size=32, + latent_context_overlap=latent_context_overlap, + latent_context_batch_size=2, + global_weight=1.0, + global_steps=1, + global_decay=0.5, + ) + + +def _segment(mask: torch.Tensor, confidence: float) -> Segment: + """Return a full-canvas SEG for a test mask.""" + + height, width = mask.shape + crop = CropRegion(0, 0, width, height) + return Segment( + cropped_image=None, + cropped_mask=mask, + confidence=confidence, + crop_region=crop, + bbox=BoundingBox(*crop), + label="region", + ) diff --git a/tests/test_contextual_diffusion_sampling.py b/tests/test_contextual_diffusion_sampling.py new file mode 100644 index 0000000..92291a0 --- /dev/null +++ b/tests/test_contextual_diffusion_sampling.py @@ -0,0 +1,289 @@ +# 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 contextual diffusion prediction fusion.""" + +from __future__ import annotations + +from typing import Any + +import pytest +import torch + +from simple_syrup.domain.contextual_diffusion import ( + ContextualDiffusionControls, + build_contextual_diffusion_plan, +) +from simple_syrup.runtime import ( + contextual_diffusion_sampling, + sampling_samplers, + sampling_schedulers, +) +from simple_syrup.runtime.contextual_diffusion_sampling import ( + ContextualDiffusionModelWrapper, +) + +comfy_sample = contextual_diffusion_sampling._comfy_sample() +comfy_utils = contextual_diffusion_sampling._comfy_utils() +latent_preview = contextual_diffusion_sampling._latent_preview() + + +def test_global_prediction_owns_low_frequency_while_tiles_keep_detail() -> None: + """The native view replaces tile-level scene intent without blurring detail.""" + + 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, + ) + calls: list[tuple[int, int]] = [] + + def apply_model( + x: torch.Tensor, + timestep: torch.Tensor, + **conditioning: object, + ) -> torch.Tensor: + """Return high-frequency tile detail and a constant whole-image intent.""" + + del timestep, conditioning + calls.append((int(x.shape[-2]), int(x.shape[-1]))) + if x.shape[-2:] == (8, 16): + return torch.full_like(x, 3.0) + rows = torch.arange(x.shape[-2], device=x.device).reshape(1, 1, -1, 1) + return torch.where(rows % 2 == 0, 1.0, -1.0).expand_as(x) + + output = wrapper( + apply_model, + { + "input": torch.zeros((1, 1, 16, 32)), + "timestep": torch.tensor([1.0]), + "c": {}, + }, + ) + + assert calls == [(16, 16), (8, 16)] + assert torch.allclose(output[:, :, 0::2], torch.full((1, 1, 8, 32), 4.0)) + assert torch.allclose(output[:, :, 1::2], torch.full((1, 1, 8, 32), 2.0)) + + +def test_global_prediction_decays_then_stops_after_configured_steps() -> None: + """Global authority decays before late denoising becomes local-only.""" + + controls = _controls(global_steps=2) + 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.75, 0.5, 0.25, 0.0]), + existing_wrapper=None, + ) + calls: list[tuple[int, int]] = [] + + def apply_model( + x: torch.Tensor, + timestep: torch.Tensor, + **conditioning: object, + ) -> torch.Tensor: + """Record whether each prediction uses local or global context.""" + + del timestep, conditioning + calls.append((int(x.shape[-2]), int(x.shape[-1]))) + value = 3.0 if x.shape[-2:] == (8, 16) else 1.0 + return torch.full_like(x, value) + + base_args = { + "input": torch.zeros((1, 1, 16, 32)), + "c": {}, + } + output = wrapper(apply_model, base_args | {"timestep": torch.tensor([0.75])}) + assert calls == [(16, 16), (8, 16)] + assert torch.allclose(output, torch.full_like(output, 2.0)) + + calls.clear() + output = wrapper(apply_model, base_args | {"timestep": torch.tensor([0.5])}) + assert calls == [(16, 16)] + assert torch.allclose(output, torch.ones_like(output)) + + +def test_native_sized_canvas_delegates_to_one_original_model_call() -> None: + """A canvas already inside the model view limit behaves like normal sampling.""" + + controls = _controls() + plan = build_contextual_diffusion_plan( + latent_width=16, + latent_height=16, + controls=controls, + segs=None, + ) + wrapper = ContextualDiffusionModelWrapper( + plan=plan, + controls=controls, + sigmas=torch.tensor([1.0, 0.0]), + existing_wrapper=None, + ) + calls = 0 + + def apply_model( + x: torch.Tensor, + timestep: torch.Tensor, + **conditioning: object, + ) -> torch.Tensor: + """Count direct model evaluations.""" + + del timestep, conditioning + nonlocal calls + calls += 1 + return x + 2.0 + + x = torch.zeros((1, 1, 16, 16)) + output = wrapper( + apply_model, + {"input": x, "timestep": torch.tensor([1.0]), "c": {}}, + ) + + assert calls == 1 + assert torch.equal(output, x + 2.0) + + +def test_runtime_delegates_sampling_to_comfy_with_wrapped_clone( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The vertical runtime path preserves KSampler sampling and latent metadata.""" + + model = _FakeModel() + sampler = object() + latent_samples = torch.zeros((1, 4, 16, 32)) + sampled = torch.ones_like(latent_samples) + latent = { + "samples": latent_samples, + "downscale_ratio_spacial": 2, + "kept": "metadata", + } + controls = _controls() + plan = build_contextual_diffusion_plan( + latent_width=32, + latent_height=16, + controls=controls, + segs=None, + ) + calls: dict[str, Any] = {} + + monkeypatch.setattr( + sampling_samplers, + "resolve_sampler", + lambda _name: sampler, + ) + monkeypatch.setattr( + sampling_schedulers, + "calculate_sigmas", + lambda **_kwargs: torch.tensor([1.0, 0.0]), + ) + monkeypatch.setattr( + comfy_sample, + "fix_empty_latent_channels", + lambda _model, samples, _ratio: samples, + ) + monkeypatch.setattr( + comfy_sample, + "prepare_noise", + lambda samples, _seed, _batch_inds=None: torch.ones_like(samples), + ) + monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None) + monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", False) + + def fake_sample_custom( + sampling_model: _FakeModel, + noise: torch.Tensor, + cfg: float, + received_sampler: object, + sigmas: torch.Tensor, + positive: object, + negative: object, + latent_image: torch.Tensor, + **kwargs: object, + ) -> torch.Tensor: + """Capture the final Comfy boundary call.""" + + del noise, cfg, sigmas, positive, negative, latent_image, kwargs + calls["model"] = sampling_model + calls["sampler"] = received_sampler + return sampled + + monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) + + output = contextual_diffusion_sampling.sample_contextual_diffusion( + model=model, + seed=7, + steps=2, + cfg=1.0, + sampler_name="euler", + scheduler="simple", + positive=[], + negative=[], + latent_image=latent, + denoise=0.8, + diffusion_mode="mixture_of_diffusers", + controls=controls, + plan=plan, + ) + + assert calls["model"] is not model + assert isinstance(calls["model"].wrapper, ContextualDiffusionModelWrapper) + assert calls["model"].wrapper.diffusion_mode == "mixture_of_diffusers" + assert calls["sampler"] is sampler + assert output["samples"] is sampled + assert output["kept"] == "metadata" + assert "downscale_ratio_spacial" not in output + + +def _controls( + *, + global_steps: int = 1, + global_decay: float = 0.5, +) -> ContextualDiffusionControls: + """Return a small two-tile test configuration.""" + + return ContextualDiffusionControls( + latent_context_size=16, + latent_context_overlap=0, + latent_context_batch_size=2, + global_weight=1.0, + global_steps=global_steps, + global_decay=global_decay, + ) + + +class _FakeModel: + """Provide the ModelPatcher surface used by the semantic runtime.""" + + def __init__(self, model_options: dict[str, Any] | None = None) -> None: + """Create a CPU-backed fake model patcher.""" + + self.load_device = torch.device("cpu") + self.model_options = {} if model_options is None else model_options + self.wrapper: object | None = None + + def clone(self) -> _FakeModel: + """Return a clone with copied model options.""" + + return _FakeModel(self.model_options.copy()) + + def set_model_unet_function_wrapper(self, wrapper: object) -> None: + """Capture the installed wrapper.""" + + self.wrapper = wrapper + self.model_options["model_function_wrapper"] = wrapper diff --git a/tests/test_contextual_diffusion_sampling_service.py b/tests/test_contextual_diffusion_sampling_service.py new file mode 100644 index 0000000..a7522ec --- /dev/null +++ b/tests/test_contextual_diffusion_sampling_service.py @@ -0,0 +1,178 @@ +# 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 contextual diffusion sampling orchestration.""" + +from __future__ import annotations + +from typing import Any + +import pytest +import torch + +from simple_syrup.domain.segs import BoundingBox, CropRegion, Segment +from simple_syrup.services import ( + contextual_diffusion_sampling_service as service_module, +) +from simple_syrup.services.contextual_diffusion_sampling_service import ( + ContextualDiffusionSamplingService, +) + + +def test_service_builds_global_context_and_segs_guided_tile_plan( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Connected SEGS reach the runtime through the existing guided tile planner.""" + + calls: list[dict[str, Any]] = [] + + def fake_sample_contextual_diffusion(**kwargs: Any) -> dict[str, Any]: + """Record the completed plan and return the input 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, + ) + latent = {"samples": torch.zeros((1, 4, 64, 96))} + + result = ContextualDiffusionSamplingService().sample( + **_sample_kwargs(latent=latent, segs=_segs(512, 768)) + ) + + assert torch.equal(result["samples"], latent["samples"]) + assert len(calls) == 1 + assert calls[0]["diffusion_mode"] == "mixture_of_diffusers" + plan = calls[0]["plan"] + assert ( + plan.global_context.context_width, + plan.global_context.context_height, + ) == (64, 44) + assert len(plan.tile_plan.tiles) > 1 + assert all(tile.weight_mask is not None for tile in plan.tile_plan.tiles) + + +def test_service_rejects_segs_batch_mismatch_before_runtime( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Ambiguous per-image semantic guidance fails before model sampling.""" + + called = False + + def fake_sample_contextual_diffusion(**kwargs: Any) -> dict[str, Any]: + """Mark unexpected runtime entry.""" + + del kwargs + nonlocal called + called = True + return {} + + monkeypatch.setattr( + service_module, + "sample_contextual_diffusion", + fake_sample_contextual_diffusion, + ) + latent = {"samples": torch.zeros((2, 4, 64, 96))} + + with pytest.raises(ValueError, match="one SEGS payload or one per latent"): + ContextualDiffusionSamplingService().sample( + **_sample_kwargs( + latent=latent, + segs=[_segs(512, 768), _segs(512, 768), _segs(512, 768)], + ) + ) + + assert not called + + +def test_service_rejects_invalid_controls_before_runtime( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Invalid context geometry fails without sampler side effects.""" + + monkeypatch.setattr( + service_module, + "sample_contextual_diffusion", + lambda **_kwargs: pytest.fail("runtime should not be called"), + ) + + with pytest.raises(ValueError, match="latent_context_overlap"): + ContextualDiffusionSamplingService().sample( + **( + _sample_kwargs( + latent={"samples": torch.zeros((1, 4, 64, 96))}, + segs=None, + ) + | {"latent_context_overlap": 64} + ) + ) + + +def test_service_rejects_unknown_diffusion_mode_before_runtime( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Unknown tile blending policies fail before model sampling begins.""" + + monkeypatch.setattr( + service_module, + "sample_contextual_diffusion", + lambda **_kwargs: pytest.fail("runtime should not be called"), + ) + + with pytest.raises(ValueError, match="diffusion_mode must be one of"): + ContextualDiffusionSamplingService().sample( + **( + _sample_kwargs( + latent={"samples": torch.zeros((1, 4, 64, 96))}, + segs=None, + ) + | {"diffusion_mode": "unknown"} + ) + ) + + +def _sample_kwargs(*, latent: dict[str, Any], segs: object | None) -> dict[str, Any]: + """Return one valid service request.""" + + return { + "model": object(), + "seed": 1, + "steps": 4, + "cfg": 1.0, + "sampler_name": "euler", + "scheduler": "simple", + "positive": [], + "negative": [], + "latent_image": latent, + "denoise": 0.5, + "diffusion_mode": "mixture_of_diffusers", + "latent_context_size": 64, + "latent_context_overlap": 8, + "latent_context_batch_size": 2, + "global_weight": 1.0, + "global_steps": 1, + "global_decay": 0.5, + "segs": segs, + } + + +def _segs(height: int, width: int) -> object: + """Return one full-image Impact-compatible SEG payload.""" + + crop = CropRegion(0, 0, width, height) + segment = Segment( + cropped_image=None, + cropped_mask=torch.ones((height, width), dtype=torch.float32), + confidence=1.0, + crop_region=crop, + bbox=BoundingBox(*crop), + label="subject", + ) + return ((height, width), (segment,)) diff --git a/tests/test_ksampler_contextual_diffusion_node.py b/tests/test_ksampler_contextual_diffusion_node.py new file mode 100644 index 0000000..944945b --- /dev/null +++ b/tests/test_ksampler_contextual_diffusion_node.py @@ -0,0 +1,152 @@ +# 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 the KSampler Contextual Diffusion node contract.""" + +from __future__ import annotations + +from typing import Any + +import pytest +import torch + +from simple_syrup.nodes.ksampler_contextual_diffusion import ( + KSamplerContextualDiffusion, +) +from simple_syrup.runtime import sampling_samplers, sampling_schedulers + + +def test_input_types_expose_concise_klein_oriented_controls( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The node keeps KSampler inputs and bounded contextual settings.""" + + monkeypatch.setattr(sampling_samplers, "available_samplers", lambda: ("euler",)) + monkeypatch.setattr( + sampling_schedulers, + "available_schedulers", + lambda: ("simple",), + ) + + declared = KSamplerContextualDiffusion.INPUT_TYPES() + required = declared["required"] + + assert tuple(required) == ( + "model", + "seed", + "steps", + "cfg", + "sampler_name", + "scheduler", + "positive", + "negative", + "latent_image", + "denoise", + "diffusion_mode", + "latent_context_size", + "latent_context_overlap", + "latent_context_batch_size", + "global_weight", + "global_steps", + "global_decay", + ) + assert required["steps"][1]["default"] == 4 + assert required["cfg"][1]["default"] == 1.0 + assert required["diffusion_mode"][0] == [ + "multidiffusion", + "mixture_of_diffusers", + ] + assert required["diffusion_mode"][1]["default"] == "multidiffusion" + assert required["latent_context_size"][1]["default"] == 96 + assert required["latent_context_overlap"][1]["default"] == 32 + assert required["latent_context_batch_size"][1]["default"] == 4 + assert required["global_weight"][1]["default"] == 1.0 + assert required["global_steps"][1]["default"] == 1 + assert required["global_decay"][1]["default"] == 0.5 + assert declared["optional"]["segs"][0] == "SEGS" + + +def test_node_metadata_matches_separate_sampler_contract() -> None: + """Contextual Diffusion remains a distinct sampler with latent output.""" + + assert KSamplerContextualDiffusion.RETURN_TYPES == ("LATENT",) + assert KSamplerContextualDiffusion.FUNCTION == "sample" + assert KSamplerContextualDiffusion.CATEGORY == "SimpleSyrup/Sampling" + assert "composition" in KSamplerContextualDiffusion.DESCRIPTION + + +def test_sample_delegates_every_control_to_service( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The node API owns no contextual planning or runtime behavior.""" + + fake_service = _FakeContextualDiffusionService() + monkeypatch.setattr( + KSamplerContextualDiffusion, + "service_class", + staticmethod(lambda: fake_service), + ) + latent = {"samples": torch.zeros((1, 4, 32, 48))} + segs = object() + + (result,) = KSamplerContextualDiffusion().sample( + model="model", + seed=12, + steps=8, + cfg=1.0, + sampler_name="euler", + scheduler="simple", + positive="positive", + negative="negative", + latent_image=latent, + denoise=0.7, + diffusion_mode="mixture_of_diffusers", + latent_context_size=96, + latent_context_overlap=12, + latent_context_batch_size=3, + global_weight=0.9, + global_steps=2, + global_decay=0.4, + segs=segs, + ) + + assert result is fake_service.output + assert fake_service.calls == [ + { + "model": "model", + "seed": 12, + "steps": 8, + "cfg": 1.0, + "sampler_name": "euler", + "scheduler": "simple", + "positive": "positive", + "negative": "negative", + "latent_image": latent, + "denoise": 0.7, + "diffusion_mode": "mixture_of_diffusers", + "latent_context_size": 96, + "latent_context_overlap": 12, + "latent_context_batch_size": 3, + "global_weight": 0.9, + "global_steps": 2, + "global_decay": 0.4, + "segs": segs, + } + ] + + +class _FakeContextualDiffusionService: + """Record node delegation without entering the Comfy runtime.""" + + def __init__(self) -> None: + """Create a stable output and empty call history.""" + + self.output: dict[str, Any] = {"samples": torch.ones((1, 4, 32, 48))} + self.calls: list[dict[str, Any]] = [] + + def sample(self, **kwargs: Any) -> dict[str, Any]: + """Record one call and return the stable latent.""" + + self.calls.append(kwargs) + return self.output diff --git a/tests/test_ksampler_extras_node.py b/tests/test_ksampler_extras_node.py index dd65dee..e33f2c6 100644 --- a/tests/test_ksampler_extras_node.py +++ b/tests/test_ksampler_extras_node.py @@ -111,6 +111,7 @@ def test_scheduler_options_include_extras_and_exclude_svd() -> None: assert "GITS" in scheduler_options assert "beta57" in scheduler_options assert "automatic_a1111" in scheduler_options + assert "Flux2" in scheduler_options assert "AYS SVD" not in scheduler_options @@ -152,6 +153,8 @@ def test_sample_delegates_to_runtime_helpers( sampler_name: str, steps: int, denoise: float, + *, + view: sampling_schedulers.SchedulerView, ) -> torch.Tensor: """Record scheduler calculation.""" @@ -161,6 +164,7 @@ def test_sample_delegates_to_runtime_helpers( "sampler_name": sampler_name, "steps": steps, "denoise": denoise, + "view": view, } return fixed_sigmas @@ -286,6 +290,10 @@ def test_sample_delegates_to_runtime_helpers( "sampler_name": "lcm", "steps": 2, "denoise": 0.8, + "view": sampling_schedulers.SchedulerView( + latent_width=8, + latent_height=8, + ), } assert calls["prepare_noise"]["batch_inds"] == [0] assert calls["sample_custom"]["noise_mask"] is latent_image["noise_mask"] diff --git a/tests/test_mixture_of_diffusers_sampling.py b/tests/test_mixture_of_diffusers_sampling.py index 19c8c60..d452eeb 100644 --- a/tests/test_mixture_of_diffusers_sampling.py +++ b/tests/test_mixture_of_diffusers_sampling.py @@ -533,6 +533,9 @@ def test_sample_delegates_to_comfy_sampling_with_cloned_wrapped_model( assert calls["sample_custom"]["sampler"] is sampler assert calls["sample_custom"]["sigmas"] is fixed_sigmas assert calls["sample_custom"]["disable_pbar"] is True + assert calls["calculate_sigmas"]["view"] == ( + sampling_schedulers.SchedulerView(latent_width=4, latent_height=4) + ) def test_sample_accepts_singleton_depth_5d_latent( diff --git a/tests/test_multidiffusion_sampling.py b/tests/test_multidiffusion_sampling.py index 8f6ac87..194b219 100644 --- a/tests/test_multidiffusion_sampling.py +++ b/tests/test_multidiffusion_sampling.py @@ -678,6 +678,9 @@ def test_sample_delegates_to_comfy_sampling_with_cloned_wrapped_model( assert calls["sample_custom"]["sampler"] is sampler assert calls["sample_custom"]["noise"] is fixed_noise assert calls["sample_custom"]["disable_pbar"] is True + assert calls["calculate_sigmas"]["view"] == ( + sampling_schedulers.SchedulerView(latent_width=4, latent_height=4) + ) def test_sample_accepts_singleton_depth_5d_latent( diff --git a/tests/test_persisted_widget_order_contract.py b/tests/test_persisted_widget_order_contract.py index c748c32..05a52c4 100644 --- a/tests/test_persisted_widget_order_contract.py +++ b/tests/test_persisted_widget_order_contract.py @@ -19,6 +19,9 @@ from simple_syrup.nodes.detail_segs_by_scale_factor_tiled_diffusion import ( from simple_syrup.nodes.detect_segs_with_ultralytics import DetectSEGSWithUltralytics from simple_syrup.nodes.grounding_dino_model_loader import GroundingDINOModelLoader from simple_syrup.nodes.image_resize_to_target import ResizeImageToTarget +from simple_syrup.nodes.ksampler_contextual_diffusion import ( + KSamplerContextualDiffusion, +) from simple_syrup.nodes.ksampler_extras import KSamplerExtras from simple_syrup.nodes.ksampler_tiled_diffusion import KSamplerTiledDiffusion from simple_syrup.nodes.load_ultralytics_model import LoadUltralyticsModel @@ -146,6 +149,24 @@ _PERSISTED_WIDGET_PREFIXES: tuple[tuple[type[_ClassicNode], tuple[str, ...]], .. "latent_tile_batch_size", ), ), + ( + KSamplerContextualDiffusion, + ( + "seed", + "steps", + "cfg", + "sampler_name", + "scheduler", + "denoise", + "diffusion_mode", + "latent_context_size", + "latent_context_overlap", + "latent_context_batch_size", + "global_weight", + "global_steps", + "global_decay", + ), + ), (LoadUltralyticsModel, ("model_name",)), (PromptEncodeStyle, ("encode_style",)), ( diff --git a/tests/test_registration.py b/tests/test_registration.py index a127c4f..c3b842b 100644 --- a/tests/test_registration.py +++ b/tests/test_registration.py @@ -33,6 +33,7 @@ BASE_NODE_IDS = [ "SimpleSyrup.KSamplerExtras", "SimpleSyrup.KSamplerPromptByRegion", "SimpleSyrup.KSamplerPromptByTiledRegion", + "SimpleSyrup.KSamplerContextualDiffusion", "SimpleSyrup.KSamplerTiledDiffusion", "SimpleSyrup.LatentDiagnostics", "SimpleSyrup.LayerStyleSAMModelsAdapter", diff --git a/tests/test_sampling_schedulers.py b/tests/test_sampling_schedulers.py index 37d8d4b..6eb4fd6 100644 --- a/tests/test_sampling_schedulers.py +++ b/tests/test_sampling_schedulers.py @@ -8,6 +8,7 @@ from __future__ import annotations import math from collections.abc import Sequence +from importlib import import_module from typing import Any import comfy.samplers @@ -16,6 +17,7 @@ import torch from simple_syrup.runtime import sampling_schedulers from simple_syrup.runtime.sampling_schedulers import ( + SchedulerView, available_schedulers, calculate_sigmas, ) @@ -36,6 +38,32 @@ class FakeModel: return self.model_sampling +class FakeLatentFormat: + """Expose the spatial compression used to recover image dimensions.""" + + def __init__(self, spacial_downscale_ratio: int) -> None: + """Store one deterministic latent-to-image scale.""" + + self.spacial_downscale_ratio = spacial_downscale_ratio + + +class FakeModelWithLatentFormat(FakeModel): + """Provide model-sampling and latent-format objects for Flux2 tests.""" + + def __init__(self, spacial_downscale_ratio: int) -> None: + """Create a fake model with the requested spatial compression.""" + + super().__init__() + self.latent_format = FakeLatentFormat(spacial_downscale_ratio) + + def get_model_object(self, name: str) -> object: + """Return the requested fake model object.""" + + if name == "latent_format": + return self.latent_format + return super().get_model_object(name) + + class FakeDiscreteModelSampling: """Provide k-diffusion-style discrete sigma conversion for tests.""" @@ -199,15 +227,59 @@ def test_available_schedulers_includes_core_and_extras() -> None: for scheduler in comfy.samplers.KSampler.SCHEDULERS: assert scheduler in schedulers - assert schedulers[-5:] == ( + assert schedulers[-6:] == ( "AYS SD1", "AYS SDXL", "GITS", "beta57", "automatic_a1111", + "Flux2", ) +def test_flux2_schedule_matches_comfy_for_model_view_resolution() -> None: + """Flux2 delegates to ComfyUI using the effective model-view resolution.""" + + model = FakeModelWithLatentFormat(spacial_downscale_ratio=16) + sigmas = calculate_sigmas( + model, + "Flux2", + "euler", + 4, + 1.0, + view=SchedulerView(latent_width=64, latent_height=64), + ) + + flux_nodes = import_module("comfy_extras.nodes_flux") + expected = torch.as_tensor(flux_nodes.get_schedule(4, 4096), dtype=torch.float32) + assert torch.allclose(sigmas, expected, atol=1e-6, rtol=1e-6) + + +def test_flux2_schedule_requires_a_model_view() -> None: + """Flux2 fails clearly when a caller omits its resolution context.""" + + with pytest.raises(ValueError, match="requires a model view"): + calculate_sigmas(FakeModel(), "Flux2", "euler", 4, 1.0) + + +def test_flux2_schedule_uses_model_latent_downscale_ratio() -> None: + """Flux2 remains selectable for models with non-Flux latent formats.""" + + model = FakeModelWithLatentFormat(spacial_downscale_ratio=8) + sigmas = calculate_sigmas( + model, + "Flux2", + "euler", + 4, + 1.0, + view=SchedulerView(latent_width=128, latent_height=128), + ) + + flux_nodes = import_module("comfy_extras.nodes_flux") + expected = torch.as_tensor(flux_nodes.get_schedule(4, 4096), dtype=torch.float32) + assert torch.allclose(sigmas, expected, atol=1e-6, rtol=1e-6) + + def test_available_schedulers_deduplicates_beta57_when_globally_patched( monkeypatch: pytest.MonkeyPatch, ) -> None: diff --git a/tests/test_segs_tiled_diffusion.py b/tests/test_segs_tiled_diffusion.py index ca35e3f..d9c5081 100644 --- a/tests/test_segs_tiled_diffusion.py +++ b/tests/test_segs_tiled_diffusion.py @@ -11,8 +11,9 @@ import torch from simple_syrup.domain.segs import BoundingBox, CropRegion, Segment from simple_syrup.domain.segs_tiled_diffusion import ( - _segment_mask_to_latent, build_segs_guided_tiled_diffusion_plan, + segment_mask_to_latent, + segment_weight_to_latent, ) @@ -127,7 +128,7 @@ def test_crop_local_mask_projects_directly_to_latent_space() -> None: label="small_region", ) - latent_mask = _segment_mask_to_latent( + latent_mask = segment_mask_to_latent( segment, source_height=4096, source_width=4096, @@ -141,6 +142,30 @@ def test_crop_local_mask_projects_directly_to_latent_space() -> None: assert not bool(latent_mask[:, :64].any()) +def test_crop_local_weight_projection_preserves_soft_mask_values() -> None: + """Semantic consumers can retain fractional SAM write ownership.""" + + crop = CropRegion(0, 0, 8, 8) + segment = Segment( + cropped_image=None, + cropped_mask=torch.full((8, 8), 0.25, dtype=torch.float32), + confidence=1.0, + crop_region=crop, + bbox=BoundingBox(*crop), + label="soft_region", + ) + + latent_weight = segment_weight_to_latent( + segment, + source_height=8, + source_width=8, + latent_width=8, + latent_height=8, + ) + + assert torch.allclose(latent_weight, torch.full((8, 8), 0.25)) + + def test_guided_plan_rejects_mismatched_image_aspect_ratio() -> None: """SEGS from a different image fail before tiled sampling begins.""" diff --git a/tests/test_tiled_sampling_runtime.py b/tests/test_tiled_sampling_runtime.py index 36fbe78..dde1ef5 100644 --- a/tests/test_tiled_sampling_runtime.py +++ b/tests/test_tiled_sampling_runtime.py @@ -9,6 +9,7 @@ from __future__ import annotations import pytest import torch +from simple_syrup.domain.contextual_diffusion import SpatialContext from simple_syrup.domain.tiled_diffusion import LatentTile, build_tiled_diffusion_plan from simple_syrup.runtime import tiled_sampling @@ -170,6 +171,57 @@ def test_tile_transformer_options_repeats_model_metadata() -> None: assert torch.equal(tiled["sample_sigmas"], torch.tensor([1.0, 0.0])) +def test_spatial_context_args_resize_latent_and_canvas_conditioning() -> None: + """A global context resizes spatial tensors and repeats aligned metadata.""" + + x = torch.arange(2 * 1 * 8 * 12, dtype=torch.float32).reshape((2, 1, 8, 12)) + timestep = torch.tensor([0.5, 0.75]) + context = SpatialContext(0, 0, 12, 8, 6, 4) + + transformed = tiled_sampling.make_spatial_context_model_args( + args={ + "input": x, + "timestep": timestep, + "cond_or_uncond": [0, 1], + "c": { + "c_concat": x.clone(), + "c_crossattn": torch.ones((2, 3, 1)), + "transformer_options": {"cond_or_uncond": [0, 1]}, + }, + }, + contexts=(context,), + input_batch_size=2, + latent_height=8, + latent_width=12, + ) + + assert transformed["input"].shape == (2, 1, 4, 6) + assert transformed["c"]["c_concat"].shape == (2, 1, 4, 6) + assert transformed["c"]["c_crossattn"].shape == (2, 3, 1) + assert transformed["c"]["transformer_options"]["cond_or_uncond"] == [0, 1] + + +def test_spatial_context_args_batch_equal_shapes_for_5d_latents() -> None: + """Contexts batch across the spatial axes of singleton-depth latents.""" + + x = torch.zeros((1, 16, 1, 8, 16)) + contexts = ( + SpatialContext(0, 0, 8, 8, 8, 8), + SpatialContext(8, 0, 8, 8, 8, 8), + ) + + transformed = tiled_sampling.make_spatial_context_model_args( + args={"input": x, "timestep": torch.tensor([1.0]), "c": {}}, + contexts=contexts, + input_batch_size=1, + latent_height=8, + latent_width=16, + ) + + assert transformed["input"].shape == (2, 16, 1, 8, 8) + assert transformed["timestep"].shape == (2,) + + def test_new_spatial_weight_buffer_broadcasts_over_spatial_axes() -> None: """Spatial weight buffers broadcast over BCHW and BCDHW model outputs.""" @@ -209,6 +261,37 @@ def test_semantic_tile_weight_cache_reuses_resident_weights() -> None: assert accumulation_weight.dtype == torch.float32 +def test_tile_prediction_accumulator_selects_multidiffusion_or_mod_weights() -> None: + """One shared accumulator preserves the distinct overlap blending policies.""" + + plan = build_tiled_diffusion_plan(6, 4, 4, 4, 2, 2) + x = torch.arange(6, dtype=torch.float32).reshape(1, 1, 1, 6).expand(1, 1, 4, 6) + args = {"input": x, "timestep": torch.tensor([1.0]), "c": {}} + + def evaluate(tiled_args: dict[str, object]) -> torch.Tensor: + """Return a distinct constant prediction for each source tile.""" + + tiled_input = tiled_args["input"] + if not isinstance(tiled_input, torch.Tensor): + raise TypeError("Test evaluator expected a tiled input tensor.") + means = tiled_input.mean(dim=(-2, -1), keepdim=True) + return means.expand_as(tiled_input) + + multidiffusion = tiled_sampling.TilePredictionAccumulator( + plan, + diffusion_mode="multidiffusion", + ).predict(args=args, x=x, evaluate=evaluate) + mixture = tiled_sampling.TilePredictionAccumulator( + plan, + diffusion_mode="mixture_of_diffusers", + ).predict(args=args, x=x, evaluate=evaluate) + + assert torch.allclose(multidiffusion[:, :, :, 2:4], torch.full((1, 1, 4, 2), 2.5)) + assert not torch.allclose(mixture[:, :, :, 2:4], multidiffusion[:, :, :, 2:4]) + assert torch.allclose(mixture[:, :, :, :2], multidiffusion[:, :, :, :2]) + assert torch.allclose(mixture[:, :, :, 4:], multidiffusion[:, :, :, 4:]) + + def test_contains_unsupported_conditioning_key_finds_nested_values() -> None: """Unsupported regional and control keys are detected recursively."""