Phase 3 of the modernization plan: move the framework-free math out of the comfy-coupled modules so Tier-0 tests can run on plain Linux without ComfyUI, coremltools, or python_coreml_stable_diffusion. New pure-core package (no comfy / coreml / mps imports): - coreml_suite.core.latents: chunk_batch, merge_chunks - coreml_suite.core.controlnet: expand_inputs, no_control, extract_residual_kwargs, chunk_control - coreml_suite.core.inputs: CoreMLInputs (chunks + coreml_kwargs) - coreml_suite.core.sdxl: is_sdxl / is_sdxl_base / is_sdxl_refiner, build_sdxl_time_ids (base len 6, refiner len 5), build_sdxl_text_embeds, sdxl_model_function_wrapper - coreml_suite.core.naming: compose_out_name, lora_names_from_params Thin adapters keep the public import paths: - coreml_suite.latents / coreml_suite.controlnet: re-export from core - coreml_suite.models: CoreMLModelWrapper, CoreMLModelWrapperLCM, add_sdxl_model_options (now uses the pure builders from core.sdxl), get_latent_image, get_model_patcher remain framework-coupled - coreml_suite.nodes: CoreMLConverter.convert now delegates the out_name composition to core.naming.compose_out_name Test infra: - tests/unit/* re-pointed at coreml_suite.core.* - test_chunks.py dropped `from comfy.model_management import ...` and the dead `model_config` fixture (Phase 1 left it broken; Phase 3 removes it entirely) - test_characterization_sdxl_options now targets the pure builders directly via inspect.getclosurevars on the wrapper closure - test_characterization_out_name now calls compose_out_name without the heavy CoreMLConverter monkey-patching that Phase 2 needed - tests/unit/test_tier0_purity.py: new gate that fails if comfy / coremltools / etc leak into sys.modules during a pure `-m unit` run (skipped in mixed runs where m2 / integration legitimately import them) - tests/__init__.py + top-level conftest.py + pyproject addopts `--import-mode=importlib --confcutdir=tests` together stop pytest from importing the repo-root `__init__.py` (the ComfyUI custom-node entry pulls in comfy) - tests/conftest.py adds tier-aware collect_ignore so `-m unit` skips tests/m2 + tests/integration at collection time Verification: - `pytest -m unit tests/` → 88 passed in ~2s; deterministic across runs - Tier-0 purity gate confirms no comfy/coreml/etc in sys.modules - m2 golden image (Phase 2 anchor) still hashes identical → refactor produced bit-for-bit unchanged output - `git diff main -- __init__.py coreml_suite/nodes.py` shows zero churn to NODE_CLASS_MAPPINGS keys or INPUT_TYPES field names (public workflow contract intact)
128 lines
4.4 KiB
Python
128 lines
4.4 KiB
Python
"""Phase 2 characterization tests, Phase 3 re-pointed.
|
|
|
|
After Phase 3 the SDXL time_ids / text_embeds math lives in
|
|
coreml_suite.core.sdxl as pure builders. The framework adapter
|
|
add_sdxl_model_options (in models.py) is exercised separately by the m2
|
|
golden image test; here we just lock the pure math.
|
|
"""
|
|
import inspect
|
|
|
|
import pytest
|
|
import torch
|
|
|
|
from coreml_suite.core.sdxl import (
|
|
build_sdxl_text_embeds,
|
|
build_sdxl_time_ids,
|
|
sdxl_model_function_wrapper,
|
|
)
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _deterministic_seed():
|
|
torch.manual_seed(0)
|
|
|
|
|
|
# ---------- build_sdxl_time_ids: base (len 6) -------------------------------
|
|
|
|
|
|
def test_build_time_ids_base_defaults():
|
|
out = build_sdxl_time_ids({}, {}, is_base=True, is_refiner=False)
|
|
expected = torch.tensor([[768, 768, 0, 0, 768, 768], [768, 768, 0, 0, 768, 768]])
|
|
assert out.shape == (2, 6)
|
|
assert torch.equal(out, expected)
|
|
|
|
|
|
def test_build_time_ids_base_respects_overrides():
|
|
pos = {"height": 1024, "width": 512, "crop_h": 8, "crop_w": 4,
|
|
"target_height": 1024, "target_width": 1024}
|
|
neg = {"height": 256, "width": 256, "crop_h": 0, "crop_w": 0,
|
|
"target_height": 256, "target_width": 256}
|
|
out = build_sdxl_time_ids(pos, neg, is_base=True, is_refiner=False)
|
|
expected = torch.tensor([[1024, 512, 8, 4, 1024, 1024], [256, 256, 0, 0, 256, 256]])
|
|
assert torch.equal(out, expected)
|
|
|
|
|
|
# ---------- build_sdxl_time_ids: refiner (len 5) ----------------------------
|
|
|
|
|
|
def test_build_time_ids_refiner_defaults():
|
|
out = build_sdxl_time_ids({}, {}, is_base=False, is_refiner=True)
|
|
expected = torch.tensor([[768, 768, 0, 0, 6.0], [768, 768, 0, 0, 2.5]])
|
|
assert out.shape == (2, 5)
|
|
assert torch.equal(out, expected)
|
|
|
|
|
|
def test_build_time_ids_refiner_respects_aesthetic_score():
|
|
pos = {"aesthetic_score": 8.5}
|
|
neg = {"aesthetic_score": 1.5}
|
|
out = build_sdxl_time_ids(pos, neg, is_base=False, is_refiner=True)
|
|
expected = torch.tensor([[768, 768, 0, 0, 8.5], [768, 768, 0, 0, 1.5]])
|
|
assert torch.equal(out, expected)
|
|
|
|
|
|
# ---------- build_sdxl_time_ids: edge case ----------------------------------
|
|
|
|
|
|
def test_build_time_ids_neither_base_nor_refiner_returns_len4():
|
|
out = build_sdxl_time_ids({}, {}, is_base=False, is_refiner=False)
|
|
assert out.shape == (2, 4)
|
|
|
|
|
|
# ---------- build_sdxl_text_embeds ------------------------------------------
|
|
|
|
|
|
def test_text_embeds_concat_pos_then_neg():
|
|
pos = torch.full((1, 1280), 1.0)
|
|
neg = torch.full((1, 1280), -1.0)
|
|
out = build_sdxl_text_embeds(pos, neg)
|
|
assert out.shape == (2, 1280)
|
|
assert torch.equal(out[0], pos[0])
|
|
assert torch.equal(out[1], neg[0])
|
|
|
|
|
|
# ---------- sdxl_model_function_wrapper closure -----------------------------
|
|
|
|
|
|
def test_wrapper_captures_time_ids_text_embeds_refiner_via_closure():
|
|
time_ids = torch.zeros(2, 6)
|
|
text_embeds = torch.zeros(2, 1280)
|
|
wrapper = sdxl_model_function_wrapper(time_ids, text_embeds, refiner=False)
|
|
closure = inspect.getclosurevars(wrapper).nonlocals
|
|
assert closure["time_ids"] is time_ids
|
|
assert closure["text_embeds"] is text_embeds
|
|
assert closure["refiner"] is False
|
|
|
|
|
|
def test_wrapper_returns_zero_when_context_missing():
|
|
"""When c_crossattn is None the wrapper short-circuits to zeros_like(x).
|
|
Locked here because Phase 3 mustn't change this default."""
|
|
wrapper = sdxl_model_function_wrapper(torch.zeros(2, 6), torch.zeros(2, 1280))
|
|
x = torch.randn(2, 4, 16, 16)
|
|
out = wrapper(
|
|
model_function=lambda *a, **kw: pytest.fail("model_function must not run"),
|
|
params={"input": x, "timestep": torch.zeros(2), "c": {}},
|
|
)
|
|
assert torch.equal(out, torch.zeros_like(x))
|
|
|
|
|
|
def test_wrapper_refiner_truncates_context_to_g_clip():
|
|
"""refiner=True slices c_crossattn[:, :, 768:] before forwarding."""
|
|
captured = {}
|
|
|
|
def fake_model(x, t, **c):
|
|
captured["context_shape"] = c["c_crossattn"].shape
|
|
captured["time_ids_shape"] = c["time_ids"].shape
|
|
return x
|
|
|
|
wrapper = sdxl_model_function_wrapper(
|
|
torch.zeros(2, 5), torch.zeros(2, 1280), refiner=True
|
|
)
|
|
x = torch.randn(2, 4, 16, 16)
|
|
context = torch.randn(2, 77, 2048) # 768 + 1280 dims
|
|
wrapper(
|
|
model_function=fake_model,
|
|
params={"input": x, "timestep": torch.zeros(2), "c": {"c_crossattn": context}},
|
|
)
|
|
assert captured["context_shape"] == (2, 77, 1280)
|
|
assert captured["time_ids_shape"] == (2, 5)
|