feat(sampling): add contextual diffusion sampler

This commit is contained in:
Artificial Sweetener
2026-08-02 02:43:26 -04:00
parent 36081d7771
commit b58539097a
31 changed files with 2420 additions and 209 deletions
+146
View File
@@ -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
+28 -8
View File
@@ -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
@@ -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,)
+1 -1
View File
@@ -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,)
@@ -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": (
+31
View File
@@ -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 "
+2
View File
@@ -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,
@@ -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",
@@ -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")
+8 -8
View File
@@ -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
)
@@ -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,
+14 -50
View File
@@ -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."""
@@ -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(
+87 -3
View File
@@ -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."""
+298 -1
View File
@@ -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],
@@ -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,
)
@@ -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,
+56
View File
@@ -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
@@ -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."""
+98
View File
@@ -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",
)
+289
View File
@@ -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
@@ -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,))
@@ -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
+8
View File
@@ -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"]
@@ -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(
+3
View File
@@ -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(
@@ -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",)),
(
+1
View File
@@ -33,6 +33,7 @@ BASE_NODE_IDS = [
"SimpleSyrup.KSamplerExtras",
"SimpleSyrup.KSamplerPromptByRegion",
"SimpleSyrup.KSamplerPromptByTiledRegion",
"SimpleSyrup.KSamplerContextualDiffusion",
"SimpleSyrup.KSamplerTiledDiffusion",
"SimpleSyrup.LatentDiagnostics",
"SimpleSyrup.LayerStyleSAMModelsAdapter",
+73 -1
View File
@@ -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:
+27 -2
View File
@@ -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."""
+83
View File
@@ -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."""