test(phase2): add characterization tests + M2 golden image anchor

Phase 2 of the modernization plan: lock the current behavior of the pure
math so the Phase 3 refactor cannot silently change it.

Unit characterization tests (Tier 0, 64 new):
- test_characterization_latents.py: chunk_batch / merge_chunks padding,
  truncation, and identity contracts.
- test_characterization_controlnet.py: expand_inputs / no_control /
  extract_residual_kwargs / chunk_control shape, dtype, and zero-fill
  behavior, including the [None]*target contract.
- test_characterization_inputs.py: CoreMLInputs.chunks / coreml_kwargs
  for SD1.5, SDXL base (time_ids len 6), SDXL refiner (time_ids len 5),
  and LCM (timestep_cond).
- test_characterization_sdxl_options.py: add_sdxl_model_options time_ids
  / text_embeds assembly via a SimpleNamespace fake ModelPatcher and
  inspect.getclosurevars on the returned model_function_wrapper.
- test_characterization_out_name.py: CoreMLConverter out_name encoding
  for attn_impl suffix, batch/size, ControlNet, LoRA (sorted), SDXL.

M2 [Tier 2] golden image anchor (1 new):
- test_golden_image.py: posts the SD1.5+CoreML workflow to a local
  ComfyUI server (auto-skips if unreachable), asserts SHA256 of the
  generated PNG against tests/m2/goldens/sd15_seed42.sha256; falls back
  to PSNR >= 40 dB if the hash drifts.

