237 lines
9.0 KiB
Python
237 lines
9.0 KiB
Python
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
|
# Copyright (C) 2026 Artificial Sweetener and contributors
|
|
# SPDX-License-Identifier: AGPL-3.0-or-later
|
|
|
|
"""Compile capability configuration into one authoritative sampling execution."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from typing import Any, TypedDict
|
|
|
|
import torch
|
|
|
|
from ..domain.attention_coupling_request import (
|
|
AttentionCouplingRequestMode,
|
|
classify_attention_coupling_request,
|
|
)
|
|
from ..domain.sampler_options import SamplerOptions, TilingOptions
|
|
from ..runtime import sampling_schedulers
|
|
from ..runtime.noise_inversion import validate_inversion_target
|
|
from ..runtime.tiled_sampling_validation import validate_sampling_controls
|
|
from ..shared.logging import get_logger
|
|
from .attention_coupling_sampling_service import AttentionCouplingSamplingService
|
|
from .contextual_attention_coupling_sampling_service import (
|
|
ContextualAttentionCouplingSamplingService,
|
|
)
|
|
from .contextual_diffusion_sampling_service import ContextualDiffusionSamplingService
|
|
from .ksampler_sampling_service import KSamplerSamplingService
|
|
from .tiled_attention_coupling_sampling_service import (
|
|
TiledAttentionCouplingSamplingService,
|
|
)
|
|
from .tiled_diffusion_sampling_service import TiledDiffusionSamplingService
|
|
|
|
LOGGER = get_logger(__name__)
|
|
|
|
|
|
class SamplingArguments(TypedDict):
|
|
"""Narrow controls while retaining dynamic host MODEL and tensor payloads."""
|
|
|
|
model: Any
|
|
seed: int
|
|
steps: int
|
|
cfg: float
|
|
sampler_name: str
|
|
scheduler: str
|
|
positive: object
|
|
negative: object
|
|
latent_image: dict[str, Any]
|
|
denoise: float
|
|
|
|
|
|
class ContextualArguments(TypedDict):
|
|
"""Describe global context and the sole local tile plan's execution controls."""
|
|
|
|
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
|
|
|
|
|
|
class TiledArguments(TypedDict):
|
|
"""Describe local geometry and mask-dependent denoising policy."""
|
|
|
|
diffusion_mode: str
|
|
latent_tile_width: int
|
|
latent_tile_height: int
|
|
latent_tile_overlap: int
|
|
latent_tile_batch_size: int
|
|
differential_diffusion: bool
|
|
|
|
|
|
class SamplerOptionsSamplingService:
|
|
"""Route configuration independently of connection order or node placement."""
|
|
|
|
def sample(
|
|
self,
|
|
*,
|
|
model: Any,
|
|
seed: int,
|
|
steps: int,
|
|
cfg: float,
|
|
sampler_name: str,
|
|
scheduler: str,
|
|
positive: object,
|
|
negative: object,
|
|
latent_image: dict[str, Any],
|
|
denoise: float,
|
|
options: SamplerOptions | None = None,
|
|
segs: object | None = None,
|
|
region_masks: object | None = None,
|
|
) -> dict[str, Any]:
|
|
"""Admit connected region data only through its enabled sampling capability."""
|
|
if options is not None and not isinstance(options, SamplerOptions):
|
|
raise TypeError(
|
|
"KSampler options must come from SimpleSyrup options nodes."
|
|
)
|
|
configured = options if options is not None else SamplerOptions()
|
|
context = configured.contextual_diffusion
|
|
tiling = configured.tiling
|
|
if context is not None:
|
|
if tiling is not None:
|
|
LOGGER.warning(
|
|
"Contextual Diffusion takes precedence; Tiling Options ignored.",
|
|
extra={"node_id": "SimpleSyrup.KSampler"},
|
|
)
|
|
tiling = context.local_tiling()
|
|
arguments: SamplingArguments = {
|
|
"model": model,
|
|
"seed": seed,
|
|
"steps": steps,
|
|
"cfg": cfg,
|
|
"sampler_name": sampler_name,
|
|
"scheduler": scheduler,
|
|
"positive": positive,
|
|
"negative": negative,
|
|
"latent_image": latent_image,
|
|
"denoise": denoise,
|
|
}
|
|
self._preflight(arguments, configured, tiling)
|
|
attention = configured.attention_coupling
|
|
if (
|
|
attention is not None
|
|
and classify_attention_coupling_request(
|
|
positive=positive, negative=negative, region_masks=region_masks
|
|
)
|
|
is AttentionCouplingRequestMode.BYPASS
|
|
):
|
|
attention = None
|
|
inversion = configured.noise_inversion
|
|
if context is not None:
|
|
assert tiling is not None
|
|
contextual_arguments: ContextualArguments = {
|
|
"diffusion_mode": tiling.diffusion_mode,
|
|
"latent_context_size": context.context_size,
|
|
"latent_context_overlap": tiling.overlap,
|
|
"latent_context_batch_size": tiling.batch_size,
|
|
"global_weight": context.global_weight,
|
|
"global_steps": context.global_steps,
|
|
"global_decay": context.global_decay,
|
|
}
|
|
if attention is not None:
|
|
result = ContextualAttentionCouplingSamplingService().sample(
|
|
**arguments,
|
|
**contextual_arguments,
|
|
region_masks=region_masks,
|
|
regional_prompt_weight=attention.regional_prompt_weight,
|
|
region_mask_feather=attention.region_mask_feather,
|
|
segs=segs,
|
|
tiling=tiling,
|
|
noise_inversion=inversion,
|
|
)
|
|
else:
|
|
result = ContextualDiffusionSamplingService().sample(
|
|
**arguments,
|
|
**contextual_arguments,
|
|
segs=segs,
|
|
tiling=tiling,
|
|
noise_inversion=inversion,
|
|
)
|
|
return result.latent
|
|
if tiling is not None:
|
|
tiled_arguments: TiledArguments = {
|
|
"diffusion_mode": tiling.diffusion_mode,
|
|
"latent_tile_width": tiling.width,
|
|
"latent_tile_height": tiling.height,
|
|
"latent_tile_overlap": tiling.overlap,
|
|
"latent_tile_batch_size": tiling.batch_size,
|
|
"differential_diffusion": tiling.differential_diffusion,
|
|
}
|
|
if attention is not None:
|
|
return TiledAttentionCouplingSamplingService().sample(
|
|
**arguments,
|
|
**tiled_arguments,
|
|
region_masks=region_masks,
|
|
regional_prompt_weight=attention.regional_prompt_weight,
|
|
region_mask_feather=attention.region_mask_feather,
|
|
segs=segs,
|
|
noise_inversion=inversion,
|
|
)
|
|
return TiledDiffusionSamplingService().sample(
|
|
**arguments,
|
|
**tiled_arguments,
|
|
segs=segs,
|
|
noise_inversion=inversion,
|
|
)
|
|
if attention is not None:
|
|
return AttentionCouplingSamplingService().sample(
|
|
**arguments,
|
|
region_masks=region_masks,
|
|
regional_prompt_weight=attention.regional_prompt_weight,
|
|
region_mask_feather=attention.region_mask_feather,
|
|
noise_inversion=inversion,
|
|
)
|
|
return KSamplerSamplingService().sample(**arguments, noise_inversion=inversion)
|
|
|
|
def _preflight(
|
|
self,
|
|
arguments: SamplingArguments,
|
|
options: SamplerOptions,
|
|
tiling: TilingOptions | None,
|
|
) -> None:
|
|
"""Reject unsupported schedules and endpoints before preparing models."""
|
|
samples = arguments["latent_image"].get("samples")
|
|
if not isinstance(samples, torch.Tensor):
|
|
raise TypeError("KSampler latent samples must be a torch.Tensor.")
|
|
validate_sampling_controls(
|
|
steps=arguments["steps"],
|
|
denoise=arguments["denoise"],
|
|
latent_tile_width=tiling.width if tiling is not None else 16,
|
|
latent_tile_height=tiling.height if tiling is not None else 16,
|
|
latent_tile_batch_size=tiling.batch_size if tiling is not None else 1,
|
|
)
|
|
incompatible_unipc = options.contextual_diffusion is not None or (
|
|
tiling is not None and tiling.diffusion_mode == "multidiffusion"
|
|
)
|
|
if incompatible_unipc and arguments["sampler_name"] in {"uni_pc", "uni_pc_bh2"}:
|
|
raise ValueError(
|
|
"Tiling and Contextual Diffusion do not support UniPC samplers."
|
|
)
|
|
if options.noise_inversion is not None:
|
|
view = (
|
|
sampling_schedulers.SchedulerView(tiling.width, tiling.height)
|
|
if tiling is not None
|
|
else sampling_schedulers.SchedulerView.from_tensor(samples)
|
|
)
|
|
sigmas = sampling_schedulers.calculate_sigmas(
|
|
model=arguments["model"],
|
|
scheduler_name=arguments["scheduler"],
|
|
sampler_name=arguments["sampler_name"],
|
|
steps=arguments["steps"],
|
|
denoise=arguments["denoise"],
|
|
view=view,
|
|
)
|
|
validate_inversion_target(arguments["model"], sigmas)
|