290 lines
8.5 KiB
Python
290 lines
8.5 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
|
|
|
|
"""Tests for contextual diffusion prediction fusion."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from typing import Any
|
|
|
|
import pytest
|
|
import torch
|
|
|
|
from simple_syrup.domain.contextual_diffusion import (
|
|
ContextualDiffusionControls,
|
|
build_contextual_diffusion_plan,
|
|
)
|
|
from simple_syrup.runtime import (
|
|
contextual_diffusion_sampling,
|
|
sampling_samplers,
|
|
sampling_schedulers,
|
|
)
|
|
from simple_syrup.runtime.contextual_diffusion_sampling import (
|
|
ContextualDiffusionModelWrapper,
|
|
)
|
|
|
|
comfy_sample = contextual_diffusion_sampling._comfy_sample()
|
|
comfy_utils = contextual_diffusion_sampling._comfy_utils()
|
|
latent_preview = contextual_diffusion_sampling._latent_preview()
|
|
|
|
|
|
def test_global_prediction_owns_low_frequency_while_tiles_keep_detail() -> None:
|
|
"""The native view replaces tile-level scene intent without blurring detail."""
|
|
|
|
controls = _controls()
|
|
plan = build_contextual_diffusion_plan(
|
|
latent_width=32,
|
|
latent_height=16,
|
|
controls=controls,
|
|
segs=None,
|
|
)
|
|
wrapper = ContextualDiffusionModelWrapper(
|
|
plan=plan,
|
|
controls=controls,
|
|
sigmas=torch.tensor([1.0, 0.0]),
|
|
existing_wrapper=None,
|
|
)
|
|
calls: list[tuple[int, int]] = []
|
|
|
|
def apply_model(
|
|
x: torch.Tensor,
|
|
timestep: torch.Tensor,
|
|
**conditioning: object,
|
|
) -> torch.Tensor:
|
|
"""Return high-frequency tile detail and a constant whole-image intent."""
|
|
|
|
del timestep, conditioning
|
|
calls.append((int(x.shape[-2]), int(x.shape[-1])))
|
|
if x.shape[-2:] == (8, 16):
|
|
return torch.full_like(x, 3.0)
|
|
rows = torch.arange(x.shape[-2], device=x.device).reshape(1, 1, -1, 1)
|
|
return torch.where(rows % 2 == 0, 1.0, -1.0).expand_as(x)
|
|
|
|
output = wrapper(
|
|
apply_model,
|
|
{
|
|
"input": torch.zeros((1, 1, 16, 32)),
|
|
"timestep": torch.tensor([1.0]),
|
|
"c": {},
|
|
},
|
|
)
|
|
|
|
assert calls == [(16, 16), (8, 16)]
|
|
assert torch.allclose(output[:, :, 0::2], torch.full((1, 1, 8, 32), 4.0))
|
|
assert torch.allclose(output[:, :, 1::2], torch.full((1, 1, 8, 32), 2.0))
|
|
|
|
|
|
def test_global_prediction_decays_then_stops_after_configured_steps() -> None:
|
|
"""Global authority decays before late denoising becomes local-only."""
|
|
|
|
controls = _controls(global_steps=2)
|
|
plan = build_contextual_diffusion_plan(
|
|
latent_width=32,
|
|
latent_height=16,
|
|
controls=controls,
|
|
segs=None,
|
|
)
|
|
wrapper = ContextualDiffusionModelWrapper(
|
|
plan=plan,
|
|
controls=controls,
|
|
sigmas=torch.tensor([1.0, 0.75, 0.5, 0.25, 0.0]),
|
|
existing_wrapper=None,
|
|
)
|
|
calls: list[tuple[int, int]] = []
|
|
|
|
def apply_model(
|
|
x: torch.Tensor,
|
|
timestep: torch.Tensor,
|
|
**conditioning: object,
|
|
) -> torch.Tensor:
|
|
"""Record whether each prediction uses local or global context."""
|
|
|
|
del timestep, conditioning
|
|
calls.append((int(x.shape[-2]), int(x.shape[-1])))
|
|
value = 3.0 if x.shape[-2:] == (8, 16) else 1.0
|
|
return torch.full_like(x, value)
|
|
|
|
base_args = {
|
|
"input": torch.zeros((1, 1, 16, 32)),
|
|
"c": {},
|
|
}
|
|
output = wrapper(apply_model, base_args | {"timestep": torch.tensor([0.75])})
|
|
assert calls == [(16, 16), (8, 16)]
|
|
assert torch.allclose(output, torch.full_like(output, 2.0))
|
|
|
|
calls.clear()
|
|
output = wrapper(apply_model, base_args | {"timestep": torch.tensor([0.5])})
|
|
assert calls == [(16, 16)]
|
|
assert torch.allclose(output, torch.ones_like(output))
|
|
|
|
|
|
def test_native_sized_canvas_delegates_to_one_original_model_call() -> None:
|
|
"""A canvas already inside the model view limit behaves like normal sampling."""
|
|
|
|
controls = _controls()
|
|
plan = build_contextual_diffusion_plan(
|
|
latent_width=16,
|
|
latent_height=16,
|
|
controls=controls,
|
|
segs=None,
|
|
)
|
|
wrapper = ContextualDiffusionModelWrapper(
|
|
plan=plan,
|
|
controls=controls,
|
|
sigmas=torch.tensor([1.0, 0.0]),
|
|
existing_wrapper=None,
|
|
)
|
|
calls = 0
|
|
|
|
def apply_model(
|
|
x: torch.Tensor,
|
|
timestep: torch.Tensor,
|
|
**conditioning: object,
|
|
) -> torch.Tensor:
|
|
"""Count direct model evaluations."""
|
|
|
|
del timestep, conditioning
|
|
nonlocal calls
|
|
calls += 1
|
|
return x + 2.0
|
|
|
|
x = torch.zeros((1, 1, 16, 16))
|
|
output = wrapper(
|
|
apply_model,
|
|
{"input": x, "timestep": torch.tensor([1.0]), "c": {}},
|
|
)
|
|
|
|
assert calls == 1
|
|
assert torch.equal(output, x + 2.0)
|
|
|
|
|
|
def test_runtime_delegates_sampling_to_comfy_with_wrapped_clone(
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
) -> None:
|
|
"""The vertical runtime path preserves KSampler sampling and latent metadata."""
|
|
|
|
model = _FakeModel()
|
|
sampler = object()
|
|
latent_samples = torch.zeros((1, 4, 16, 32))
|
|
sampled = torch.ones_like(latent_samples)
|
|
latent = {
|
|
"samples": latent_samples,
|
|
"downscale_ratio_spacial": 2,
|
|
"kept": "metadata",
|
|
}
|
|
controls = _controls()
|
|
plan = build_contextual_diffusion_plan(
|
|
latent_width=32,
|
|
latent_height=16,
|
|
controls=controls,
|
|
segs=None,
|
|
)
|
|
calls: dict[str, Any] = {}
|
|
|
|
monkeypatch.setattr(
|
|
sampling_samplers,
|
|
"resolve_sampler",
|
|
lambda _name: sampler,
|
|
)
|
|
monkeypatch.setattr(
|
|
sampling_schedulers,
|
|
"calculate_sigmas",
|
|
lambda **_kwargs: torch.tensor([1.0, 0.0]),
|
|
)
|
|
monkeypatch.setattr(
|
|
comfy_sample,
|
|
"fix_empty_latent_channels",
|
|
lambda _model, samples, _ratio: samples,
|
|
)
|
|
monkeypatch.setattr(
|
|
comfy_sample,
|
|
"prepare_noise",
|
|
lambda samples, _seed, _batch_inds=None: torch.ones_like(samples),
|
|
)
|
|
monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None)
|
|
monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", False)
|
|
|
|
def fake_sample_custom(
|
|
sampling_model: _FakeModel,
|
|
noise: torch.Tensor,
|
|
cfg: float,
|
|
received_sampler: object,
|
|
sigmas: torch.Tensor,
|
|
positive: object,
|
|
negative: object,
|
|
latent_image: torch.Tensor,
|
|
**kwargs: object,
|
|
) -> torch.Tensor:
|
|
"""Capture the final Comfy boundary call."""
|
|
|
|
del noise, cfg, sigmas, positive, negative, latent_image, kwargs
|
|
calls["model"] = sampling_model
|
|
calls["sampler"] = received_sampler
|
|
return sampled
|
|
|
|
monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom)
|
|
|
|
output = contextual_diffusion_sampling.sample_contextual_diffusion(
|
|
model=model,
|
|
seed=7,
|
|
steps=2,
|
|
cfg=1.0,
|
|
sampler_name="euler",
|
|
scheduler="simple",
|
|
positive=[],
|
|
negative=[],
|
|
latent_image=latent,
|
|
denoise=0.8,
|
|
diffusion_mode="mixture_of_diffusers",
|
|
controls=controls,
|
|
plan=plan,
|
|
)
|
|
|
|
assert calls["model"] is not model
|
|
assert isinstance(calls["model"].wrapper, ContextualDiffusionModelWrapper)
|
|
assert calls["model"].wrapper.diffusion_mode == "mixture_of_diffusers"
|
|
assert calls["sampler"] is sampler
|
|
assert output["samples"] is sampled
|
|
assert output["kept"] == "metadata"
|
|
assert "downscale_ratio_spacial" not in output
|
|
|
|
|
|
def _controls(
|
|
*,
|
|
global_steps: int = 1,
|
|
global_decay: float = 0.5,
|
|
) -> ContextualDiffusionControls:
|
|
"""Return a small two-tile test configuration."""
|
|
|
|
return ContextualDiffusionControls(
|
|
latent_context_size=16,
|
|
latent_context_overlap=0,
|
|
latent_context_batch_size=2,
|
|
global_weight=1.0,
|
|
global_steps=global_steps,
|
|
global_decay=global_decay,
|
|
)
|
|
|
|
|
|
class _FakeModel:
|
|
"""Provide the ModelPatcher surface used by the semantic runtime."""
|
|
|
|
def __init__(self, model_options: dict[str, Any] | None = None) -> None:
|
|
"""Create a CPU-backed fake model patcher."""
|
|
|
|
self.load_device = torch.device("cpu")
|
|
self.model_options = {} if model_options is None else model_options
|
|
self.wrapper: object | None = None
|
|
|
|
def clone(self) -> _FakeModel:
|
|
"""Return a clone with copied model options."""
|
|
|
|
return _FakeModel(self.model_options.copy())
|
|
|
|
def set_model_unet_function_wrapper(self, wrapper: object) -> None:
|
|
"""Capture the installed wrapper."""
|
|
|
|
self.wrapper = wrapper
|
|
self.model_options["model_function_wrapper"] = wrapper
|