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:
@@ -9,3 +9,4 @@ coremlsuite-venv/
|
||||
experiment_results/
|
||||
test_results/
|
||||
bench/scripts/*.log
|
||||
tests/m2/_latest_generated.png
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 451 KiB |
@@ -0,0 +1 @@
|
||||
9e12ba1f98969a99c17e84a87a77c4561fcac3ca2ddd6338a7ae89d324df6994
|
||||
@@ -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))
|
||||
@@ -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))
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user