fix(sampling): normalize model-specific latent layouts
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user