Files
Artificial-Sweetener-Simple…/simple_syrup/services/contextual_diffusion_sampling_service.py
T

282 lines
9.9 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
"""Application service for contextual diffusion latent sampling."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, ClassVar, TypeAlias
import torch
from ..domain.conditioning_batch import ConditioningBatch, select_conditioning
from ..domain.context_segs import (
ContextSegs,
context_segs_from_tile_plan,
merge_context_segs,
)
from ..domain.contextual_diffusion import (
ContextualDiffusionControls,
build_contextual_diffusion_plan,
)
from ..domain.regional_features import (
CONTEXTUAL_DIFFUSION_REGIONAL_SAMPLER_CAPABILITIES,
EMPTY_REGIONAL_FEATURE_REQUEST,
RegionalCapabilityAdmission,
RegionalFeature,
RegionalFeatureRequest,
)
from ..domain.segs import NativeSegs, coerce_segs_group
from ..domain.tiled_diffusion import validate_tiled_diffusion_mode
from ..runtime.contextual_diffusion_sampling import sample_contextual_diffusion
from ..runtime.latent_geometry import decoded_image_dimensions
from .regional_capability_admission_service import (
RegionalCapabilityAdmissionService,
)
from .regional_sampling_preparation_service import (
RegionalSamplingPreparationService,
)
from .sampling_batch import (
combine_latent_outputs,
latent_batch_size,
single_item_latent,
)
Latent: TypeAlias = dict[str, Any]
@dataclass(frozen=True)
class ContextualDiffusionSamplingResult:
"""Return the sampled latent and lazy non-global context SEGS."""
latent: Latent
contexts: ContextSegs
class ContextualDiffusionSamplingService:
"""Plan and execute composition-preserving contextual diffusion."""
regional_preparation_service_class: ClassVar[
type[RegionalSamplingPreparationService]
] = RegionalSamplingPreparationService
capability_admission_service_class: ClassVar[
type[RegionalCapabilityAdmissionService]
] = RegionalCapabilityAdmissionService
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,
region_masks: object | None = None,
regional_prompt_weight: float = 0.5,
region_mask_feather: int = 0,
feature_request: RegionalFeatureRequest = EMPTY_REGIONAL_FEATURE_REQUEST,
planning_region_masks: torch.Tensor | None = None,
) -> ContextualDiffusionSamplingResult:
"""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()
regional = self.regional_preparation_service_class().prepare(
positive=positive,
negative=negative,
latent_image=latent_image,
region_masks=region_masks,
regional_prompt_weight=regional_prompt_weight,
region_mask_feather=region_mask_feather,
)
if planning_region_masks is not None and regional.mask_bank is not None:
raise ValueError(
"Contextual Diffusion cannot combine explicit planning masks with "
"legacy Regional Conditioning masks."
)
planning_masks = planning_region_masks
if planning_masks is None and regional.mask_bank is not None:
planning_masks = regional.mask_bank.planning_masks
effective_request = feature_request
if regional.active:
effective_request = effective_request.with_feature(
RegionalFeature.FULL_CONTEXT_MASKED_CONDITIONING
)
capability_admission = self.capability_admission_service_class().admit(
request=effective_request,
sampler_capabilities=CONTEXTUAL_DIFFUSION_REGIONAL_SAMPLER_CAPABILITIES,
model=model,
)
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}."
)
image_height, image_width = (
segs_group[0][0]
if segs_group
else decoded_image_dimensions(
model=model,
latent_image=latent_image,
latent_height=int(latent_image["samples"].shape[-2]),
latent_width=int(latent_image["samples"].shape[-1]),
)
)
split_batch = (
bool(segs_group)
or regional.active
or isinstance(regional.positive, ConditioningBatch)
or isinstance(regional.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=regional.positive,
negative=regional.negative,
latent_image=latent_image,
denoise=denoise,
diffusion_mode=diffusion_mode,
controls=controls,
segs=None,
image_height=image_height,
image_width=image_width,
region_masks=planning_masks,
capability_admission=capability_admission,
)
outputs: list[torch.Tensor] = []
contexts: list[ContextSegs] = []
for index in range(batch_size):
item_result = self._sample_item(
model=model,
seed=seed,
steps=steps,
cfg=cfg,
sampler_name=sampler_name,
scheduler=scheduler,
positive=(
select_conditioning(regional.positive, index)
if isinstance(regional.positive, ConditioningBatch)
else regional.positive
),
negative=(
select_conditioning(regional.negative, index)
if isinstance(regional.negative, ConditioningBatch)
else regional.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
),
image_height=image_height,
image_width=image_width,
region_masks=planning_masks,
capability_admission=capability_admission,
)
samples = item_result.latent.get("samples")
if not isinstance(samples, torch.Tensor):
raise TypeError(
"Contextual Diffusion output samples must be a torch.Tensor."
)
outputs.append(samples)
contexts.append(item_result.contexts)
return ContextualDiffusionSamplingResult(
latent=combine_latent_outputs(latent_image, outputs),
contexts=merge_context_segs(contexts),
)
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,
image_height: int,
image_width: int,
region_masks: torch.Tensor | None,
capability_admission: RegionalCapabilityAdmission,
) -> ContextualDiffusionSamplingResult:
"""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,
region_masks=region_masks,
)
latent = 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,
capability_admission=capability_admission,
)
return ContextualDiffusionSamplingResult(
latent=latent,
contexts=context_segs_from_tile_plan(
plan.tile_plan,
image_height=image_height,
image_width=image_width,
),
)