diff --git a/.gitignore b/.gitignore index f22f155..b578bb9 100644 --- a/.gitignore +++ b/.gitignore @@ -9,3 +9,4 @@ coremlsuite-venv/ experiment_results/ test_results/ bench/scripts/*.log +tests/m2/_latest_generated.png diff --git a/pyproject.toml b/pyproject.toml index b1906e8..b35c783 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"] diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..0754a0c --- /dev/null +++ b/tests/conftest.py @@ -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 diff --git a/tests/m2/__init__.py b/tests/m2/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/m2/goldens/sd15_seed42.png b/tests/m2/goldens/sd15_seed42.png new file mode 100644 index 0000000..a561e16 Binary files /dev/null and b/tests/m2/goldens/sd15_seed42.png differ diff --git a/tests/m2/goldens/sd15_seed42.sha256 b/tests/m2/goldens/sd15_seed42.sha256 new file mode 100644 index 0000000..6ea161e --- /dev/null +++ b/tests/m2/goldens/sd15_seed42.sha256 @@ -0,0 +1 @@ +9e12ba1f98969a99c17e84a87a77c4561fcac3ca2ddd6338a7ae89d324df6994 diff --git a/tests/m2/test_golden_image.py b/tests/m2/test_golden_image.py new file mode 100644 index 0000000..d9d70e8 --- /dev/null +++ b/tests/m2/test_golden_image.py @@ -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}" + ) diff --git a/tests/unit/test_characterization_controlnet.py b/tests/unit/test_characterization_controlnet.py new file mode 100644 index 0000000..1c20756 --- /dev/null +++ b/tests/unit/test_characterization_controlnet.py @@ -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)) diff --git a/tests/unit/test_characterization_inputs.py b/tests/unit/test_characterization_inputs.py new file mode 100644 index 0000000..c1769ca --- /dev/null +++ b/tests/unit/test_characterization_inputs.py @@ -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)) diff --git a/tests/unit/test_characterization_latents.py b/tests/unit/test_characterization_latents.py new file mode 100644 index 0000000..ff6f136 --- /dev/null +++ b/tests/unit/test_characterization_latents.py @@ -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 diff --git a/tests/unit/test_characterization_out_name.py b/tests/unit/test_characterization_out_name.py new file mode 100644 index 0000000..cb2f22d --- /dev/null +++ b/tests/unit/test_characterization_out_name.py @@ -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" diff --git a/tests/unit/test_characterization_sdxl_options.py b/tests/unit/test_characterization_sdxl_options.py new file mode 100644 index 0000000..7b5f9f0 --- /dev/null +++ b/tests/unit/test_characterization_sdxl_options.py @@ -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"