feat(sampling): add contextual diffusion sampler
This commit is contained in:
@@ -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
|
||||
@@ -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,)
|
||||
@@ -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": (
|
||||
|
||||
@@ -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 "
|
||||
|
||||
@@ -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")
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
@@ -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
|
||||
@@ -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(
|
||||
|
||||
@@ -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",)),
|
||||
(
|
||||
|
||||
@@ -33,6 +33,7 @@ BASE_NODE_IDS = [
|
||||
"SimpleSyrup.KSamplerExtras",
|
||||
"SimpleSyrup.KSamplerPromptByRegion",
|
||||
"SimpleSyrup.KSamplerPromptByTiledRegion",
|
||||
"SimpleSyrup.KSamplerContextualDiffusion",
|
||||
"SimpleSyrup.KSamplerTiledDiffusion",
|
||||
"SimpleSyrup.LatentDiagnostics",
|
||||
"SimpleSyrup.LayerStyleSAMModelsAdapter",
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user