Test infra:
- pyproject.toml [tool.pytest.ini_options]: unit / m2 / smoke markers,
  testpaths=tests; rootdir is now this package (was ComfyUI's pytest.ini).
- tests/conftest.py: bootstraps sys.path for comfy imports, auto-marks
  tests by directory, and ignores the maintainer's WIP scaffolds
  (test_experiments / test_unet_conversion / standalone_test) so they
  don't break collection.

All 85 collected tests pass; two consecutive runs produced identical
results (run1: 3.55s, run2: 3.40s).
This commit is contained in:
aszc-dev
2026-05-22 15:38:09 +02:00
parent cf6d7c6855
commit 04911d0052
12 changed files with 1084 additions and 0 deletions
+1
View File
@@ -9,3 +9,4 @@ coremlsuite-venv/
experiment_results/
test_results/
bench/scripts/*.log
tests/m2/_latest_generated.png
+16
View File
@@ -37,3 +37,19 @@ dev = [
"pillow>=12.2.0",
"psutil>=7.2.2",
]
[tool.pytest.ini_options]
# Phase 2: tier markers gate which environment a test needs.
# - unit: framework-free pure-logic tests (Tier 0; future Phase 3 will run
# these without ComfyUI on Linux). For now they may still transitively
# import comfy.
# - m2: needs an Apple Silicon Mac with the Neural Engine (Tier 2),
# typically the maintainer's self-hosted runner or local M-series box.
# - smoke: lightweight checks that need Apple Silicon + coremltools but no
# ANE/real model (reserved for Phase 4 Tier 1 smoke harness).
markers = [
"unit: framework-free unit test (Tier 0)",
"m2: requires Apple Silicon + Neural Engine (Tier 2)",
"smoke: macOS-ARM smoke test on a synthetic micro-model (Tier 1)",
]
testpaths = ["tests"]
+50
View File
@@ -0,0 +1,50 @@
"""Pytest bootstrap for ComfyUI-CoreMLSuite tests.
- Adds the ComfyUI checkout to sys.path so production modules that
transitively import `comfy.*` resolve when pytest is invoked from this
package's root. Phase 3 will split pure logic into a comfy-free core and
this hack can go.
- Auto-applies tier markers based on the directory a test lives in, so
individual files don't have to repeat @pytest.mark.unit / .m2.
- Skips the maintainer's in-progress test scaffolds so they don't break
collection (they reference modules / venv layouts that aren't part of
Phase 2 scope).
"""
import sys
from pathlib import Path
import pytest
REPO_ROOT = Path(__file__).resolve().parents[1]
COMFY_DIR = REPO_ROOT.parents[1]
for p in (str(COMFY_DIR), str(REPO_ROOT)):
if p not in sys.path:
sys.path.insert(0, p)
# Skip WIP test scaffolds left in the tree by the maintainer; they import
# modules (coreml_suite.experiments, convert_apple) that are not part of
# Phase 2 scope.
collect_ignore_glob = [
"unit/test_experiments.py",
"unit/test_unet_conversion.py",
"unit/standalone_test.py",
]
_TIER_BY_DIR = {
"tests/unit": "unit",
"tests/m2": "m2",
"tests/integration": "m2",
"tests/smoke": "smoke",
}
def pytest_collection_modifyitems(config, items):
for item in items:
path = str(item.fspath).replace("\\", "/")
for fragment, marker in _TIER_BY_DIR.items():
if f"/{fragment}/" in path:
item.add_marker(getattr(pytest.mark, marker))
break
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 451 KiB

+1
View File
@@ -0,0 +1 @@
9e12ba1f98969a99c17e84a87a77c4561fcac3ca2ddd6338a7ae89d324df6994
+162
View File
@@ -0,0 +1,162 @@
"""Phase 2 [M2-ANE] golden-image anchor.
Runs the e2e SD1.5 + CoreML workflow against a local ComfyUI server, fetches
the generated PNG, and asserts both:
- byte-identical SHA256 against the stored golden, OR
- PSNR >= GOLDEN_PSNR_MIN_DB against the stored golden PNG.
The hash is the strict gate (a Phase 3 refactor that doesn't touch the math
should hit it). PSNR is the soft gate that tolerates a sub-bit drift in
sampler ordering or kernel selection — anything below the threshold is
treated as a regression.
Skips entirely on non-Apple-Silicon hosts or when the server / converted
model is missing, so the unit lane on Linux still passes.
The first run with no golden writes one and fails so it's reviewed before
being committed.
"""
import hashlib
import json
import os
import platform
import shutil
import time
import urllib.error
import urllib.request
from pathlib import Path
import numpy as np
import pytest
from PIL import Image
REPO_ROOT = Path(__file__).resolve().parents[2]
COMFY_DIR = Path(os.environ.get("COMFY_DIR", REPO_ROOT.parents[1])).resolve()
COMFY_HOST = os.environ.get("COMFY_HOST", "localhost")
COMFY_PORT = int(os.environ.get("COMFY_PORT", "8188"))
COMFY_URL = f"http://{COMFY_HOST}:{COMFY_PORT}"
CKPT_NAME = os.environ.get("CKPT_NAME", "v1-5-pruned-emaonly.safetensors")
WORKFLOW_PATH = (
REPO_ROOT / "tests" / "integration" / "workflows" / "e2e-1.5-basic-conversion.json"
)
GOLDEN_DIR = Path(__file__).parent / "goldens"
GOLDEN_HASH_PATH = GOLDEN_DIR / "sd15_seed42.sha256"
GOLDEN_PNG_PATH = GOLDEN_DIR / "sd15_seed42.png"
GOLDEN_PSNR_MIN_DB = float(os.environ.get("GOLDEN_PSNR_MIN_DB", "40"))
SEED = 42
def _server_reachable() -> bool:
try:
with urllib.request.urlopen(f"{COMFY_URL}/prompt", timeout=3) as r:
return r.status == 200
except (urllib.error.URLError, urllib.error.HTTPError, ConnectionError):
return False
@pytest.fixture(scope="module")
def comfy_server():
if platform.machine() != "arm64":
pytest.skip("requires Apple Silicon")
if not _server_reachable():
pytest.skip(f"ComfyUI server not reachable at {COMFY_URL}")
return COMFY_URL
def _http_post_json(path: str, payload: dict) -> dict:
data = json.dumps(payload).encode("utf-8")
req = urllib.request.Request(
f"{COMFY_URL}{path}", data=data,
headers={"Content-Type": "application/json"}, method="POST",
)
with urllib.request.urlopen(req, timeout=300) as r:
return json.loads(r.read().decode())
def _http_get_json(path: str, timeout: int = 300) -> dict:
"""ComfyUI runs UNet inference on its single asyncio loop, so GET /prompt
blocks while the queued prompt is executing. Use a generous timeout."""
with urllib.request.urlopen(f"{COMFY_URL}{path}", timeout=timeout) as r:
return json.loads(r.read().decode())
def _drain_queue(timeout_s: int = 600) -> None:
deadline = time.time() + timeout_s
while time.time() < deadline:
try:
q = _http_get_json("/prompt")
except (urllib.error.URLError, TimeoutError):
# Transient block while server executes; retry until our overall
# deadline expires.
continue
if q.get("exec_info", {}).get("queue_remaining", -1) == 0:
return
time.sleep(2)
raise TimeoutError(f"queue did not drain within {timeout_s}s")
def _post_workflow_and_collect_png() -> bytes:
workflow = json.loads(WORKFLOW_PATH.read_text())
for nid in ("4", "10"):
if nid in workflow:
workflow[nid]["inputs"]["ckpt_name"] = CKPT_NAME
for nid in ("3", "11"):
if nid in workflow and "seed" in workflow[nid].get("inputs", {}):
workflow[nid]["inputs"]["seed"] = SEED
# Drop the MPS reference branch — torch 2.0.1 MPS path is broken on macOS 26
# (see Phase 1 gate report). Only the Core ML pipeline is needed here.
for nid in ("3", "8", "9"):
workflow.pop(nid, None)
_http_post_json("/prompt", {"prompt": workflow})
_drain_queue()
comfy_out = COMFY_DIR / "output"
matches = sorted(comfy_out.glob("E2E-1.5-CoreML_*.png"), reverse=True)
if not matches:
raise FileNotFoundError(f"no Core ML image under {comfy_out}")
return matches[0].read_bytes()
def _psnr(a: np.ndarray, b: np.ndarray) -> float:
mse = float(np.mean((a.astype(np.float64) - b.astype(np.float64)) ** 2))
if mse == 0:
return 100.0
return 20.0 * float(np.log10(255.0 / np.sqrt(mse)))
def test_sd15_seed42_image_matches_golden(comfy_server):
GOLDEN_DIR.mkdir(parents=True, exist_ok=True)
png_bytes = _post_workflow_and_collect_png()
sha = hashlib.sha256(png_bytes).hexdigest()
if not GOLDEN_HASH_PATH.exists() or not GOLDEN_PNG_PATH.exists():
GOLDEN_HASH_PATH.write_text(sha + "\n")
# Persist the PNG too for visual diffing + PSNR.
tmp_path = Path(__file__).parent / "_latest_generated.png"
tmp_path.write_bytes(png_bytes)
shutil.copy2(tmp_path, GOLDEN_PNG_PATH)
pytest.fail(
f"No golden present; wrote {GOLDEN_HASH_PATH.name} and "
f"{GOLDEN_PNG_PATH.name}. Review the image and re-run."
)
expected_hash = GOLDEN_HASH_PATH.read_text().strip()
if sha == expected_hash:
return
# Hash drift: fall back to PSNR to distinguish a refactor-safe rounding
# change from a real regression.
a = np.array(Image.open(GOLDEN_PNG_PATH).convert("RGB"))
b_path = Path(__file__).parent / "_latest_generated.png"
b_path.write_bytes(png_bytes)
b = np.array(Image.open(b_path).convert("RGB"))
if a.shape != b.shape:
pytest.fail(f"shape mismatch: golden={a.shape} actual={b.shape}")
psnr_db = _psnr(a, b)
assert psnr_db >= GOLDEN_PSNR_MIN_DB, (
f"hash drifted (got {sha[:12]}.., expected {expected_hash[:12]}..) and "
f"PSNR {psnr_db:.2f} dB < {GOLDEN_PSNR_MIN_DB} dB threshold; "
f"diff PNG at {b_path}"
)
@@ -0,0 +1,186 @@
"""Phase 2 characterization tests for coreml_suite.controlnet.
Locks shapes + dtypes + zero-fill behavior of expand_inputs / no_control /
extract_residual_kwargs / chunk_control. These pure helpers feed the Core ML
UNet's additional_residual_N inputs; any drift here silently breaks
ControlNet-based workflows.
"""
import numpy as np
import pytest
import torch
from coreml_suite.controlnet import (
chunk_control,
expand_inputs,
extract_residual_kwargs,
no_control,
)
@pytest.fixture(autouse=True)
def _deterministic_seed():
torch.manual_seed(0)
np.random.seed(0)
SD15_RESIDUAL_SPEC = {
"additional_residual_0": {"shape": (2, 320, 64, 64)},
"additional_residual_1": {"shape": (2, 640, 32, 32)},
"additional_residual_2": {"shape": (2, 1280, 8, 8)},
}
NON_RESIDUAL_SPEC = {
"sample": {"shape": (2, 4, 64, 64)},
"encoder_hidden_states": {"shape": (2, 768, 1, 77)},
}
# ---------- expand_inputs ----------------------------------------------------
def test_expand_inputs_doubles_singleton_numpy():
inputs = {"a": np.ones((1, 4), dtype=np.float32)}
out = expand_inputs(inputs)
assert out["a"].shape == (2, 4)
assert np.array_equal(out["a"], np.ones((2, 4)))
def test_expand_inputs_doubles_singleton_torch():
inputs = {"a": torch.ones(1, 4)}
out = expand_inputs(inputs)
assert out["a"].shape == (2, 4)
assert torch.equal(out["a"], torch.ones(2, 4))
def test_expand_inputs_doubles_singleton_list():
inputs = {"a": [42]}
out = expand_inputs(inputs)
assert out["a"] == [42, 42]
def test_expand_inputs_skips_already_batched():
"""batch > 1 inputs are returned unchanged (same object identity)."""
arr = np.ones((2, 4), dtype=np.float32)
tensor = torch.ones(3, 4)
lst = [1, 2]
out = expand_inputs({"a": arr, "b": tensor, "c": lst})
assert out["a"] is arr
assert out["b"] is tensor
assert out["c"] is lst
def test_expand_inputs_preserves_unknown_value_types():
# Strings/None pass through untouched — locks current permissive contract.
inputs = {"s": "hello", "none": None, "int": 7}
out = expand_inputs(inputs)
assert out == {"s": "hello", "none": None, "int": 7}
# ---------- no_control -------------------------------------------------------
def test_no_control_returns_zero_fp16_for_residuals():
out = no_control({**SD15_RESIDUAL_SPEC, **NON_RESIDUAL_SPEC})
# Only additional_residual_* keys are produced.
assert set(out.keys()) == set(SD15_RESIDUAL_SPEC.keys())
for key, spec in SD15_RESIDUAL_SPEC.items():
arr = out[key]
assert arr.shape == spec["shape"]
assert arr.dtype == np.float16
assert np.all(arr == 0)
def test_no_control_returns_empty_when_no_residuals():
out = no_control(NON_RESIDUAL_SPEC)
assert out == {}
# ---------- extract_residual_kwargs -----------------------------------------
def test_extract_residual_kwargs_empty_when_model_has_no_residual_inputs():
out = extract_residual_kwargs(NON_RESIDUAL_SPEC, control={"output": [], "middle": []})
assert out == {}
def test_extract_residual_kwargs_none_control_returns_no_control_shapes():
out = extract_residual_kwargs(SD15_RESIDUAL_SPEC, control=None)
assert set(out.keys()) == set(SD15_RESIDUAL_SPEC.keys())
for key, spec in SD15_RESIDUAL_SPEC.items():
assert out[key].shape == spec["shape"]
assert out[key].dtype == np.float16
assert np.all(out[key] == 0)
def test_extract_residual_kwargs_flattens_output_then_middle_and_casts_fp16():
"""output residuals come first (indexed 0..N-1), then middle residuals
(indexed N..M-1). Values come out of CPU as fp16 numpy arrays."""
control = {
"output": [torch.ones(2, 320, 64, 64) * 0.5, torch.ones(2, 640, 32, 32) * 2.0],
"middle": [torch.ones(2, 1280, 8, 8) * -1.0],
}
out = extract_residual_kwargs(SD15_RESIDUAL_SPEC, control)
assert set(out.keys()) == {"additional_residual_0", "additional_residual_1", "additional_residual_2"}
assert out["additional_residual_0"].shape == (2, 320, 64, 64)
assert out["additional_residual_1"].shape == (2, 640, 32, 32)
assert out["additional_residual_2"].shape == (2, 1280, 8, 8)
for arr in out.values():
assert arr.dtype == np.float16
# Locked order: index 0 == first output residual (0.5), index 2 == middle (-1.0).
assert np.allclose(out["additional_residual_0"], 0.5)
assert np.allclose(out["additional_residual_1"], 2.0)
assert np.allclose(out["additional_residual_2"], -1.0)
# ---------- chunk_control ----------------------------------------------------
def test_chunk_control_none_returns_list_of_nones_with_length_target():
"""`no_control` path: when there's no control, you get [None] * target_size
(NOT [None, None] regardless of target — this is the contract today)."""
assert chunk_control(None, 1) == [None]
assert chunk_control(None, 2) == [None, None]
assert chunk_control(None, 4) == [None, None, None, None]
@pytest.mark.parametrize(
"batch,target,expected_chunks",
[(1, 2, 1), (2, 2, 1), (3, 2, 2), (4, 2, 2), (5, 3, 2), (9, 4, 3)],
)
def test_chunk_control_shapes_after_chunking(batch, target, expected_chunks):
cn = {
"output": [
torch.randn(batch, 320, 64, 64),
torch.randn(batch, 640, 32, 32),
],
"middle": [torch.randn(batch, 1280, 8, 8)],
}
chunks = chunk_control(cn, target)
assert len(chunks) == expected_chunks
for c in chunks:
assert c["output"][0].shape == (target, 320, 64, 64)
assert c["output"][1].shape == (target, 640, 32, 32)
assert c["middle"][0].shape == (target, 1280, 8, 8)
def test_chunk_control_preserves_keys_order():
"""Output dicts contain exactly {"output", "middle"} in that order."""
cn = {
"output": [torch.zeros(2, 4, 4, 4)],
"middle": [torch.zeros(2, 4, 4, 4)],
}
chunks = chunk_control(cn, 2)
assert list(chunks[0].keys()) == ["output", "middle"]
def test_chunk_control_zero_pads_remainder():
"""A batch=3, target=2 split puts the third row alongside a zero row."""
cn = {
"output": [torch.arange(3 * 4).reshape(3, 1, 2, 2).float()],
"middle": [torch.arange(3 * 4).reshape(3, 1, 2, 2).float()],
}
chunks = chunk_control(cn, 2)
assert len(chunks) == 2
last_out = chunks[-1]["output"][0]
# First row is the original third row; second row is padding zeros.
assert torch.equal(last_out[0], cn["output"][0][2])
assert torch.equal(last_out[1], torch.zeros(1, 2, 2))
+228
View File
@@ -0,0 +1,228 @@
"""Phase 2 characterization tests for coreml_suite.models.CoreMLInputs.
Locks the shape transforms applied by chunks() and coreml_kwargs() for the
four model variants the suite supports: SD1.5, LCM (SD1.5 + timestep_cond),
SDXL base (time_ids len 6), and SDXL refiner (time_ids len 5).
These contracts feed the Core ML UNet at runtime; if Phase 3 silently
re-shapes them, generation breaks.
"""
import numpy as np
import pytest
import torch
from coreml_suite.models import CoreMLInputs
@pytest.fixture(autouse=True)
def _deterministic_seed():
torch.manual_seed(0)
np.random.seed(0)
# ---------- expected_inputs fixtures (mirror real model expectations) -------
SD15_EXPECTED = {
"sample": {"shape": (2, 4, 64, 64)},
"timestep": {"shape": (2,)},
"encoder_hidden_states": {"shape": (2, 768, 1, 77)},
}
SD15_WITH_CN = {
**SD15_EXPECTED,
"additional_residual_0": {"shape": (2, 320, 64, 64)},
"additional_residual_1": {"shape": (2, 640, 32, 32)},
}
LCM_EXPECTED = {
**SD15_EXPECTED,
"timestep_cond": {"shape": (2, 256)},
}
SDXL_BASE_EXPECTED = {
"sample": {"shape": (2, 4, 128, 128)},
"timestep": {"shape": (2,)},
"encoder_hidden_states": {"shape": (2, 2048, 1, 77)},
"time_ids": {"shape": (2, 6)},
"text_embeds": {"shape": (2, 1280)},
}
SDXL_REFINER_EXPECTED = {
"sample": {"shape": (2, 4, 128, 128)},
"timestep": {"shape": (2,)},
"encoder_hidden_states": {"shape": (2, 1280, 1, 77)},
"time_ids": {"shape": (2, 5)},
"text_embeds": {"shape": (2, 1280)},
}
def _sd15_inputs(batch=1, with_control=False, with_ts_cond=False):
x = torch.randn(batch, 4, 64, 64)
t = torch.full((batch,), 999.0)
context = torch.randn(batch, 77, 768)
control = None
if with_control:
control = {
"output": [torch.randn(batch, 320, 64, 64), torch.randn(batch, 640, 32, 32)],
"middle": [],
}
kwargs = {}
if with_ts_cond:
kwargs["timestep_cond"] = torch.randn(batch, 256)
return CoreMLInputs(x, t, context, control, **kwargs)
def _sdxl_inputs(batch=1, refiner=False):
x = torch.randn(batch, 4, 128, 128)
t = torch.full((batch,), 999.0)
ctx_dim = 1280 if refiner else 2048
context = torch.randn(batch, 77, ctx_dim)
time_ids_dim = 5 if refiner else 6
time_ids = torch.randn(batch, time_ids_dim)
text_embeds = torch.randn(batch, 1280)
return CoreMLInputs(
x, t, context, control=None, time_ids=time_ids, text_embeds=text_embeds
)
# ---------- coreml_kwargs ---------------------------------------------------
def test_coreml_kwargs_sd15_shapes_and_fp16():
out = _sd15_inputs(batch=1).coreml_kwargs(SD15_EXPECTED)
assert set(out.keys()) == {"sample", "encoder_hidden_states", "timestep"}
assert out["sample"].shape == (1, 4, 64, 64)
assert out["sample"].dtype == np.float16
# encoder_hidden_states is transposed (b, seq, dim) -> (b, dim, 1, seq).
assert out["encoder_hidden_states"].shape == (1, 768, 1, 77)
assert out["encoder_hidden_states"].dtype == np.float16
assert out["timestep"].shape == (1,)
assert out["timestep"].dtype == np.float16
def test_coreml_kwargs_sd15_with_controlnet_emits_residuals():
inputs = _sd15_inputs(batch=1, with_control=True)
out = inputs.coreml_kwargs(SD15_WITH_CN)
assert "additional_residual_0" in out
assert "additional_residual_1" in out
assert out["additional_residual_0"].shape == (1, 320, 64, 64)
assert out["additional_residual_1"].shape == (1, 640, 32, 32)
def test_coreml_kwargs_sd15_without_controlnet_zero_fills_residuals():
inputs = _sd15_inputs(batch=1, with_control=False)
out = inputs.coreml_kwargs(SD15_WITH_CN)
assert np.all(out["additional_residual_0"] == 0)
assert np.all(out["additional_residual_1"] == 0)
def test_coreml_kwargs_lcm_adds_timestep_cond():
inputs = _sd15_inputs(batch=1, with_ts_cond=True)
out = inputs.coreml_kwargs(LCM_EXPECTED)
assert "timestep_cond" in out
assert out["timestep_cond"].shape == (1, 256)
assert out["timestep_cond"].dtype == np.float16
def test_coreml_kwargs_lcm_skips_timestep_cond_when_not_provided():
"""timestep_cond is only forwarded when the input supplied one — even if
the model's expected_inputs lists it."""
inputs = _sd15_inputs(batch=1, with_ts_cond=False)
out = inputs.coreml_kwargs(LCM_EXPECTED)
assert "timestep_cond" not in out
def test_coreml_kwargs_sdxl_base_emits_time_ids_and_text_embeds():
out = _sdxl_inputs(batch=1, refiner=False).coreml_kwargs(SDXL_BASE_EXPECTED)
assert out["time_ids"].shape == (1, 6)
assert out["text_embeds"].shape == (1, 1280)
assert out["time_ids"].dtype == np.float16
assert out["text_embeds"].dtype == np.float16
def test_coreml_kwargs_sdxl_refiner_uses_len5_time_ids():
out = _sdxl_inputs(batch=1, refiner=True).coreml_kwargs(SDXL_REFINER_EXPECTED)
assert out["time_ids"].shape == (1, 5)
# ---------- chunks ----------------------------------------------------------
def test_chunks_sd15_pad_to_batch2_returns_one_chunk():
chunked = _sd15_inputs(batch=1).chunks(SD15_EXPECTED)
assert len(chunked) == 1
c = chunked[0]
assert c.x.shape == (2, 4, 64, 64)
assert c.t.shape == (2,)
# context shape: (b, seq, dim) padded along batch dim.
assert c.context.shape == (2, 77, 768)
assert c.control is None
assert c.ts_cond is None
assert c.time_ids is None
assert c.text_embeds is None
def test_chunks_sd15_with_controlnet_chunks_residuals_too():
chunked = _sd15_inputs(batch=1, with_control=True).chunks(SD15_EXPECTED)
assert len(chunked) == 1
cn = chunked[0].control
assert cn is not None
assert cn["output"][0].shape == (2, 320, 64, 64)
assert cn["output"][1].shape == (2, 640, 32, 32)
def test_chunks_lcm_carries_timestep_cond_per_chunk():
chunked = _sd15_inputs(batch=1, with_ts_cond=True).chunks(LCM_EXPECTED)
assert len(chunked) == 1
assert chunked[0].ts_cond is not None
assert chunked[0].ts_cond.shape == (2, 256)
def test_chunks_sdxl_base_propagates_time_ids_and_text_embeds():
chunked = _sdxl_inputs(batch=1, refiner=False).chunks(SDXL_BASE_EXPECTED)
assert len(chunked) == 1
c = chunked[0]
assert c.time_ids is not None and c.time_ids.shape == (2, 6)
assert c.text_embeds is not None and c.text_embeds.shape == (2, 1280)
def test_chunks_sdxl_refiner_uses_len5_time_ids():
chunked = _sdxl_inputs(batch=1, refiner=True).chunks(SDXL_REFINER_EXPECTED)
assert chunked[0].time_ids.shape == (2, 5)
def test_chunks_sdxl_synthesizes_zero_time_ids_when_caller_omits():
"""If the model expects time_ids but caller passed nothing, the suite
fabricates a zero-filled tensor. Lock that fallback."""
x = torch.randn(1, 4, 128, 128)
t = torch.full((1,), 999.0)
context = torch.randn(1, 77, 2048)
inputs = CoreMLInputs(x, t, context, control=None)
chunked = inputs.chunks(SDXL_BASE_EXPECTED)
assert chunked[0].time_ids.shape == (2, 6)
assert torch.equal(chunked[0].time_ids, torch.zeros(2, 6))
assert chunked[0].text_embeds.shape == (2, 1280)
assert torch.equal(chunked[0].text_embeds, torch.zeros(2, 1280))
def test_chunks_splits_batch_into_multiple_target2_chunks():
"""batch=5 with target_batch=2 -> 3 chunks (last padded)."""
chunked = _sd15_inputs(batch=5).chunks(SD15_EXPECTED)
assert len(chunked) == 3
for c in chunked:
assert c.x.shape == (2, 4, 64, 64)
assert c.context.shape == (2, 77, 768)
# Last chunk's second batch row is the zero-pad.
assert torch.equal(chunked[-1].x[1], torch.zeros(4, 64, 64))
def test_chunks_timestep_is_broadcast_from_first_value():
"""t is rebuilt from t[0] across all chunks: locks current behavior that
discards any per-row timestep variation."""
x = torch.randn(2, 4, 64, 64)
t = torch.tensor([42.0, 99.0]) # the second value will be lost
context = torch.randn(2, 77, 768)
inputs = CoreMLInputs(x, t, context, control=None)
chunked = inputs.chunks(SD15_EXPECTED)
assert chunked[0].t.shape == (2,)
assert torch.equal(chunked[0].t, torch.full((2,), 42.0))
+118
View File
@@ -0,0 +1,118 @@
"""Phase 2 characterization tests for coreml_suite.latents.
Locks the *current* behavior of chunk_batch / merge_chunks — including the
zero-pad regions and the truncation in merge — so the Phase 3 refactor
cannot silently shift either contract.
"""
import pytest
import torch
from coreml_suite.latents import chunk_batch, merge_chunks
@pytest.fixture(autouse=True)
def _deterministic_seed():
torch.manual_seed(0)
def _const_tensor(batch, *rest):
return torch.arange(batch * 4 * 8 * 8, dtype=torch.float32).reshape(batch, 4, 8, 8)
# ---------- chunk_batch ------------------------------------------------------
def test_chunk_batch_passthrough_when_shape_matches():
x = _const_tensor(2)
out = chunk_batch(x, (2, 4, 8, 8))
assert len(out) == 1
# passthrough: the same object identity is returned (no copy).
assert out[0] is x
def test_chunk_batch_pads_single_chunk_when_input_smaller():
"""batch=1, target=2 -> one padded chunk; the second row is exact zero."""
x = _const_tensor(1)
out = chunk_batch(x, (2, 4, 8, 8))
assert len(out) == 1
assert out[0].shape == (2, 4, 8, 8)
assert torch.equal(out[0][0], x[0])
assert torch.equal(out[0][1], torch.zeros(4, 8, 8))
def test_chunk_batch_splits_exact_multiple():
"""batch=4, target=2 -> two chunks, no padding."""
x = _const_tensor(4)
out = chunk_batch(x, (2, 4, 8, 8))
assert len(out) == 2
assert out[0].shape == (2, 4, 8, 8)
assert out[1].shape == (2, 4, 8, 8)
assert torch.equal(out[0], x[:2])
assert torch.equal(out[1], x[2:])
def test_chunk_batch_pads_remainder_chunk():
"""batch=5, target=2 -> chunks=[x[0:2], x[2:4]] then [x[4], 0]."""
x = _const_tensor(5)
out = chunk_batch(x, (2, 4, 8, 8))
assert len(out) == 3
assert torch.equal(out[0], x[0:2])
assert torch.equal(out[1], x[2:4])
last = out[-1]
assert last.shape == (2, 4, 8, 8)
assert torch.equal(last[0], x[4])
# The remainder row is zero-padded; lock that exact contract.
assert torch.equal(last[1], torch.zeros(4, 8, 8))
assert last[1].sum() == 0
@pytest.mark.parametrize(
"batch_size,target,expected_chunks",
[
(1, 4, 1),
(3, 2, 2),
(5, 3, 2),
(9, 4, 3),
],
)
def test_chunk_batch_pad_region_is_zero(batch_size, target, expected_chunks):
x = _const_tensor(batch_size)
out = chunk_batch(x, (target, 4, 8, 8))
assert len(out) == expected_chunks
mod = batch_size % target
if mod == 0 and batch_size >= target:
return
last = out[-1]
pad_rows = target - (mod if (mod != 0 and batch_size >= target) else batch_size)
pad_region = last[-pad_rows:]
assert torch.equal(pad_region, torch.zeros_like(pad_region))
# ---------- merge_chunks -----------------------------------------------------
def test_merge_chunks_exact_concat():
x = _const_tensor(4)
chunks = chunk_batch(x, (2, 4, 8, 8))
merged = merge_chunks(chunks, x.shape)
assert merged.shape == x.shape
assert torch.equal(merged, x)
def test_merge_chunks_truncates_padding():
"""Round-trip with a padded last chunk drops the pad rows."""
x = _const_tensor(5)
chunks = chunk_batch(x, (2, 4, 8, 8))
merged = merge_chunks(chunks, x.shape)
assert merged.shape == x.shape
assert torch.equal(merged, x)
def test_merge_chunks_singleton_returns_equal_copy_when_shape_matches():
"""A singleton chunk list still goes through torch.cat, so we get a new
tensor equal to the input — locked here because Phase 3 might be tempted
to short-circuit and accidentally return the same object."""
x = _const_tensor(2)
out = merge_chunks([x], x.shape)
assert torch.equal(out, x)
assert out is not x
@@ -0,0 +1,175 @@
"""Phase 2 characterization tests for CoreMLConverter.out_name composition.
out_name is encoded into the .mlpackage filename and therefore drives the
"have we converted this combo already?" cache check. Drift here silently
invalidates user caches and breaks workflow node references.
Tests intercept converter.get_out_path to capture the composed string,
and stub out the heavy conversion + Core ML model load.
"""
import pytest
import folder_paths
from coreml_suite import converter
from coreml_suite import nodes as nodes_mod
from coreml_suite.nodes import CoreMLConverter
@pytest.fixture
def capture_out_name(monkeypatch):
captured = {}
def fake_get_out_path(submodule, name):
captured["submodule"] = submodule
captured["out_name"] = name
return f"/tmp/fake/{name}.mlpackage"
monkeypatch.setattr(converter, "get_out_path", fake_get_out_path)
monkeypatch.setattr(converter, "convert", lambda **kwargs: None)
monkeypatch.setattr(
converter,
"compile_model",
lambda out_path, out_name, submodule_name: f"/tmp/fake/{out_name}_{submodule_name}.mlmodelc",
)
monkeypatch.setattr(
folder_paths,
"get_full_path",
lambda kind, name: f"/tmp/ckpts/{name}" if kind == "checkpoints" else None,
)
monkeypatch.setattr(nodes_mod, "CoreMLModel", lambda *a, **kw: object())
return captured
def _convert(
*,
ckpt_name="dreamshaper_8.safetensors",
model_version="SD15",
height=512,
width=512,
batch_size=1,
attention_implementation="SPLIT_EINSUM",
compute_unit="CPU_AND_NE",
controlnet_support=False,
lora_params=None,
):
node = CoreMLConverter()
node.convert(
ckpt_name=ckpt_name,
model_version=model_version,
height=height,
width=width,
batch_size=batch_size,
attention_implementation=attention_implementation,
compute_unit=compute_unit,
controlnet_support=controlnet_support,
lora_params=lora_params,
)
# ---------- basic golden strings --------------------------------------------
@pytest.mark.parametrize(
"attn_name,suffix",
[
("SPLIT_EINSUM", "se"),
("SPLIT_EINSUM_V2", "se2"),
("ORIGINAL", "orig"),
],
)
def test_out_name_attention_suffix(capture_out_name, attn_name, suffix):
_convert(attention_implementation=attn_name)
assert capture_out_name["out_name"] == f"dreamshaper_8_1x512x512_{suffix}"
def test_out_name_includes_batch_and_size(capture_out_name):
_convert(batch_size=4, width=768, height=1024)
assert capture_out_name["out_name"] == "dreamshaper_8_4x768x1024_se"
def test_out_name_appends_cn_suffix_when_controlnet_support_true(capture_out_name):
_convert(controlnet_support=True)
assert capture_out_name["out_name"] == "dreamshaper_8_1x512x512_cn_se"
def test_out_name_drops_dot_extension_only_at_first_period(capture_out_name):
"""`ckpt_name.split('.')[0]` — first '.' wins; locked behaviour."""
_convert(ckpt_name="my.checkpoint.v2.safetensors")
assert capture_out_name["out_name"] == "my_1x512x512_se"
def test_out_name_replaces_spaces_with_underscores(capture_out_name):
_convert(ckpt_name="dream shaper 8.safetensors")
assert capture_out_name["out_name"] == "dream_shaper_8_1x512x512_se"
# ---------- LoRA suffixes ---------------------------------------------------
def test_out_name_with_single_lora(capture_out_name, monkeypatch):
monkeypatch.setattr(
folder_paths,
"get_full_path",
lambda kind, name: f"/tmp/{kind}/{name}",
)
_convert(lora_params={"epi_noiseoffset.safetensors": (0.8,)})
assert (
capture_out_name["out_name"]
== "dreamshaper_8_epi_noiseoffset_1x512x512_se"
)
def test_out_name_with_multiple_loras_sorted(capture_out_name, monkeypatch):
"""LoRAs are sorted by name then joined with '_' — locks the order."""
monkeypatch.setattr(
folder_paths,
"get_full_path",
lambda kind, name: f"/tmp/{kind}/{name}",
)
_convert(
lora_params={
"zoom.safetensors": (1.0,),
"alpha.safetensors": (0.5,),
"moody.safetensors": (0.3,),
}
)
assert (
capture_out_name["out_name"]
== "dreamshaper_8_alpha_moody_zoom_1x512x512_se"
)
def test_out_name_lora_plus_controlnet(capture_out_name, monkeypatch):
monkeypatch.setattr(
folder_paths,
"get_full_path",
lambda kind, name: f"/tmp/{kind}/{name}",
)
_convert(
lora_params={"a.safetensors": (1.0,)},
controlnet_support=True,
)
assert capture_out_name["out_name"] == "dreamshaper_8_a_1x512x512_cn_se"
# ---------- sdxl combinations -----------------------------------------------
def test_out_name_sdxl_1024(capture_out_name):
_convert(
ckpt_name="sd_xl_base_1.0.safetensors",
model_version="SDXL",
width=1024,
height=1024,
attention_implementation="ORIGINAL",
compute_unit="CPU_AND_GPU",
)
assert capture_out_name["out_name"] == "sd_xl_base_1_1x1024x1024_orig"
# ---------- submodule path --------------------------------------------------
def test_get_out_path_invoked_with_unet_submodule(capture_out_name):
_convert()
assert capture_out_name["submodule"] == "unet"
@@ -0,0 +1,147 @@
"""Phase 2 characterization tests for add_sdxl_model_options.
Locks the SDXL time_ids / text_embeds assembly: base produces a (2, 6)
time_ids vector (h, w, crop_h, crop_w, target_h, target_w), refiner
produces (2, 5) (h, w, crop_h, crop_w, aesthetic_score). Both stack pos
then neg along the batch dim, and text_embeds is cat(pos_pooled,
neg_pooled).
Uses a SimpleNamespace-shaped fake ModelPatcher because exercising the
real comfy.model_patcher.ModelPatcher here is overkill — only three
attribute paths are read by the SUT.
"""
import inspect
from types import SimpleNamespace
import pytest
import torch
from coreml_suite.models import add_sdxl_model_options
@pytest.fixture(autouse=True)
def _deterministic_seed():
torch.manual_seed(0)
def _fake_patcher(is_base: bool, is_refiner: bool):
diffusion = SimpleNamespace(is_sdxl_base=is_base, is_sdxl_refiner=is_refiner)
model = SimpleNamespace(diffusion_model=diffusion)
patcher = SimpleNamespace(model=model, model_options={})
patcher.clone = lambda: patcher # in-place: simplest mock that matches contract
return patcher
def _cond(pooled, **overrides):
base = {"pooled_output": pooled}
base.update(overrides)
return [(None, base)]
def _closure_vars(wrapper):
return inspect.getclosurevars(wrapper).nonlocals
# ---------- base (len 6) -----------------------------------------------------
def test_add_sdxl_model_options_base_produces_len6_time_ids_with_defaults():
pos_pooled = torch.randn(1, 1280)
neg_pooled = torch.randn(1, 1280)
patcher = _fake_patcher(is_base=True, is_refiner=False)
out = add_sdxl_model_options(patcher, _cond(pos_pooled), _cond(neg_pooled))
wrapper = out.model_options["model_function_wrapper"]
closure = _closure_vars(wrapper)
assert closure["time_ids"].shape == (2, 6)
# Defaults: 768 height/width, 0 crop, 768 target.
expected = torch.tensor([[768, 768, 0, 0, 768, 768], [768, 768, 0, 0, 768, 768]])
assert torch.equal(closure["time_ids"], expected)
assert closure["refiner"] is False
def test_add_sdxl_model_options_base_respects_overrides():
pos_pooled = torch.randn(1, 1280)
neg_pooled = torch.randn(1, 1280)
patcher = _fake_patcher(is_base=True, is_refiner=False)
out = add_sdxl_model_options(
patcher,
_cond(pos_pooled, height=1024, width=512, crop_h=8, crop_w=4, target_height=1024, target_width=1024),
_cond(neg_pooled, height=256, width=256, crop_h=0, crop_w=0, target_height=256, target_width=256),
)
closure = _closure_vars(out.model_options["model_function_wrapper"])
expected = torch.tensor([[1024, 512, 8, 4, 1024, 1024], [256, 256, 0, 0, 256, 256]])
assert torch.equal(closure["time_ids"], expected)
# ---------- refiner (len 5) -------------------------------------------------
def test_add_sdxl_model_options_refiner_produces_len5_time_ids():
pos_pooled = torch.randn(1, 1280)
neg_pooled = torch.randn(1, 1280)
patcher = _fake_patcher(is_base=False, is_refiner=True)
out = add_sdxl_model_options(patcher, _cond(pos_pooled), _cond(neg_pooled))
closure = _closure_vars(out.model_options["model_function_wrapper"])
assert closure["time_ids"].shape == (2, 5)
# Defaults: pos aesthetic_score=6, neg aesthetic_score=2.5.
expected = torch.tensor(
[[768, 768, 0, 0, 6.0], [768, 768, 0, 0, 2.5]],
)
assert torch.equal(closure["time_ids"], expected)
assert closure["refiner"] is True
def test_add_sdxl_model_options_refiner_respects_aesthetic_score_overrides():
pos_pooled = torch.randn(1, 1280)
neg_pooled = torch.randn(1, 1280)
patcher = _fake_patcher(is_base=False, is_refiner=True)
out = add_sdxl_model_options(
patcher,
_cond(pos_pooled, aesthetic_score=8.5),
_cond(neg_pooled, aesthetic_score=1.5),
)
closure = _closure_vars(out.model_options["model_function_wrapper"])
expected = torch.tensor([[768, 768, 0, 0, 8.5], [768, 768, 0, 0, 1.5]])
assert torch.equal(closure["time_ids"], expected)
# ---------- text_embeds ----------------------------------------------------
def test_text_embeds_is_concat_pos_then_neg_along_batch():
pos_pooled = torch.full((1, 1280), 1.0)
neg_pooled = torch.full((1, 1280), -1.0)
patcher = _fake_patcher(is_base=True, is_refiner=False)
out = add_sdxl_model_options(patcher, _cond(pos_pooled), _cond(neg_pooled))
closure = _closure_vars(out.model_options["model_function_wrapper"])
embeds = closure["text_embeds"]
assert embeds.shape == (2, 1280)
assert torch.equal(embeds[0], pos_pooled[0])
assert torch.equal(embeds[1], neg_pooled[0])
def test_neither_base_nor_refiner_yields_len4_time_ids():
"""Locked edge case: if both is_sdxl_base and is_sdxl_refiner are False,
no extra entries are appended -> time_ids is only the 4 shared fields."""
pos_pooled = torch.randn(1, 1280)
neg_pooled = torch.randn(1, 1280)
patcher = _fake_patcher(is_base=False, is_refiner=False)
out = add_sdxl_model_options(patcher, _cond(pos_pooled), _cond(neg_pooled))
closure = _closure_vars(out.model_options["model_function_wrapper"])
assert closure["time_ids"].shape == (2, 4)
# ---------- patcher contract ------------------------------------------------
def test_model_options_dict_is_merged_into_clone_not_original():
"""Verify SUT writes to the clone's model_options. Our fake reuses the
same instance, so the merge should still leave the model_function_wrapper
key present after the call."""
pos_pooled = torch.randn(1, 1280)
neg_pooled = torch.randn(1, 1280)
patcher = _fake_patcher(is_base=True, is_refiner=False)
patcher.model_options["pre_existing"] = "kept"
out = add_sdxl_model_options(patcher, _cond(pos_pooled), _cond(neg_pooled))
assert "model_function_wrapper" in out.model_options
assert out.model_options["pre_existing"] == "kept"