fix(sampling): normalize model-specific latent layouts

This commit is contained in:
Artificial Sweetener
2026-08-30 21:00:58 -04:00
parent e8cbc6f724
commit c35191de31
5 changed files with 189 additions and 9 deletions
@@ -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()
@@ -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:
@@ -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,
@@ -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()
+43
View File
@@ -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)