From c35191de31291520fe4067b24158e34ff75caa0c Mon Sep 17 00:00:00 2001 From: Artificial Sweetener Date: Sun, 30 Aug 2026 21:00:58 -0400 Subject: [PATCH] fix(sampling): normalize model-specific latent layouts --- .../runtime/comfy_latent_normalization.py | 42 ++++++++++ ...tion_coupling_model_preparation_service.py | 16 ++++ .../services/ksampler_sampling_service.py | 20 +++-- ...tion_coupling_model_preparation_service.py | 77 ++++++++++++++++++- tests/test_comfy_latent_normalization.py | 43 +++++++++++ 5 files changed, 189 insertions(+), 9 deletions(-) create mode 100644 simple_syrup/runtime/comfy_latent_normalization.py create mode 100644 tests/test_comfy_latent_normalization.py diff --git a/simple_syrup/runtime/comfy_latent_normalization.py b/simple_syrup/runtime/comfy_latent_normalization.py new file mode 100644 index 0000000..286d924 --- /dev/null +++ b/simple_syrup/runtime/comfy_latent_normalization.py @@ -0,0 +1,42 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Normalize latent tensors through ComfyUI's model-aware host contract.""" + +from __future__ import annotations + +from importlib import import_module + +import torch + + +class ComfyLatentNormalizer: + """Apply ComfyUI's channel, spatial, and temporal latent normalization.""" + + def normalize( + self, + *, + model: object, + samples: torch.Tensor, + spatial_downscale_ratio: object | None = None, + temporal_downscale_ratio: object | None = None, + ) -> torch.Tensor: + """Return samples in the latent layout required by the supplied model.""" + + normalize = import_module("comfy.sample").fix_empty_latent_channels + if temporal_downscale_ratio is None: + normalized = normalize(model, samples, spatial_downscale_ratio) + else: + normalized = normalize( + model, + samples, + spatial_downscale_ratio, + temporal_downscale_ratio, + ) + if not isinstance(normalized, torch.Tensor): + raise TypeError("ComfyUI latent normalization must return a torch.Tensor.") + return normalized + + +COMFY_LATENT_NORMALIZER = ComfyLatentNormalizer() diff --git a/simple_syrup/services/attention_coupling_model_preparation_service.py b/simple_syrup/services/attention_coupling_model_preparation_service.py index 321820e..be7d5a4 100644 --- a/simple_syrup/services/attention_coupling_model_preparation_service.py +++ b/simple_syrup/services/attention_coupling_model_preparation_service.py @@ -25,6 +25,7 @@ from ..runtime.comfy_conditioning_model_loader import ComfyConditioningModelLoad from ..runtime.comfy_conditioning_processing import ( ComfyRegionalConditioningProcessor, ) +from ..runtime.comfy_latent_normalization import ComfyLatentNormalizer from ..runtime.regional_lora_conditioning_adapter import ( RegionalLoraConditioningAdapter, ) @@ -78,6 +79,9 @@ class AttentionCouplingModelPreparationService: interop_validator_class: ClassVar[type[RegionalModelPatchInteropValidator]] = ( RegionalModelPatchInteropValidator ) + latent_normalizer_class: ClassVar[type[ComfyLatentNormalizer]] = ( + ComfyLatentNormalizer + ) model_family_selector_class: ClassVar[ type[AttentionCouplingModelFamilySelector] ] = AttentionCouplingModelFamilySelector @@ -113,6 +117,18 @@ class AttentionCouplingModelPreparationService: interop_validator = self.interop_validator_class() interop_report = interop_validator.validate(model, capabilities) model_family = self.model_family_selector_class().select(capabilities) + samples = self.latent_normalizer_class().normalize( + model=model, + samples=samples, + spatial_downscale_ratio=latent_image.get( + "downscale_ratio_spacial", + None, + ), + temporal_downscale_ratio=latent_image.get( + "downscale_ratio_temporal", + None, + ), + ) model_family.validate_latent(samples) def prepare_uncached() -> PreparedAttentionCouplingModel: diff --git a/simple_syrup/services/ksampler_sampling_service.py b/simple_syrup/services/ksampler_sampling_service.py index 86066b3..e861ae2 100644 --- a/simple_syrup/services/ksampler_sampling_service.py +++ b/simple_syrup/services/ksampler_sampling_service.py @@ -13,6 +13,7 @@ import torch from ..domain.conditioning_batch import ConditioningBatch, select_conditioning from ..runtime import sampling_samplers, sampling_schedulers +from ..runtime.comfy_latent_normalization import COMFY_LATENT_NORMALIZER from ..shared.logging import get_logger Latent: TypeAlias = dict[str, Any] @@ -45,15 +46,18 @@ class KSamplerSamplingService: comfy_sample = import_module("comfy.sample") comfy_utils = import_module("comfy.utils") - latent_samples = comfy_sample.fix_empty_latent_channels( - model, - latent_samples, - latent_image.get("downscale_ratio_spacial", None), + latent_samples = COMFY_LATENT_NORMALIZER.normalize( + model=model, + samples=latent_samples, + spatial_downscale_ratio=latent_image.get( + "downscale_ratio_spacial", + None, + ), + temporal_downscale_ratio=latent_image.get( + "downscale_ratio_temporal", + None, + ), ) - if not isinstance(latent_samples, torch.Tensor): - raise TypeError( - "KSampler normalized latent samples must be a torch.Tensor." - ) sigmas = sampling_schedulers.calculate_sigmas( model=model, scheduler_name=scheduler, diff --git a/tests/test_attention_coupling_model_preparation_service.py b/tests/test_attention_coupling_model_preparation_service.py index f797aef..3c3cbbc 100644 --- a/tests/test_attention_coupling_model_preparation_service.py +++ b/tests/test_attention_coupling_model_preparation_service.py @@ -7,7 +7,7 @@ from __future__ import annotations from types import SimpleNamespace -from typing import Any, ClassVar +from typing import Any, ClassVar, cast from uuid import uuid4 import pytest @@ -212,6 +212,33 @@ class _ModelFamilySelector: return _ModelFamily() +class _IdentityLatentNormalizer: + """Preserve synthetic test latents outside normalization-specific coverage.""" + + def normalize(self, **kwargs: object) -> torch.Tensor: + """Return the supplied tensor unchanged.""" + + samples = kwargs["samples"] + if not isinstance(samples, torch.Tensor): + raise TypeError("test latent samples must be a tensor") + return samples + + +class _AnimaLatentNormalizer: + """Adapt an ordinary image latent to the installed Anima layout.""" + + calls: ClassVar[list[dict[str, object]]] = [] + + def normalize(self, **kwargs: object) -> torch.Tensor: + """Record the request and add Anima's singleton temporal axis.""" + + type(self).calls.append(kwargs) + samples = kwargs["samples"] + if not isinstance(samples, torch.Tensor): + raise TypeError("test latent samples must be a tensor") + return samples.unsqueeze(2) + + def test_model_preparation_runs_each_shared_owner_once() -> None: """Prepare masks, contexts, model residency, and one selected backend once.""" @@ -260,6 +287,50 @@ def test_model_preparation_runs_each_shared_owner_once() -> None: assert _ModelFamily.derive_calls[0]["interop_report"] is _InteropValidator.report +def test_normalizes_standard_image_latent_before_anima_family_validation() -> None: + """Preparation presents BCHW empty latents to Anima as BC1HW tensors.""" + + originals = _install_fakes() + original_normalizer = ( + AttentionCouplingModelPreparationService.latent_normalizer_class + ) + AttentionCouplingModelPreparationService.latent_normalizer_class = cast( + Any, + _AnimaLatentNormalizer, + ) + samples = torch.zeros((1, 4, 8, 12)) + _reset_calls() + try: + AttentionCouplingModelPreparationService().prepare( + model=SimpleNamespace(model=object(), load_device="cpu"), + positive=ConditioningBatch((_conditioning(1.0), _conditioning(2.0))), + negative=ConditioningBatch((_conditioning(-1.0), _conditioning(-2.0))), + region_masks=torch.ones((1, 8, 12)), + regional_prompt_weight=1.0, + region_mask_feather=0, + latent_image={ + "samples": samples, + "downscale_ratio_spacial": 8, + }, + execution_mode=RegionalAttentionExecutionMode.FULL, + ) + finally: + AttentionCouplingModelPreparationService.latent_normalizer_class = ( + original_normalizer + ) + _restore_fakes(originals) + + assert _AnimaLatentNormalizer.calls == [ + { + "model": _AnimaLatentNormalizer.calls[0]["model"], + "samples": samples, + "spatial_downscale_ratio": 8, + "temporal_downscale_ratio": None, + } + ] + assert _ModelFamily.latent_calls[0].shape == (1, 4, 1, 8, 12) + + def test_source_model_is_restored_before_adapter_graph_discovery() -> None: """Load the source patcher before inspecting adapter target ownership.""" @@ -389,6 +460,7 @@ def _install_fakes() -> tuple[type[Any], ...]: service.conditioning_processor_class, service.model_family_selector_class, service.interop_validator_class, + service.latent_normalizer_class, ) service.capability_service_class = _CapabilityService # type: ignore[assignment] service.lora_adapter_class = _LoraAdapter # type: ignore[assignment] @@ -397,6 +469,7 @@ def _install_fakes() -> tuple[type[Any], ...]: service.conditioning_processor_class = _ConditioningProcessor # type: ignore[assignment] service.model_family_selector_class = _ModelFamilySelector # type: ignore[assignment] service.interop_validator_class = _InteropValidator # type: ignore[assignment] + service.latent_normalizer_class = _IdentityLatentNormalizer # type: ignore[assignment] return originals @@ -411,6 +484,7 @@ def _restore_fakes(originals: tuple[type[Any], ...]) -> None: service.conditioning_processor_class = originals[4] service.model_family_selector_class = originals[5] service.interop_validator_class = originals[6] + service.latent_normalizer_class = originals[7] def _reset_calls() -> None: @@ -429,4 +503,5 @@ def _reset_calls() -> None: _ModelFamily.derive_calls = [] _ModelFamily.latent_error = None _ModelFamily.adaptation_error = None + _AnimaLatentNormalizer.calls = [] _PREPARATION_EVENTS.clear() diff --git a/tests/test_comfy_latent_normalization.py b/tests/test_comfy_latent_normalization.py new file mode 100644 index 0000000..f6c30db --- /dev/null +++ b/tests/test_comfy_latent_normalization.py @@ -0,0 +1,43 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Verify model-aware latent normalization through the installed ComfyUI host.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import torch + +from simple_syrup.runtime.comfy_latent_normalization import ComfyLatentNormalizer + + +class _AnimaModel: + """Expose the latent-format values used by an Anima model patcher.""" + + def get_model_object(self, name: str) -> object: + """Return a three-dimensional sixteen-channel latent format.""" + + if name != "latent_format": + raise KeyError(name) + return SimpleNamespace( + latent_channels=16, + latent_dimensions=3, + spacial_downscale_ratio=8, + temporal_downscale_ratio=4, + ) + + +def test_normalizes_ordinary_empty_image_latent_for_anima() -> None: + """A standard BCHW empty latent becomes Anima's BC1HW latent.""" + + samples = torch.zeros((1, 4, 8, 12)) + + normalized = ComfyLatentNormalizer().normalize( + model=_AnimaModel(), + samples=samples, + spatial_downscale_ratio=8, + ) + + assert normalized.shape == (1, 16, 1, 8, 12)