Compare commits
18
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9af24733d4 | ||
|
|
f03db59ef8 | ||
|
|
aceba439d1 | ||
|
|
31774e3324 | ||
|
|
6a0ccbeb60 | ||
|
|
29b493454a | ||
|
|
60f2be86e6 | ||
|
|
085509d83e | ||
|
|
1240524201 | ||
|
|
3adb216763 | ||
|
|
0bbd8d8e0d | ||
|
|
2b649e6606 | ||
|
|
1e5791d108 | ||
|
|
8382b13598 | ||
|
|
5dafd261b7 | ||
|
|
04911d0052 | ||
|
|
cf6d7c6855 | ||
|
|
ef2a18cff3 |
@@ -0,0 +1,33 @@
|
|||||||
|
name: Tier 0 — Unit (Linux)
|
||||||
|
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
branches: [main]
|
||||||
|
pull_request:
|
||||||
|
|
||||||
|
# Minimal-deps run: Tier 0 must work without ComfyUI, coremltools, or
|
||||||
|
# python_coreml_stable_diffusion (Linux CI image won't have them). The
|
||||||
|
# in-tree purity gate (tests/unit/test_tier0_purity.py) double-checks
|
||||||
|
# that the suite hasn't started leaking framework imports.
|
||||||
|
jobs:
|
||||||
|
unit:
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
timeout-minutes: 10
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
|
||||||
|
- uses: actions/setup-python@v5
|
||||||
|
with:
|
||||||
|
python-version: "3.11"
|
||||||
|
|
||||||
|
- name: Install Tier 0 deps
|
||||||
|
run: |
|
||||||
|
python -m pip install --upgrade pip
|
||||||
|
# Tier 0 only needs torch + numpy + pytest; everything else
|
||||||
|
# is Mac-only.
|
||||||
|
python -m pip install \
|
||||||
|
"torch==2.0.1" "numpy<1.25" \
|
||||||
|
"pytest>=8" "pytest-xdist"
|
||||||
|
|
||||||
|
- name: Run Tier 0
|
||||||
|
run: pytest -m unit tests/ -v
|
||||||
@@ -0,0 +1,31 @@
|
|||||||
|
name: Tier 1 — Smoke (macOS-ARM)
|
||||||
|
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
branches: [main]
|
||||||
|
pull_request:
|
||||||
|
# Gate behind the run-tier1 label too, so external PRs that touch
|
||||||
|
# only docs don't burn a minute of macOS-ARM time. Maintainers can
|
||||||
|
# always re-run via the run-tier1 label.
|
||||||
|
types: [opened, synchronize, reopened, labeled]
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
smoke:
|
||||||
|
if: |
|
||||||
|
github.event_name == 'push' ||
|
||||||
|
github.event.action != 'labeled' ||
|
||||||
|
contains(github.event.pull_request.labels.*.name, 'run-tier1')
|
||||||
|
runs-on: macos-14 # M1, Apple Silicon hosted runner
|
||||||
|
timeout-minutes: 20
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
|
||||||
|
- uses: astral-sh/setup-uv@v3
|
||||||
|
with:
|
||||||
|
enable-cache: true
|
||||||
|
|
||||||
|
- name: uv sync
|
||||||
|
run: uv sync --no-install-project
|
||||||
|
|
||||||
|
- name: Run Tier 1 (synthetic micro-UNet smoke)
|
||||||
|
run: uv run pytest -m smoke tests/ -v
|
||||||
@@ -0,0 +1,124 @@
|
|||||||
|
name: Tier 2 — M2 / ANE (self-hosted)
|
||||||
|
|
||||||
|
on:
|
||||||
|
pull_request:
|
||||||
|
# `labeled` fires when run-m2 is first added; `synchronize`/`reopened`
|
||||||
|
# re-run on every subsequent push while the label is present, so the
|
||||||
|
# result tracks the PR head instead of going stale. The `if` below keeps
|
||||||
|
# the run gated on the run-m2 label for all pull_request events.
|
||||||
|
types: [labeled, synchronize, reopened]
|
||||||
|
schedule:
|
||||||
|
# Nightly at 04:00 UTC (~05/06 in PL). Keeps the M2 path honest
|
||||||
|
# without burning the runner on every PR.
|
||||||
|
- cron: "0 4 * * *"
|
||||||
|
workflow_dispatch:
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
m2:
|
||||||
|
if: |
|
||||||
|
github.event_name == 'schedule' ||
|
||||||
|
github.event_name == 'workflow_dispatch' ||
|
||||||
|
(github.event_name == 'pull_request' &&
|
||||||
|
contains(github.event.pull_request.labels.*.name, 'run-m2'))
|
||||||
|
# Self-hosted Apple Silicon runner. Prerequisites: COMFY_DIR pointing at
|
||||||
|
# a runner-owned ComfyUI clone, plus a cached SD1.5 checkpoint.
|
||||||
|
runs-on: [self-hosted, macOS, ARM64, coreml]
|
||||||
|
timeout-minutes: 90
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
|
||||||
|
# Hybrid ComfyUI strategy:
|
||||||
|
# - schedule (nightly) -> latest origin/master + ComfyUI's own
|
||||||
|
# requirements.txt (constrained). Canary for upstream API breakage.
|
||||||
|
# - PR label / dispatch -> the pinned requires-comfyui SHA + the frozen
|
||||||
|
# `comfy` uv group. Reproducible merge gate, immune to overnight drift.
|
||||||
|
- name: Resolve ComfyUI ref + mode
|
||||||
|
run: |
|
||||||
|
if [ "$GITHUB_EVENT_NAME" = "schedule" ]; then
|
||||||
|
echo "COMFY_MODE=latest" >> "$GITHUB_ENV"
|
||||||
|
echo "COMFY_REF=master" >> "$GITHUB_ENV"
|
||||||
|
else
|
||||||
|
PIN="$(sed -nE 's/^requires-comfyui *= *"==?([0-9a-f]+)".*/\1/p' pyproject.toml)"
|
||||||
|
if [ -z "$PIN" ]; then echo "could not parse requires-comfyui from pyproject.toml"; exit 1; fi
|
||||||
|
echo "COMFY_MODE=pinned" >> "$GITHUB_ENV"
|
||||||
|
echo "COMFY_REF=$PIN" >> "$GITHUB_ENV"
|
||||||
|
fi
|
||||||
|
|
||||||
|
- name: Set up ComfyUI checkout
|
||||||
|
# COMFY_DIR is exported by the self-hosted runner's .env and MUST be a
|
||||||
|
# runner-owned ComfyUI clone (never your dev checkout — this step does
|
||||||
|
# git reset --hard and rewrites custom_nodes). Cloned on first run.
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
if [ -z "${COMFY_DIR:-}" ]; then echo "COMFY_DIR unset"; exit 1; fi
|
||||||
|
# Init-in-place rather than `git clone`: COMFY_DIR may already hold the
|
||||||
|
# cached checkpoint (models/checkpoints) or converted .mlmodelc, and
|
||||||
|
# `git clone` refuses a non-empty target. init + fetch + `checkout -f`
|
||||||
|
# populates the ComfyUI tree while leaving untracked files (the
|
||||||
|
# checkpoint, the cached models) untouched — so setup order is free.
|
||||||
|
if [ ! -d "$COMFY_DIR/.git" ]; then
|
||||||
|
echo "initialising ComfyUI repo in $COMFY_DIR"
|
||||||
|
mkdir -p "$COMFY_DIR"
|
||||||
|
git -C "$COMFY_DIR" init -q
|
||||||
|
fi
|
||||||
|
git -C "$COMFY_DIR" remote get-url origin >/dev/null 2>&1 \
|
||||||
|
|| git -C "$COMFY_DIR" remote add origin https://github.com/comfyanonymous/ComfyUI.git
|
||||||
|
git -C "$COMFY_DIR" fetch --quiet origin
|
||||||
|
if [ "$COMFY_MODE" = "latest" ]; then
|
||||||
|
git -C "$COMFY_DIR" checkout -f -B master origin/master
|
||||||
|
else
|
||||||
|
git -C "$COMFY_DIR" checkout -f "$COMFY_REF"
|
||||||
|
fi
|
||||||
|
COMFY_SHA="$(git -C "$COMFY_DIR" rev-parse HEAD)"
|
||||||
|
echo "COMFY_SHA=$COMFY_SHA" >> "$GITHUB_ENV"
|
||||||
|
echo "Tier 2 mode=$COMFY_MODE, ComfyUI \`$COMFY_SHA\`" >> "$GITHUB_STEP_SUMMARY"
|
||||||
|
|
||||||
|
# Point ComfyUI's custom-node loader at this checkout. Refresh the
|
||||||
|
# symlink only; refuse to clobber a real directory (guards against a
|
||||||
|
# COMFY_DIR that is accidentally a dev checkout).
|
||||||
|
NODE_LINK="$COMFY_DIR/custom_nodes/ComfyUI-CoreMLSuite"
|
||||||
|
if [ -e "$NODE_LINK" ] && [ ! -L "$NODE_LINK" ]; then
|
||||||
|
echo "ERROR: $NODE_LINK is a real directory, not a symlink."
|
||||||
|
echo "COMFY_DIR must be a runner-owned ComfyUI, not your dev checkout."
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
mkdir -p "$COMFY_DIR/custom_nodes"
|
||||||
|
ln -sfn "$GITHUB_WORKSPACE" "$NODE_LINK"
|
||||||
|
|
||||||
|
- name: Install dependencies
|
||||||
|
run: |
|
||||||
|
set -euo pipefail
|
||||||
|
if [ "$COMFY_MODE" = "latest" ]; then
|
||||||
|
# Node deps (our coremltools-9 toolchain), then ComfyUI's own
|
||||||
|
# requirements for the pulled SHA, capped by the toolchain ceiling.
|
||||||
|
uv sync
|
||||||
|
uv pip install -r "$COMFY_DIR/requirements.txt" \
|
||||||
|
-c constraints/comfy-ceiling.txt
|
||||||
|
else
|
||||||
|
# Pinned gate: the frozen group mirrors the known-good pinned SHA.
|
||||||
|
uv sync --group comfy
|
||||||
|
fi
|
||||||
|
|
||||||
|
- name: Start ComfyUI server (background)
|
||||||
|
run: |
|
||||||
|
cd "$COMFY_DIR"
|
||||||
|
nohup "$GITHUB_WORKSPACE/.venv/bin/python" main.py --port 8188 --cpu-vae > /tmp/comfyui-ci.log 2>&1 &
|
||||||
|
# Poll the HTTP endpoint for readiness — robust to startup-banner
|
||||||
|
# wording / colored-log changes in a floating-latest ComfyUI.
|
||||||
|
for _ in $(seq 1 90); do
|
||||||
|
if curl -sf -o /dev/null http://127.0.0.1:8188/system_stats; then
|
||||||
|
echo "comfy ready (ComfyUI ${COMFY_SHA:-unknown})"; exit 0
|
||||||
|
fi
|
||||||
|
sleep 2
|
||||||
|
done
|
||||||
|
echo "comfy failed to start"; tail -100 /tmp/comfyui-ci.log; exit 1
|
||||||
|
|
||||||
|
- name: Run Tier 2 (m2 marker)
|
||||||
|
# The golden-image workflow drives the Core ML Converter node, so the
|
||||||
|
# UNet is converted on demand on the first run and reused from the
|
||||||
|
# runner-local cache afterwards.
|
||||||
|
run: uv run --no-sync pytest -m m2 tests/ -v
|
||||||
|
|
||||||
|
- name: Stop ComfyUI server
|
||||||
|
if: always()
|
||||||
|
run: pkill -f "main.py.*8188" || true
|
||||||
+4
-1
@@ -1,3 +1,6 @@
|
|||||||
playground/
|
playground/
|
||||||
experiments/
|
|
||||||
__pycache__/
|
__pycache__/
|
||||||
|
models/
|
||||||
|
.venv/
|
||||||
|
test_results/
|
||||||
|
tests/m2/_latest_generated.png
|
||||||
|
|||||||
@@ -370,6 +370,47 @@ The models used in this workflow are available at the following links:
|
|||||||
|
|
||||||

|

|
||||||
|
|
||||||
|
## Quantization (opt-in)
|
||||||
|
|
||||||
|
The `Core ML Converter` and `Core ML LCM Converter` nodes accept an
|
||||||
|
optional `quantize_nbits` dropdown that runs k-means weight palettization
|
||||||
|
(`coremltools.optimize.coreml.palettize_weights`) on the UNet before save.
|
||||||
|
|
||||||
|
Values: `none` (default — no quantization, identical to unquantized
|
||||||
|
behavior and filenames), `8`, `6`, `4`. The number is appended to the
|
||||||
|
.mlpackage stem as `_q<bits>` so quantized and unquantized variants
|
||||||
|
coexist on disk and in cache.
|
||||||
|
|
||||||
|
### SD1.5 1×512×512 SPLIT_EINSUM tradeoffs (M2 Pro, ANE)
|
||||||
|
|
||||||
|
Measured with 20 UNet forward passes at a fixed seed for the PSNR
|
||||||
|
comparison:
|
||||||
|
|
||||||
|
| nbits | size (MB) | size vs none | fwd median (ms) | PSNR vs `none` (dB) |
|
||||||
|
|---|---:|---:|---:|---:|
|
||||||
|
| none | 1641 | 1.000 | 197.1 | — |
|
||||||
|
| 8 | 822 | 0.501 | 186.6 | 53.5 |
|
||||||
|
| 6 | 617 | 0.376 | 183.0 | 40.2 |
|
||||||
|
| 4 | 412 | 0.251 | 179.8 | 27.5 |
|
||||||
|
|
||||||
|
PSNR here is computed on the raw `noise_pred` output of a single UNet
|
||||||
|
forward at a fixed seed, not on the final decoded image — it isolates
|
||||||
|
the quantization-induced drift from sampler / VAE noise. Final-image
|
||||||
|
PSNR is comfortably higher (the sampler averages over 20 steps).
|
||||||
|
|
||||||
|
### Recommended settings per chip / RAM
|
||||||
|
|
||||||
|
- **8 GB RAM (M1 base, M2 base):** `nbits=4`. ~4× smaller model, still
|
||||||
|
loads, PSNR 27 dB is visually identical at SD1.5 sizes.
|
||||||
|
- **16 GB RAM (M1/M2/M3 Pro):** `nbits=6` is the sweet spot — ~2.7×
|
||||||
|
smaller, PSNR 40 dB, no perceptible quality drop.
|
||||||
|
- **32 GB+ RAM (Max / Ultra):** `nbits=8` if you want the safety
|
||||||
|
margin, `none` if you want bit-identical output for golden testing.
|
||||||
|
|
||||||
|
The default stays `none` so existing workflows produce byte-for-byte
|
||||||
|
identical output — the golden-image anchor (`tests/m2/test_golden_image.py`)
|
||||||
|
verifies this on every Tier 2 run.
|
||||||
|
|
||||||
## Limitations
|
## Limitations
|
||||||
|
|
||||||
- Core ML models are fixed in terms of their inputs and outputs.
|
- Core ML models are fixed in terms of their inputs and outputs.
|
||||||
|
|||||||
@@ -0,0 +1,4 @@
|
|||||||
|
"""Top-level conftest: prevent pytest from importing the repo-root
|
||||||
|
__init__.py (the ComfyUI custom-node entry point pulls in comfy + nodes,
|
||||||
|
which breaks the Tier-0 'no-framework' promise)."""
|
||||||
|
collect_ignore = ["__init__.py"]
|
||||||
@@ -0,0 +1,18 @@
|
|||||||
|
# Toolchain ceiling for installing a floating-latest ComfyUI's requirements.txt
|
||||||
|
# in the Tier 2 nightly canary (.github/workflows/tier2.yml, latest mode).
|
||||||
|
#
|
||||||
|
# ComfyUI's requirements.txt requests bare `torch`/`torchvision`/`torchaudio`
|
||||||
|
# and `numpy>=1.25.0`, which would float past the versions coremltools 9 /
|
||||||
|
# apple-ml-stable-diffusion have been validated against.
|
||||||
|
# These constraints cap the resolution so the canary keeps testing the same
|
||||||
|
# toolchain the suite actually ships.
|
||||||
|
#
|
||||||
|
# If upstream ComfyUI ever hard-requires something beyond these bounds, the
|
||||||
|
# install FAILS — and that failure is the signal we want: it means the host
|
||||||
|
# outgrew the pinned toolchain and coremltools / ml-stable-diffusion need a
|
||||||
|
# deliberate bump, not a silent float.
|
||||||
|
torch>=2.7,<2.8
|
||||||
|
torchvision>=0.22,<0.23
|
||||||
|
torchaudio>=2.7,<2.8
|
||||||
|
numpy>=1.25,<2
|
||||||
|
coremltools>=9,<10
|
||||||
+13
-61
@@ -1,62 +1,14 @@
|
|||||||
from itertools import chain
|
"""Compatibility shim — re-exports from coreml_suite.core.controlnet."""
|
||||||
from math import ceil
|
from coreml_suite.core.controlnet import (
|
||||||
|
chunk_control,
|
||||||
|
expand_inputs,
|
||||||
|
extract_residual_kwargs,
|
||||||
|
no_control,
|
||||||
|
)
|
||||||
|
|
||||||
import numpy as np
|
__all__ = [
|
||||||
import torch
|
"chunk_control",
|
||||||
|
"expand_inputs",
|
||||||
from coreml_suite.latents import chunk_batch
|
"extract_residual_kwargs",
|
||||||
|
"no_control",
|
||||||
|
]
|
||||||
def expand_inputs(inputs):
|
|
||||||
expanded = inputs.copy()
|
|
||||||
for k, v in inputs.items():
|
|
||||||
if isinstance(v, np.ndarray):
|
|
||||||
expanded[k] = np.concatenate([v] * 2) if v.shape[0] == 1 else v
|
|
||||||
elif isinstance(v, torch.Tensor):
|
|
||||||
expanded[k] = torch.cat([v] * 2) if v.shape[0] == 1 else v
|
|
||||||
elif isinstance(v, list):
|
|
||||||
expanded[k] = v * 2 if len(v) == 1 else v
|
|
||||||
elif isinstance(v, dict):
|
|
||||||
expand_inputs(v)
|
|
||||||
return expanded
|
|
||||||
|
|
||||||
|
|
||||||
def extract_residual_kwargs(expected_inputs, control):
|
|
||||||
if "additional_residual_0" not in expected_inputs.keys():
|
|
||||||
return {}
|
|
||||||
if control is None:
|
|
||||||
return no_control(expected_inputs)
|
|
||||||
|
|
||||||
residual_kwargs = {
|
|
||||||
"additional_residual_{}".format(i): r.cpu().numpy().astype(np.float16)
|
|
||||||
for i, r in enumerate(chain(control["output"], control["middle"]))
|
|
||||||
}
|
|
||||||
return residual_kwargs
|
|
||||||
|
|
||||||
|
|
||||||
def no_control(expected_inputs):
|
|
||||||
shapes_dict = {
|
|
||||||
k: v["shape"] for k, v in expected_inputs.items() if k.startswith("additional")
|
|
||||||
}
|
|
||||||
residual_kwargs = {
|
|
||||||
k: torch.zeros(*shape).cpu().numpy().astype(dtype=np.float16)
|
|
||||||
for k, shape in shapes_dict.items()
|
|
||||||
}
|
|
||||||
return residual_kwargs
|
|
||||||
|
|
||||||
|
|
||||||
def chunk_control(cn, target_size):
|
|
||||||
if cn is None:
|
|
||||||
return [None] * target_size
|
|
||||||
|
|
||||||
num_chunks = ceil(cn["output"][0].shape[0] / target_size)
|
|
||||||
|
|
||||||
out = [{"output": [], "middle": []} for _ in range(num_chunks)]
|
|
||||||
|
|
||||||
for k, v in cn.items():
|
|
||||||
for i, x in enumerate(v):
|
|
||||||
chunks = chunk_batch(x, (target_size, *x.shape[1:]))
|
|
||||||
for j, chunk in enumerate(chunks):
|
|
||||||
out[j][k].append(chunk)
|
|
||||||
|
|
||||||
return out
|
|
||||||
|
|||||||
@@ -258,6 +258,7 @@ def convert_unet(
|
|||||||
batch_size: int = 1,
|
batch_size: int = 1,
|
||||||
sample_size: tuple[int, int] = (64, 64),
|
sample_size: tuple[int, int] = (64, 64),
|
||||||
controlnet_support: bool = False,
|
controlnet_support: bool = False,
|
||||||
|
quantize_nbits: str = "none",
|
||||||
):
|
):
|
||||||
coreml_unet = get_unet(model_version, ref_pipe)
|
coreml_unet = get_unet(model_version, ref_pipe)
|
||||||
ref_unet = ref_pipe.unet
|
ref_unet = ref_pipe.unet
|
||||||
@@ -305,6 +306,24 @@ def convert_unet(
|
|||||||
del traced_unet
|
del traced_unet
|
||||||
gc.collect()
|
gc.collect()
|
||||||
|
|
||||||
|
if quantize_nbits != "none":
|
||||||
|
# Opt-in k-means weight palettization. The default path
|
||||||
|
# (quantize_nbits="none") leaves the traced UNet untouched.
|
||||||
|
from coremltools.optimize.coreml import (
|
||||||
|
OpPalettizerConfig,
|
||||||
|
OptimizationConfig,
|
||||||
|
palettize_weights,
|
||||||
|
)
|
||||||
|
|
||||||
|
nbits = int(quantize_nbits)
|
||||||
|
logger.info(f"Palettizing UNet weights to {nbits}-bit (kmeans)..")
|
||||||
|
t0 = time.time()
|
||||||
|
cfg = OptimizationConfig(
|
||||||
|
global_config=OpPalettizerConfig(mode="kmeans", nbits=nbits)
|
||||||
|
)
|
||||||
|
coreml_unet = palettize_weights(coreml_unet, config=cfg)
|
||||||
|
logger.info(f"Palettization took {time.time() - t0:.1f}s")
|
||||||
|
|
||||||
coreml_unet.save(unet_out_path)
|
coreml_unet.save(unet_out_path)
|
||||||
logger.info(f"Saved unet into {unet_out_path}")
|
logger.info(f"Saved unet into {unet_out_path}")
|
||||||
|
|
||||||
@@ -319,6 +338,7 @@ def convert(
|
|||||||
lora_weights: list[tuple[Union[str, os.PathLike], float]] = None,
|
lora_weights: list[tuple[Union[str, os.PathLike], float]] = None,
|
||||||
attn_impl: str = AttentionImplementations.SPLIT_EINSUM.name,
|
attn_impl: str = AttentionImplementations.SPLIT_EINSUM.name,
|
||||||
config_path: str = None,
|
config_path: str = None,
|
||||||
|
quantize_nbits: str = "none",
|
||||||
):
|
):
|
||||||
if os.path.exists(unet_out_path):
|
if os.path.exists(unet_out_path):
|
||||||
logger.info(f"Found existing model at {unet_out_path}! Skipping..")
|
logger.info(f"Found existing model at {unet_out_path}! Skipping..")
|
||||||
@@ -344,6 +364,7 @@ def convert(
|
|||||||
batch_size,
|
batch_size,
|
||||||
sample_size,
|
sample_size,
|
||||||
controlnet_support,
|
controlnet_support,
|
||||||
|
quantize_nbits=quantize_nbits,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,10 @@
|
|||||||
|
"""Framework-free pure-logic core of ComfyUI-CoreMLSuite.
|
||||||
|
|
||||||
|
Modules under this package must NOT import `comfy`, `coremltools`,
|
||||||
|
`python_coreml_stable_diffusion`, `folder_paths`, `nodes`, or any other
|
||||||
|
ComfyUI / Apple runtime. Only `numpy` and `torch` are allowed.
|
||||||
|
|
||||||
|
The thin adapters in `coreml_suite.{latents,controlnet,models}` keep the
|
||||||
|
old public import paths working so `coreml_suite/nodes.py` and downstream
|
||||||
|
ComfyUI workflows are unchanged.
|
||||||
|
"""
|
||||||
@@ -0,0 +1,67 @@
|
|||||||
|
"""Pure helpers around the ControlNet residual inputs of the Core ML UNet.
|
||||||
|
|
||||||
|
Re-exported by coreml_suite.controlnet. Characterization tests cover
|
||||||
|
shapes, dtype (fp16), and zero-fill fallback.
|
||||||
|
"""
|
||||||
|
from itertools import chain
|
||||||
|
from math import ceil
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from coreml_suite.core.latents import chunk_batch
|
||||||
|
|
||||||
|
|
||||||
|
def expand_inputs(inputs):
|
||||||
|
expanded = inputs.copy()
|
||||||
|
for k, v in inputs.items():
|
||||||
|
if isinstance(v, np.ndarray):
|
||||||
|
expanded[k] = np.concatenate([v] * 2) if v.shape[0] == 1 else v
|
||||||
|
elif isinstance(v, torch.Tensor):
|
||||||
|
expanded[k] = torch.cat([v] * 2) if v.shape[0] == 1 else v
|
||||||
|
elif isinstance(v, list):
|
||||||
|
expanded[k] = v * 2 if len(v) == 1 else v
|
||||||
|
elif isinstance(v, dict):
|
||||||
|
expand_inputs(v)
|
||||||
|
return expanded
|
||||||
|
|
||||||
|
|
||||||
|
def extract_residual_kwargs(expected_inputs, control):
|
||||||
|
if "additional_residual_0" not in expected_inputs.keys():
|
||||||
|
return {}
|
||||||
|
if control is None:
|
||||||
|
return no_control(expected_inputs)
|
||||||
|
|
||||||
|
residual_kwargs = {
|
||||||
|
"additional_residual_{}".format(i): r.cpu().numpy().astype(np.float16)
|
||||||
|
for i, r in enumerate(chain(control["output"], control["middle"]))
|
||||||
|
}
|
||||||
|
return residual_kwargs
|
||||||
|
|
||||||
|
|
||||||
|
def no_control(expected_inputs):
|
||||||
|
shapes_dict = {
|
||||||
|
k: v["shape"] for k, v in expected_inputs.items() if k.startswith("additional")
|
||||||
|
}
|
||||||
|
residual_kwargs = {
|
||||||
|
k: torch.zeros(*shape).cpu().numpy().astype(dtype=np.float16)
|
||||||
|
for k, shape in shapes_dict.items()
|
||||||
|
}
|
||||||
|
return residual_kwargs
|
||||||
|
|
||||||
|
|
||||||
|
def chunk_control(cn, target_size):
|
||||||
|
if cn is None:
|
||||||
|
return [None] * target_size
|
||||||
|
|
||||||
|
num_chunks = ceil(cn["output"][0].shape[0] / target_size)
|
||||||
|
|
||||||
|
out = [{"output": [], "middle": []} for _ in range(num_chunks)]
|
||||||
|
|
||||||
|
for k, v in cn.items():
|
||||||
|
for i, x in enumerate(v):
|
||||||
|
chunks = chunk_batch(x, (target_size, *x.shape[1:]))
|
||||||
|
for j, chunk in enumerate(chunks):
|
||||||
|
out[j][k].append(chunk)
|
||||||
|
|
||||||
|
return out
|
||||||
@@ -0,0 +1,113 @@
|
|||||||
|
"""Pure transform from torch sampler inputs to Core ML UNet kwargs.
|
||||||
|
|
||||||
|
Characterization tests cover SD1.5 / SDXL base / SDXL refiner / LCM
|
||||||
|
variants and the chunked-batch fan-out.
|
||||||
|
"""
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from coreml_suite.core.controlnet import extract_residual_kwargs, chunk_control
|
||||||
|
from coreml_suite.core.latents import chunk_batch
|
||||||
|
|
||||||
|
|
||||||
|
class CoreMLInputs:
|
||||||
|
def __init__(self, x, t, context, control, **kwargs):
|
||||||
|
self.x = x
|
||||||
|
self.t = t
|
||||||
|
self.context = context
|
||||||
|
self.control = control
|
||||||
|
self.time_ids = kwargs.get("time_ids")
|
||||||
|
self.text_embeds = kwargs.get("text_embeds")
|
||||||
|
self.ts_cond = kwargs.get("timestep_cond")
|
||||||
|
|
||||||
|
def coreml_kwargs(self, expected_inputs):
|
||||||
|
sample = self.x.cpu().numpy().astype(np.float16)
|
||||||
|
|
||||||
|
context = self.context.cpu().numpy().astype(np.float16)
|
||||||
|
context = context.transpose(0, 2, 1)[:, :, None, :]
|
||||||
|
|
||||||
|
t = self.t.cpu().numpy().astype(np.float16)
|
||||||
|
|
||||||
|
model_input_kwargs = {
|
||||||
|
"sample": sample,
|
||||||
|
"encoder_hidden_states": context,
|
||||||
|
"timestep": t,
|
||||||
|
}
|
||||||
|
residual_kwargs = extract_residual_kwargs(expected_inputs, self.control)
|
||||||
|
model_input_kwargs |= residual_kwargs
|
||||||
|
|
||||||
|
# LCM
|
||||||
|
if self.ts_cond is not None:
|
||||||
|
model_input_kwargs["timestep_cond"] = (
|
||||||
|
self.ts_cond.cpu().numpy().astype(np.float16)
|
||||||
|
)
|
||||||
|
|
||||||
|
# SDXL
|
||||||
|
if "text_embeds" in expected_inputs:
|
||||||
|
model_input_kwargs["text_embeds"] = (
|
||||||
|
self.text_embeds.cpu().numpy().astype(np.float16)
|
||||||
|
)
|
||||||
|
if "time_ids" in expected_inputs:
|
||||||
|
model_input_kwargs["time_ids"] = (
|
||||||
|
self.time_ids.cpu().numpy().astype(np.float16)
|
||||||
|
)
|
||||||
|
|
||||||
|
return model_input_kwargs
|
||||||
|
|
||||||
|
def chunks(self, expected_inputs):
|
||||||
|
sample_shape = expected_inputs["sample"]["shape"]
|
||||||
|
timestep_shape = expected_inputs["timestep"]["shape"]
|
||||||
|
hidden_shape = expected_inputs["encoder_hidden_states"]["shape"]
|
||||||
|
context_shape = (hidden_shape[0], hidden_shape[3], hidden_shape[1])
|
||||||
|
|
||||||
|
chunked_x = chunk_batch(self.x, sample_shape)
|
||||||
|
ts = list(torch.full((len(chunked_x), timestep_shape[0]), self.t[0]))
|
||||||
|
chunked_context = chunk_batch(self.context, context_shape)
|
||||||
|
|
||||||
|
chunked_control = [None] * len(chunked_x)
|
||||||
|
if self.control is not None:
|
||||||
|
chunked_control = chunk_control(self.control, sample_shape[0])
|
||||||
|
|
||||||
|
chunked_ts_cond = [None] * len(chunked_x)
|
||||||
|
if self.ts_cond is not None:
|
||||||
|
ts_cond_shape = expected_inputs["timestep_cond"]["shape"]
|
||||||
|
chunked_ts_cond = chunk_batch(self.ts_cond, ts_cond_shape)
|
||||||
|
|
||||||
|
chunked_time_ids = [None] * len(chunked_x)
|
||||||
|
if expected_inputs.get("time_ids") is not None:
|
||||||
|
time_ids_shape = expected_inputs["time_ids"]["shape"]
|
||||||
|
if self.time_ids is None:
|
||||||
|
self.time_ids = torch.zeros(len(chunked_x), *time_ids_shape[1:]).to(
|
||||||
|
self.x.device
|
||||||
|
)
|
||||||
|
chunked_time_ids = chunk_batch(self.time_ids, time_ids_shape)
|
||||||
|
|
||||||
|
chunked_text_embeds = [None] * len(chunked_x)
|
||||||
|
if expected_inputs.get("text_embeds") is not None:
|
||||||
|
text_embeds_shape = expected_inputs["text_embeds"]["shape"]
|
||||||
|
if self.text_embeds is None:
|
||||||
|
self.text_embeds = torch.zeros(
|
||||||
|
len(chunked_x), *text_embeds_shape[1:]
|
||||||
|
).to(self.x.device)
|
||||||
|
chunked_text_embeds = chunk_batch(self.text_embeds, text_embeds_shape)
|
||||||
|
|
||||||
|
return [
|
||||||
|
CoreMLInputs(
|
||||||
|
x,
|
||||||
|
t,
|
||||||
|
context,
|
||||||
|
control,
|
||||||
|
timestep_cond=ts_cond,
|
||||||
|
time_ids=time_ids,
|
||||||
|
text_embeds=text_embeds,
|
||||||
|
)
|
||||||
|
for x, t, context, control, ts_cond, time_ids, text_embeds in zip(
|
||||||
|
chunked_x,
|
||||||
|
ts,
|
||||||
|
chunked_context,
|
||||||
|
chunked_control,
|
||||||
|
chunked_ts_cond,
|
||||||
|
chunked_time_ids,
|
||||||
|
chunked_text_embeds,
|
||||||
|
)
|
||||||
|
]
|
||||||
@@ -0,0 +1,42 @@
|
|||||||
|
"""Pure batch-chunking helpers for Core ML's fixed-shape UNet inputs.
|
||||||
|
|
||||||
|
Re-exported by coreml_suite.latents. Characterization tests cover the
|
||||||
|
contract (padding-zero regions, truncation in merge_chunks,
|
||||||
|
identity-passthrough when shape already matches).
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
def chunk_batch(input_tensor, target_shape):
|
||||||
|
if input_tensor.shape == target_shape:
|
||||||
|
return [input_tensor]
|
||||||
|
|
||||||
|
batch_size = input_tensor.shape[0]
|
||||||
|
target_batch_size = target_shape[0]
|
||||||
|
|
||||||
|
num_chunks = batch_size // target_batch_size
|
||||||
|
if num_chunks == 0:
|
||||||
|
padding = torch.zeros(target_batch_size - batch_size, *target_shape[1:]).to(
|
||||||
|
input_tensor.device
|
||||||
|
)
|
||||||
|
return [torch.cat((input_tensor, padding), dim=0)]
|
||||||
|
|
||||||
|
mod = batch_size % target_batch_size
|
||||||
|
if mod != 0:
|
||||||
|
chunks = list(torch.chunk(input_tensor[:-mod], num_chunks))
|
||||||
|
padding = torch.zeros(target_batch_size - mod, *target_shape[1:]).to(
|
||||||
|
input_tensor.device
|
||||||
|
)
|
||||||
|
padded = torch.cat((input_tensor[-mod:], padding), dim=0)
|
||||||
|
chunks.append(padded)
|
||||||
|
return chunks
|
||||||
|
|
||||||
|
chunks = list(torch.chunk(input_tensor, num_chunks))
|
||||||
|
return chunks
|
||||||
|
|
||||||
|
|
||||||
|
def merge_chunks(chunks, orig_shape):
|
||||||
|
merged = torch.cat(chunks, dim=0)
|
||||||
|
if merged.shape == orig_shape:
|
||||||
|
return merged
|
||||||
|
return merged[: orig_shape[0]]
|
||||||
@@ -0,0 +1,68 @@
|
|||||||
|
"""Pure out_name composition for the Core ML UNet artifact.
|
||||||
|
|
||||||
|
Extracted from CoreMLConverter.convert so the filename contract
|
||||||
|
can be tested + reused without instantiating the node. The string is the
|
||||||
|
cache key: every workflow that references a converted .mlpackage depends
|
||||||
|
on it staying byte-for-byte identical.
|
||||||
|
"""
|
||||||
|
from typing import Iterable, Tuple
|
||||||
|
|
||||||
|
ATTN_SUFFIX = {
|
||||||
|
"SPLIT_EINSUM": "se",
|
||||||
|
"SPLIT_EINSUM_V2": "se2",
|
||||||
|
"ORIGINAL": "orig",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Palettization bits. "none" = no quantization (default; keeps the
|
||||||
|
# unquantized filename intact so existing workflows still resolve their
|
||||||
|
# cached .mlpackage). Numeric values append a `_q<bits>` suffix.
|
||||||
|
QUANT_NBITS_VALUES = ("none", "8", "6", "4")
|
||||||
|
|
||||||
|
|
||||||
|
def compose_out_name(
|
||||||
|
*,
|
||||||
|
ckpt_name: str,
|
||||||
|
batch_size: int,
|
||||||
|
width: int,
|
||||||
|
height: int,
|
||||||
|
controlnet_support: bool,
|
||||||
|
attention_implementation: str,
|
||||||
|
lora_names: Iterable[str] = (),
|
||||||
|
quantize_nbits: str = "none",
|
||||||
|
) -> str:
|
||||||
|
"""Build the .mlpackage stem from convert() parameters.
|
||||||
|
|
||||||
|
Locked behaviour (characterization tests):
|
||||||
|
- first '.' in ckpt_name wins (`a.b.c.safetensors` -> `a`)
|
||||||
|
- spaces collapse to underscores
|
||||||
|
- LoRA names are taken stem-only, sorted, joined with '_' and
|
||||||
|
prefixed with '_' when present (caller is expected to pass a
|
||||||
|
sorted list; we sort defensively)
|
||||||
|
- controlnet adds `_cn`
|
||||||
|
- attn suffix is `_se` | `_se2` | `_orig`
|
||||||
|
|
||||||
|
Quantization:
|
||||||
|
- quantize_nbits "none" (default) appends nothing — existing
|
||||||
|
unquantized .mlpackages keep the old filename
|
||||||
|
- "4" / "6" / "8" appends `_q<bits>` after the attn suffix
|
||||||
|
"""
|
||||||
|
if quantize_nbits not in QUANT_NBITS_VALUES:
|
||||||
|
raise ValueError(
|
||||||
|
f"quantize_nbits={quantize_nbits!r} not in {QUANT_NBITS_VALUES}"
|
||||||
|
)
|
||||||
|
stem = ckpt_name.split(".")[0]
|
||||||
|
sorted_names = sorted(lora_names)
|
||||||
|
lora_str = "_" + "_".join(name.split(".")[0] for name in sorted_names) if sorted_names else ""
|
||||||
|
cn_suffix = "_cn" if controlnet_support else ""
|
||||||
|
attn_suffix = "_" + ATTN_SUFFIX[attention_implementation]
|
||||||
|
quant_suffix = f"_q{quantize_nbits}" if quantize_nbits != "none" else ""
|
||||||
|
out_name = (
|
||||||
|
f"{stem}{lora_str}_{batch_size}x{width}x{height}"
|
||||||
|
f"{cn_suffix}{attn_suffix}{quant_suffix}"
|
||||||
|
)
|
||||||
|
return out_name.replace(" ", "_")
|
||||||
|
|
||||||
|
|
||||||
|
def lora_names_from_params(lora_params: Iterable[Tuple[str, float]]) -> list[str]:
|
||||||
|
"""Mirror the sort applied inside CoreMLConverter.convert."""
|
||||||
|
return [name for name, _ in sorted(lora_params, key=lambda pair: pair[0])]
|
||||||
@@ -0,0 +1,91 @@
|
|||||||
|
"""Pure SDXL detection + time_ids/text_embeds assembly.
|
||||||
|
|
||||||
|
The framework-coupled adapter `add_sdxl_model_options` lives in models.py
|
||||||
|
and delegates the math here. Characterization tests cover base (len 6) vs
|
||||||
|
refiner (len 5) and the closure free-vars produced by
|
||||||
|
`sdxl_model_function_wrapper`.
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
def is_sdxl(coreml_model):
|
||||||
|
return (
|
||||||
|
"time_ids" in coreml_model.expected_inputs
|
||||||
|
and "text_embeds" in coreml_model.expected_inputs
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def is_sdxl_base(coreml_model):
|
||||||
|
return (
|
||||||
|
is_sdxl(coreml_model)
|
||||||
|
and coreml_model.expected_inputs["time_ids"]["shape"][1] == 6
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def is_sdxl_refiner(coreml_model):
|
||||||
|
return (
|
||||||
|
is_sdxl(coreml_model)
|
||||||
|
and coreml_model.expected_inputs["time_ids"]["shape"][1] == 5
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def build_sdxl_time_ids(pos_dict, neg_dict, *, is_base: bool, is_refiner: bool):
|
||||||
|
"""Compose the (2, N) time_ids tensor for the SDXL Core ML UNet.
|
||||||
|
|
||||||
|
- base: N=6 -> [h, w, crop_h, crop_w, target_h, target_w]
|
||||||
|
- refiner: N=5 -> [h, w, crop_h, crop_w, aesthetic_score]
|
||||||
|
- neither: N=4 -> [h, w, crop_h, crop_w] (edge case kept for parity)
|
||||||
|
"""
|
||||||
|
pos_time_ids = [
|
||||||
|
pos_dict.get("height", 768),
|
||||||
|
pos_dict.get("width", 768),
|
||||||
|
pos_dict.get("crop_h", 0),
|
||||||
|
pos_dict.get("crop_w", 0),
|
||||||
|
]
|
||||||
|
neg_time_ids = [
|
||||||
|
neg_dict.get("height", 768),
|
||||||
|
neg_dict.get("width", 768),
|
||||||
|
neg_dict.get("crop_h", 0),
|
||||||
|
neg_dict.get("crop_w", 0),
|
||||||
|
]
|
||||||
|
|
||||||
|
if is_base:
|
||||||
|
pos_time_ids += [
|
||||||
|
pos_dict.get("target_height", 768),
|
||||||
|
pos_dict.get("target_width", 768),
|
||||||
|
]
|
||||||
|
neg_time_ids += [
|
||||||
|
neg_dict.get("target_height", 768),
|
||||||
|
neg_dict.get("target_width", 768),
|
||||||
|
]
|
||||||
|
|
||||||
|
if is_refiner:
|
||||||
|
pos_time_ids += [pos_dict.get("aesthetic_score", 6)]
|
||||||
|
neg_time_ids += [neg_dict.get("aesthetic_score", 2.5)]
|
||||||
|
|
||||||
|
return torch.tensor([pos_time_ids, neg_time_ids])
|
||||||
|
|
||||||
|
|
||||||
|
def build_sdxl_text_embeds(pos_pooled, neg_pooled):
|
||||||
|
"""Concat pos then neg along the batch dim. Locked contract."""
|
||||||
|
return torch.cat((pos_pooled, neg_pooled))
|
||||||
|
|
||||||
|
|
||||||
|
def sdxl_model_function_wrapper(time_ids, text_embeds, refiner=False):
|
||||||
|
def wrapper(model_function, params):
|
||||||
|
x = params["input"]
|
||||||
|
t = params["timestep"]
|
||||||
|
c = params["c"]
|
||||||
|
|
||||||
|
context = c.get("c_crossattn")
|
||||||
|
|
||||||
|
if context is None:
|
||||||
|
return torch.zeros_like(x)
|
||||||
|
|
||||||
|
if refiner and context is not None:
|
||||||
|
# converted refiner accepts only g clip
|
||||||
|
c["c_crossattn"] = context[:, :, 768:]
|
||||||
|
|
||||||
|
return model_function(x, t, **c, time_ids=time_ids, text_embeds=text_embeds)
|
||||||
|
|
||||||
|
return wrapper
|
||||||
+3
-35
@@ -1,36 +1,4 @@
|
|||||||
import torch
|
"""Compatibility shim — re-exports from coreml_suite.core.latents."""
|
||||||
|
from coreml_suite.core.latents import chunk_batch, merge_chunks
|
||||||
|
|
||||||
|
__all__ = ["chunk_batch", "merge_chunks"]
|
||||||
def chunk_batch(input_tensor, target_shape):
|
|
||||||
if input_tensor.shape == target_shape:
|
|
||||||
return [input_tensor]
|
|
||||||
|
|
||||||
batch_size = input_tensor.shape[0]
|
|
||||||
target_batch_size = target_shape[0]
|
|
||||||
|
|
||||||
num_chunks = batch_size // target_batch_size
|
|
||||||
if num_chunks == 0:
|
|
||||||
padding = torch.zeros(target_batch_size - batch_size, *target_shape[1:]).to(
|
|
||||||
input_tensor.device
|
|
||||||
)
|
|
||||||
return [torch.cat((input_tensor, padding), dim=0)]
|
|
||||||
|
|
||||||
mod = batch_size % target_batch_size
|
|
||||||
if mod != 0:
|
|
||||||
chunks = list(torch.chunk(input_tensor[:-mod], num_chunks))
|
|
||||||
padding = torch.zeros(target_batch_size - mod, *target_shape[1:]).to(
|
|
||||||
input_tensor.device
|
|
||||||
)
|
|
||||||
padded = torch.cat((input_tensor[-mod:], padding), dim=0)
|
|
||||||
chunks.append(padded)
|
|
||||||
return chunks
|
|
||||||
|
|
||||||
chunks = list(torch.chunk(input_tensor, num_chunks))
|
|
||||||
return chunks
|
|
||||||
|
|
||||||
|
|
||||||
def merge_chunks(chunks, orig_shape):
|
|
||||||
merged = torch.cat(chunks, dim=0)
|
|
||||||
if merged.shape == orig_shape:
|
|
||||||
return merged
|
|
||||||
return merged[: orig_shape[0]]
|
|
||||||
|
|||||||
+40
-188
@@ -1,15 +1,44 @@
|
|||||||
import numpy as np
|
"""Framework-coupled glue between Core ML UNets and ComfyUI's sampler stack.
|
||||||
|
|
||||||
|
Pure math (CoreMLInputs, SDXL detection, time_ids/text_embeds assembly,
|
||||||
|
sdxl_model_function_wrapper) lives in coreml_suite.core.*.
|
||||||
|
This module is what touches comfy.*: model_base, ModelPatcher, the
|
||||||
|
diffusion_model wrapper, and the maintainer-facing add_sdxl_model_options
|
||||||
|
adapter.
|
||||||
|
"""
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from comfy import model_base
|
from comfy import model_base
|
||||||
from comfy.model_management import get_torch_device
|
from comfy.model_management import get_torch_device
|
||||||
from comfy.model_patcher import ModelPatcher
|
from comfy.model_patcher import ModelPatcher
|
||||||
|
|
||||||
from coreml_suite.config import get_model_config, ModelVersion
|
from coreml_suite.config import get_model_config, ModelVersion
|
||||||
from coreml_suite.controlnet import extract_residual_kwargs, chunk_control
|
from coreml_suite.core.inputs import CoreMLInputs
|
||||||
from coreml_suite.latents import chunk_batch, merge_chunks
|
from coreml_suite.core.latents import merge_chunks
|
||||||
|
from coreml_suite.core.sdxl import (
|
||||||
|
build_sdxl_text_embeds,
|
||||||
|
build_sdxl_time_ids,
|
||||||
|
is_sdxl,
|
||||||
|
is_sdxl_base,
|
||||||
|
is_sdxl_refiner,
|
||||||
|
sdxl_model_function_wrapper,
|
||||||
|
)
|
||||||
from coreml_suite.lcm.utils import is_lcm
|
from coreml_suite.lcm.utils import is_lcm
|
||||||
from coreml_suite.logger import logger
|
from coreml_suite.logger import logger
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"CoreMLInputs",
|
||||||
|
"CoreMLModelWrapper",
|
||||||
|
"CoreMLModelWrapperLCM",
|
||||||
|
"add_sdxl_model_options",
|
||||||
|
"get_latent_image",
|
||||||
|
"get_model_patcher",
|
||||||
|
"is_sdxl",
|
||||||
|
"is_sdxl_base",
|
||||||
|
"is_sdxl_refiner",
|
||||||
|
"sdxl_model_function_wrapper",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
class CoreMLModelWrapper:
|
class CoreMLModelWrapper:
|
||||||
def __init__(self, coreml_model):
|
def __init__(self, coreml_model):
|
||||||
@@ -68,204 +97,27 @@ class CoreMLModelWrapperLCM(CoreMLModelWrapper):
|
|||||||
self.config = None
|
self.config = None
|
||||||
|
|
||||||
|
|
||||||
class CoreMLInputs:
|
|
||||||
def __init__(self, x, t, context, control, **kwargs):
|
|
||||||
self.x = x
|
|
||||||
self.t = t
|
|
||||||
self.context = context
|
|
||||||
self.control = control
|
|
||||||
self.time_ids = kwargs.get("time_ids")
|
|
||||||
self.text_embeds = kwargs.get("text_embeds")
|
|
||||||
self.ts_cond = kwargs.get("timestep_cond")
|
|
||||||
|
|
||||||
def coreml_kwargs(self, expected_inputs):
|
|
||||||
sample = self.x.cpu().numpy().astype(np.float16)
|
|
||||||
|
|
||||||
context = self.context.cpu().numpy().astype(np.float16)
|
|
||||||
context = context.transpose(0, 2, 1)[:, :, None, :]
|
|
||||||
|
|
||||||
t = self.t.cpu().numpy().astype(np.float16)
|
|
||||||
|
|
||||||
model_input_kwargs = {
|
|
||||||
"sample": sample,
|
|
||||||
"encoder_hidden_states": context,
|
|
||||||
"timestep": t,
|
|
||||||
}
|
|
||||||
residual_kwargs = extract_residual_kwargs(expected_inputs, self.control)
|
|
||||||
model_input_kwargs |= residual_kwargs
|
|
||||||
|
|
||||||
# LCM
|
|
||||||
if self.ts_cond is not None:
|
|
||||||
model_input_kwargs["timestep_cond"] = (
|
|
||||||
self.ts_cond.cpu().numpy().astype(np.float16)
|
|
||||||
)
|
|
||||||
|
|
||||||
# SDXL
|
|
||||||
if "text_embeds" in expected_inputs:
|
|
||||||
model_input_kwargs["text_embeds"] = (
|
|
||||||
self.text_embeds.cpu().numpy().astype(np.float16)
|
|
||||||
)
|
|
||||||
if "time_ids" in expected_inputs:
|
|
||||||
model_input_kwargs["time_ids"] = (
|
|
||||||
self.time_ids.cpu().numpy().astype(np.float16)
|
|
||||||
)
|
|
||||||
|
|
||||||
return model_input_kwargs
|
|
||||||
|
|
||||||
def chunks(self, expected_inputs):
|
|
||||||
sample_shape = expected_inputs["sample"]["shape"]
|
|
||||||
timestep_shape = expected_inputs["timestep"]["shape"]
|
|
||||||
hidden_shape = expected_inputs["encoder_hidden_states"]["shape"]
|
|
||||||
context_shape = (hidden_shape[0], hidden_shape[3], hidden_shape[1])
|
|
||||||
|
|
||||||
chunked_x = chunk_batch(self.x, sample_shape)
|
|
||||||
ts = list(torch.full((len(chunked_x), timestep_shape[0]), self.t[0]))
|
|
||||||
chunked_context = chunk_batch(self.context, context_shape)
|
|
||||||
|
|
||||||
chunked_control = [None] * len(chunked_x)
|
|
||||||
if self.control is not None:
|
|
||||||
chunked_control = chunk_control(self.control, sample_shape[0])
|
|
||||||
|
|
||||||
chunked_ts_cond = [None] * len(chunked_x)
|
|
||||||
if self.ts_cond is not None:
|
|
||||||
ts_cond_shape = expected_inputs["timestep_cond"]["shape"]
|
|
||||||
chunked_ts_cond = chunk_batch(self.ts_cond, ts_cond_shape)
|
|
||||||
|
|
||||||
chunked_time_ids = [None] * len(chunked_x)
|
|
||||||
if expected_inputs.get("time_ids") is not None:
|
|
||||||
time_ids_shape = expected_inputs["time_ids"]["shape"]
|
|
||||||
if self.time_ids is None:
|
|
||||||
self.time_ids = torch.zeros(len(chunked_x), *time_ids_shape[1:]).to(
|
|
||||||
self.x.device
|
|
||||||
)
|
|
||||||
chunked_time_ids = chunk_batch(self.time_ids, time_ids_shape)
|
|
||||||
|
|
||||||
chunked_text_embeds = [None] * len(chunked_x)
|
|
||||||
if expected_inputs.get("text_embeds") is not None:
|
|
||||||
text_embeds_shape = expected_inputs["text_embeds"]["shape"]
|
|
||||||
if self.text_embeds is None:
|
|
||||||
self.text_embeds = torch.zeros(
|
|
||||||
len(chunked_x), *text_embeds_shape[1:]
|
|
||||||
).to(self.x.device)
|
|
||||||
chunked_text_embeds = chunk_batch(self.text_embeds, text_embeds_shape)
|
|
||||||
|
|
||||||
return [
|
|
||||||
CoreMLInputs(
|
|
||||||
x,
|
|
||||||
t,
|
|
||||||
context,
|
|
||||||
control,
|
|
||||||
timestep_cond=ts_cond,
|
|
||||||
time_ids=time_ids,
|
|
||||||
text_embeds=text_embeds,
|
|
||||||
)
|
|
||||||
for x, t, context, control, ts_cond, time_ids, text_embeds in zip(
|
|
||||||
chunked_x,
|
|
||||||
ts,
|
|
||||||
chunked_context,
|
|
||||||
chunked_control,
|
|
||||||
chunked_ts_cond,
|
|
||||||
chunked_time_ids,
|
|
||||||
chunked_text_embeds,
|
|
||||||
)
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
def is_sdxl(coreml_model):
|
|
||||||
return (
|
|
||||||
"time_ids" in coreml_model.expected_inputs
|
|
||||||
and "text_embeds" in coreml_model.expected_inputs
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def is_sdxl_base(coreml_model):
|
|
||||||
return (
|
|
||||||
is_sdxl(coreml_model)
|
|
||||||
and coreml_model.expected_inputs["time_ids"]["shape"][1] == 6
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def is_sdxl_refiner(coreml_model):
|
|
||||||
return (
|
|
||||||
is_sdxl(coreml_model)
|
|
||||||
and coreml_model.expected_inputs["time_ids"]["shape"][1] == 5
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def sdxl_model_function_wrapper(time_ids, text_embeds, refiner=False):
|
|
||||||
def wrapper(model_function, params):
|
|
||||||
x = params["input"]
|
|
||||||
t = params["timestep"]
|
|
||||||
c = params["c"]
|
|
||||||
|
|
||||||
context = c.get("c_crossattn")
|
|
||||||
|
|
||||||
if context is None:
|
|
||||||
return torch.zeros_like(x)
|
|
||||||
|
|
||||||
if refiner and context is not None:
|
|
||||||
# converted refiner accepts only g clip
|
|
||||||
c["c_crossattn"] = context[:, :, 768:]
|
|
||||||
|
|
||||||
return model_function(x, t, **c, time_ids=time_ids, text_embeds=text_embeds)
|
|
||||||
|
|
||||||
return wrapper
|
|
||||||
|
|
||||||
|
|
||||||
def add_sdxl_model_options(model_patcher, positive, negative):
|
def add_sdxl_model_options(model_patcher, positive, negative):
|
||||||
mp = model_patcher.clone()
|
mp = model_patcher.clone()
|
||||||
|
|
||||||
pos_dict = positive[0][1]
|
pos_dict = positive[0][1]
|
||||||
neg_dict = negative[0][1]
|
neg_dict = negative[0][1]
|
||||||
|
|
||||||
pos_pooled = pos_dict["pooled_output"]
|
is_base = model_patcher.model.diffusion_model.is_sdxl_base
|
||||||
neg_pooled = neg_dict["pooled_output"]
|
|
||||||
|
|
||||||
pos_time_ids = [
|
|
||||||
pos_dict.get("height", 768),
|
|
||||||
pos_dict.get("width", 768),
|
|
||||||
pos_dict.get("crop_h", 0),
|
|
||||||
pos_dict.get("crop_w", 0),
|
|
||||||
]
|
|
||||||
|
|
||||||
neg_time_ids = [
|
|
||||||
neg_dict.get("height", 768),
|
|
||||||
neg_dict.get("width", 768),
|
|
||||||
neg_dict.get("crop_h", 0),
|
|
||||||
neg_dict.get("crop_w", 0),
|
|
||||||
]
|
|
||||||
|
|
||||||
if model_patcher.model.diffusion_model.is_sdxl_base:
|
|
||||||
pos_time_ids += [
|
|
||||||
pos_dict.get("target_height", 768),
|
|
||||||
pos_dict.get("target_width", 768),
|
|
||||||
]
|
|
||||||
|
|
||||||
neg_time_ids += [
|
|
||||||
neg_dict.get("target_height", 768),
|
|
||||||
neg_dict.get("target_width", 768),
|
|
||||||
]
|
|
||||||
|
|
||||||
is_refiner = model_patcher.model.diffusion_model.is_sdxl_refiner
|
is_refiner = model_patcher.model.diffusion_model.is_sdxl_refiner
|
||||||
if is_refiner:
|
|
||||||
pos_time_ids += [
|
|
||||||
pos_dict.get("aesthetic_score", 6),
|
|
||||||
]
|
|
||||||
|
|
||||||
neg_time_ids += [
|
time_ids = build_sdxl_time_ids(
|
||||||
neg_dict.get("aesthetic_score", 2.5),
|
pos_dict, neg_dict, is_base=is_base, is_refiner=is_refiner
|
||||||
]
|
)
|
||||||
|
text_embeds = build_sdxl_text_embeds(
|
||||||
|
pos_dict["pooled_output"], neg_dict["pooled_output"]
|
||||||
|
)
|
||||||
|
|
||||||
time_ids = torch.tensor([pos_time_ids, neg_time_ids])
|
mp.model_options |= {
|
||||||
text_embeds = torch.cat((pos_pooled, neg_pooled))
|
|
||||||
|
|
||||||
model_options = {
|
|
||||||
"model_function_wrapper": sdxl_model_function_wrapper(
|
"model_function_wrapper": sdxl_model_function_wrapper(
|
||||||
time_ids, text_embeds, is_refiner
|
time_ids, text_embeds, is_refiner
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
mp.model_options |= model_options
|
|
||||||
|
|
||||||
return mp
|
return mp
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+22
-16
@@ -8,6 +8,11 @@ import folder_paths
|
|||||||
from coreml_suite import COREML_NODE
|
from coreml_suite import COREML_NODE
|
||||||
from coreml_suite import converter
|
from coreml_suite import converter
|
||||||
from coreml_suite.config import ModelVersion
|
from coreml_suite.config import ModelVersion
|
||||||
|
from coreml_suite.core.naming import (
|
||||||
|
QUANT_NBITS_VALUES,
|
||||||
|
compose_out_name,
|
||||||
|
lora_names_from_params,
|
||||||
|
)
|
||||||
from coreml_suite.lcm.utils import add_lcm_model_options, lcm_patch, is_lcm
|
from coreml_suite.lcm.utils import add_lcm_model_options, lcm_patch, is_lcm
|
||||||
from coreml_suite.logger import logger
|
from coreml_suite.logger import logger
|
||||||
from nodes import KSampler, LoraLoader, KSamplerAdvanced
|
from nodes import KSampler, LoraLoader, KSamplerAdvanced
|
||||||
@@ -244,6 +249,12 @@ class CoreMLConverter(COREML_NODE):
|
|||||||
"controlnet_support": ("BOOLEAN", {"default": False}),
|
"controlnet_support": ("BOOLEAN", {"default": False}),
|
||||||
},
|
},
|
||||||
"optional": {
|
"optional": {
|
||||||
|
# k-means weight palettization. Kept optional so workflows
|
||||||
|
# that omit it still validate — ComfyUI rejects a prompt that
|
||||||
|
# omits any `required` input. When omitted it defaults to
|
||||||
|
# "none", identical to unquantized behavior and filename, so
|
||||||
|
# existing cached .mlpackages still resolve.
|
||||||
|
"quantize_nbits": (list(QUANT_NBITS_VALUES), {"default": "none"}),
|
||||||
"lora_params": ("LORA_PARAMS",),
|
"lora_params": ("LORA_PARAMS",),
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -262,6 +273,7 @@ class CoreMLConverter(COREML_NODE):
|
|||||||
attention_implementation,
|
attention_implementation,
|
||||||
compute_unit,
|
compute_unit,
|
||||||
controlnet_support,
|
controlnet_support,
|
||||||
|
quantize_nbits="none",
|
||||||
lora_params=None,
|
lora_params=None,
|
||||||
):
|
):
|
||||||
"""Converts a LCM model to Core ML.
|
"""Converts a LCM model to Core ML.
|
||||||
@@ -288,24 +300,17 @@ class CoreMLConverter(COREML_NODE):
|
|||||||
h = height
|
h = height
|
||||||
w = width
|
w = width
|
||||||
sample_size = (h // 8, w // 8)
|
sample_size = (h // 8, w // 8)
|
||||||
batch_size = batch_size
|
out_name = compose_out_name(
|
||||||
cn_support_str = "_cn" if controlnet_support else ""
|
ckpt_name=ckpt_name,
|
||||||
lora_str = (
|
batch_size=batch_size,
|
||||||
"_" + "_".join(lora_param[0].split(".")[0] for lora_param in lora_params)
|
width=w,
|
||||||
if lora_params
|
height=h,
|
||||||
else ""
|
controlnet_support=controlnet_support,
|
||||||
|
attention_implementation=attention_implementation,
|
||||||
|
lora_names=lora_names_from_params(lora_params),
|
||||||
|
quantize_nbits=quantize_nbits,
|
||||||
)
|
)
|
||||||
|
|
||||||
attn_str = (
|
|
||||||
"_"
|
|
||||||
+ {"SPLIT_EINSUM": "se", "SPLIT_EINSUM_V2": "se2", "ORIGINAL": "orig"}[
|
|
||||||
attention_implementation
|
|
||||||
]
|
|
||||||
)
|
|
||||||
|
|
||||||
out_name = f"{ckpt_name.split('.')[0]}{lora_str}_{batch_size}x{w}x{h}{cn_support_str}{attn_str}"
|
|
||||||
out_name = out_name.replace(" ", "_")
|
|
||||||
|
|
||||||
logger.info(f"Converting {ckpt_name} to {out_name}")
|
logger.info(f"Converting {ckpt_name} to {out_name}")
|
||||||
logger.info(f"Batch size: {batch_size}")
|
logger.info(f"Batch size: {batch_size}")
|
||||||
logger.info(f"Width: {w}, Height: {h}")
|
logger.info(f"Width: {w}, Height: {h}")
|
||||||
@@ -335,6 +340,7 @@ class CoreMLConverter(COREML_NODE):
|
|||||||
lora_weights=lora_weights,
|
lora_weights=lora_weights,
|
||||||
attn_impl=attention_implementation,
|
attn_impl=attention_implementation,
|
||||||
config_path=config_path,
|
config_path=config_path,
|
||||||
|
quantize_nbits=quantize_nbits,
|
||||||
)
|
)
|
||||||
unet_target_path = converter.compile_model(
|
unet_target_path = converter.compile_model(
|
||||||
out_path=unet_out_path, out_name=out_name, submodule_name="unet"
|
out_path=unet_out_path, out_name=out_name, submodule_name="unet"
|
||||||
|
|||||||
+81
-2
@@ -2,8 +2,26 @@
|
|||||||
name = "comfyui-coremlsuite"
|
name = "comfyui-coremlsuite"
|
||||||
description = "This extension contains a set of custom nodes for ComfyUI that allow you to use Core ML models in your ComfyUI workflows."
|
description = "This extension contains a set of custom nodes for ComfyUI that allow you to use Core ML models in your ComfyUI workflows."
|
||||||
version = "1.0.1"
|
version = "1.0.1"
|
||||||
license = { file = "LICENSE" }
|
license = "MIT"
|
||||||
dependencies = ["git+https://github.com/apple/ml-stable-diffusion.git", "coremltools>=7.1", "overrides", "diffusers>=0.22", "peft>=0.6.2", "omegaconf>=2.3"]
|
requires-python = ">=3.12,<3.13"
|
||||||
|
packages = [{ include = "coreml_suite" }]
|
||||||
|
dependencies = [
|
||||||
|
# Python 3.12, coremltools 9, torch 2.7.
|
||||||
|
# numpy stays in the 1.24..1.x range — none of our modules need
|
||||||
|
# numpy 2, and coremltools + ml-stable-diffusion's SD UNet trace
|
||||||
|
# hit hard bugs under numpy 2 (`_cast` int(ndarray) strictness and
|
||||||
|
# `view` mixed-Var shape lists).
|
||||||
|
# torch 2.7 is the latest version coremltools 9's PyTorch frontend
|
||||||
|
# has been tested against.
|
||||||
|
"python-coreml-stable-diffusion @ git+https://github.com/apple/ml-stable-diffusion.git@e5d960c41a6a4ab200b8db379194127607b1c590",
|
||||||
|
"torch>=2.7,<2.8",
|
||||||
|
"coremltools>=9,<10",
|
||||||
|
"numpy>=1.24,<2",
|
||||||
|
"overrides",
|
||||||
|
"diffusers>=0.22",
|
||||||
|
"peft>=0.6.2",
|
||||||
|
"omegaconf>=2.3",
|
||||||
|
]
|
||||||
|
|
||||||
[project.urls]
|
[project.urls]
|
||||||
Repository = "https://github.com/aszc-dev/ComfyUI-CoreMLSuite"
|
Repository = "https://github.com/aszc-dev/ComfyUI-CoreMLSuite"
|
||||||
@@ -13,3 +31,64 @@ Repository = "https://github.com/aszc-dev/ComfyUI-CoreMLSuite"
|
|||||||
PublisherId = "aszc-dev"
|
PublisherId = "aszc-dev"
|
||||||
DisplayName = "ComfyUI-CoreMLSuite"
|
DisplayName = "ComfyUI-CoreMLSuite"
|
||||||
Icon = ""
|
Icon = ""
|
||||||
|
# Pinned to the ComfyUI commit this toolchain was validated against.
|
||||||
|
requires-comfyui = "==ab5413351eee61f3d7f10c74e75286df0058bb18"
|
||||||
|
|
||||||
|
[dependency-groups]
|
||||||
|
dev = [
|
||||||
|
"pillow>=12.2.0",
|
||||||
|
"psutil>=7.2.2",
|
||||||
|
]
|
||||||
|
# ComfyUI runtime deps that aren't part of our package's runtime contract
|
||||||
|
# but are needed to spin up the ComfyUI server for Tier 2 integration tests.
|
||||||
|
# Kept in a uv group so `uv sync --group comfy` brings them in without
|
||||||
|
# polluting the published metadata (and without re-bumping our torch pin
|
||||||
|
# via `uv pip install -r ComfyUI/requirements.txt`, which would float to
|
||||||
|
# the latest torch and break the coremltools 9 compatibility ceiling).
|
||||||
|
comfy = [
|
||||||
|
"comfyui-frontend-package==1.14.6",
|
||||||
|
"torchvision",
|
||||||
|
"torchaudio",
|
||||||
|
"torchsde",
|
||||||
|
"einops",
|
||||||
|
"tokenizers>=0.13.3",
|
||||||
|
"safetensors>=0.4.2",
|
||||||
|
"aiohttp>=3.11.8",
|
||||||
|
"yarl>=1.18.0",
|
||||||
|
"kornia>=0.7.1",
|
||||||
|
"spandrel",
|
||||||
|
"soundfile",
|
||||||
|
"sentencepiece",
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.uv]
|
||||||
|
# ml-stable-diffusion's setup.py hard-pins numpy<1.24, diffusers==0.30.2
|
||||||
|
# and transformers==4.44.2, which blocks the modern torch / coremltools
|
||||||
|
# combo on Python 3.12. Override the four blocking pins; the .unet /
|
||||||
|
# .coreml_model symbols we actually import are stable across the bumped
|
||||||
|
# versions.
|
||||||
|
override-dependencies = [
|
||||||
|
"numpy>=1.24,<2",
|
||||||
|
"diffusers>=0.30",
|
||||||
|
"transformers>=4.44",
|
||||||
|
"huggingface-hub>=0.24",
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.pytest.ini_options]
|
||||||
|
# Tier markers gate which environment a test needs.
|
||||||
|
# - unit: framework-free pure-logic tests (Tier 0; run without ComfyUI on
|
||||||
|
# Linux).
|
||||||
|
# - m2: needs an Apple Silicon Mac with the Neural Engine (Tier 2),
|
||||||
|
# typically a self-hosted runner or local M-series box.
|
||||||
|
# - smoke: lightweight checks that need Apple Silicon + coremltools but no
|
||||||
|
# ANE/real model (Tier 1).
|
||||||
|
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"]
|
||||||
|
# importlib mode keeps pytest from importing the repo-root __init__.py
|
||||||
|
# (which is the ComfyUI custom-node entry and pulls in comfy + nodes).
|
||||||
|
# Without this Tier-0 leaks the entire ComfyUI runtime on collection.
|
||||||
|
addopts = ["--import-mode=importlib", "--confcutdir=tests"]
|
||||||
|
|||||||
+4
-2
@@ -1,5 +1,7 @@
|
|||||||
git+https://github.com/apple/ml-stable-diffusion.git
|
git+https://github.com/apple/ml-stable-diffusion.git@e5d960c41a6a4ab200b8db379194127607b1c590
|
||||||
coremltools>=7.1
|
torch>=2.7,<2.8
|
||||||
|
coremltools==8.2
|
||||||
|
numpy>=2,<3
|
||||||
overrides
|
overrides
|
||||||
diffusers>=0.22
|
diffusers>=0.22
|
||||||
peft>=0.6.2
|
peft>=0.6.2
|
||||||
|
|||||||
@@ -0,0 +1,61 @@
|
|||||||
|
"""Pytest bootstrap for ComfyUI-CoreMLSuite tests.
|
||||||
|
|
||||||
|
- Adds the ComfyUI checkout to sys.path so the framework-coupled modules
|
||||||
|
that transitively import `comfy.*` resolve when pytest is invoked from
|
||||||
|
this package's root.
|
||||||
|
- Auto-applies tier markers based on the directory a test lives in, so
|
||||||
|
individual files don't have to repeat @pytest.mark.unit / .m2.
|
||||||
|
"""
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
_TIER_BY_DIR = {
|
||||||
|
"tests/unit": "unit",
|
||||||
|
"tests/m2": "m2",
|
||||||
|
"tests/integration": "m2",
|
||||||
|
"tests/smoke": "smoke",
|
||||||
|
}
|
||||||
|
|
||||||
|
# When the user asks for a single tier (-m unit / -m m2), skip the other
|
||||||
|
# directories at collection time. Tier-0 cannot afford to import tests/m2
|
||||||
|
# files because they pull in PIL + ComfyUI runtime which Linux CI won't have.
|
||||||
|
_TIER_DIRS = {
|
||||||
|
"unit": ("/tests/unit/",),
|
||||||
|
"m2": ("/tests/m2/", "/tests/integration/"),
|
||||||
|
"smoke": ("/tests/smoke/",),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def pytest_ignore_collect(collection_path, config):
|
||||||
|
expr = config.option.markexpr
|
||||||
|
if expr not in _TIER_DIRS:
|
||||||
|
return None
|
||||||
|
allowed = _TIER_DIRS[expr]
|
||||||
|
rel = str(collection_path).replace("\\", "/")
|
||||||
|
if "/tests/" not in rel:
|
||||||
|
return None
|
||||||
|
# Always allow tests/ root + the tier's own dirs.
|
||||||
|
if rel.endswith("/tests"):
|
||||||
|
return None
|
||||||
|
if any(frag in rel + "/" for frag in allowed):
|
||||||
|
return None
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
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
|
||||||
@@ -1,85 +0,0 @@
|
|||||||
import json
|
|
||||||
import os
|
|
||||||
|
|
||||||
import pytest
|
|
||||||
import requests
|
|
||||||
import torch
|
|
||||||
from PIL import Image
|
|
||||||
import numpy as np
|
|
||||||
|
|
||||||
from folder_paths import get_save_image_path, get_output_directory
|
|
||||||
|
|
||||||
IMAGE_PREFIX = "E2E-1.5"
|
|
||||||
IMAGE_PREFIX_CML = f"{IMAGE_PREFIX}-CoreML"
|
|
||||||
IMAGE_PREFIX_MPS = f"{IMAGE_PREFIX}-MPS"
|
|
||||||
|
|
||||||
|
|
||||||
class OutputImageRepository:
|
|
||||||
def __init__(self, name_prefix):
|
|
||||||
self.name_prefix = name_prefix
|
|
||||||
|
|
||||||
def list_images(self):
|
|
||||||
full_output_folder, _, _, _, _ = get_save_image_path(
|
|
||||||
self.name_prefix, get_output_directory(), 512, 512
|
|
||||||
)
|
|
||||||
return full_output_folder, os.listdir(full_output_folder)
|
|
||||||
|
|
||||||
def delete_images(self):
|
|
||||||
full_output_folder, images = self.list_images()
|
|
||||||
for image in images:
|
|
||||||
os.remove(os.path.join(full_output_folder, image))
|
|
||||||
|
|
||||||
def get_latest_image(self, prefix):
|
|
||||||
full_output_folder, images = self.list_images()
|
|
||||||
for image in sorted(images, reverse=True):
|
|
||||||
if image.startswith(prefix):
|
|
||||||
return os.path.join(full_output_folder, image)
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="function")
|
|
||||||
def output_image_repository():
|
|
||||||
repo = OutputImageRepository(IMAGE_PREFIX)
|
|
||||||
yield repo
|
|
||||||
repo.delete_images()
|
|
||||||
|
|
||||||
|
|
||||||
def test_basic_conversion_1_5(output_image_repository):
|
|
||||||
with open("tests/integration/workflows/e2e-1.5-basic-conversion.json") as f:
|
|
||||||
prompt = json.load(f)
|
|
||||||
prompt = randomize_seed_in_prompt(prompt)
|
|
||||||
queue_prompt(prompt)
|
|
||||||
|
|
||||||
coreml_img_path = output_image_repository.get_latest_image(IMAGE_PREFIX_CML)
|
|
||||||
mps_img_path = output_image_repository.get_latest_image(IMAGE_PREFIX_MPS)
|
|
||||||
|
|
||||||
coreml_image = Image.open(coreml_img_path)
|
|
||||||
mps_image = Image.open(mps_img_path)
|
|
||||||
|
|
||||||
assert psnr(np.array(coreml_image), np.array(mps_image)) > 25
|
|
||||||
|
|
||||||
|
|
||||||
def psnr(img1, img2):
|
|
||||||
mse = np.mean((img1 - img2) ** 2)
|
|
||||||
if mse == 0:
|
|
||||||
return 100
|
|
||||||
PIXEL_MAX = 255.0
|
|
||||||
return 20 * np.log10(PIXEL_MAX / np.sqrt(mse))
|
|
||||||
|
|
||||||
|
|
||||||
def queue_prompt(prompt: dict):
|
|
||||||
p = {"prompt": prompt}
|
|
||||||
data = json.dumps(p).encode("utf-8")
|
|
||||||
req = requests.post("http://localhost:8188/prompt", data=data)
|
|
||||||
assert req.status_code == 200
|
|
||||||
while True:
|
|
||||||
req = requests.get("http://localhost:8188/prompt")
|
|
||||||
if req.json()["exec_info"]["queue_remaining"] == 0:
|
|
||||||
break
|
|
||||||
|
|
||||||
|
|
||||||
def randomize_seed_in_prompt(prompt):
|
|
||||||
seed = torch.random.seed()
|
|
||||||
prompt["3"]["inputs"]["seed"] = seed
|
|
||||||
prompt["11"]["inputs"]["seed"] = seed
|
|
||||||
return prompt
|
|
||||||
Binary file not shown.
|
After Width: | Height: | Size: 448 KiB |
@@ -0,0 +1 @@
|
|||||||
|
e89344e544d4edfbd3ebe9a1c78dadb2729f53549666052b74ac7308f326f4fc
|
||||||
@@ -0,0 +1,170 @@
|
|||||||
|
"""[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 refactor that doesn't touch the math
|
||||||
|
should hit it). PSNR is the soft gate that tolerates the drift a
|
||||||
|
toolchain bump injects through different MIL graphs / kernel selection
|
||||||
|
/ fp accumulation order — anything below the threshold is treated as a
|
||||||
|
regression.
|
||||||
|
|
||||||
|
The 20 dB default absorbs Apple Neural Engine run-to-run nondeterminism:
|
||||||
|
the same model and seed can drift several dB between runs as the 20
|
||||||
|
sampling steps amplify tiny per-step UNet differences (kernel selection /
|
||||||
|
fp accumulation order). Same-scene ANE outputs have been observed at
|
||||||
|
~23 dB, so 20 leaves margin while still catching gross regressions — a
|
||||||
|
broken image lands far lower. Bump it up for pure-refactor PRs that must
|
||||||
|
not change math; down for toolchain bumps.
|
||||||
|
|
||||||
|
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", "20"))
|
||||||
|
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 — 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,122 @@
|
|||||||
|
"""Tier 1 smoke: convert a synthetic micro-UNet through coremltools and load
|
||||||
|
it back with python_coreml_stable_diffusion's CoreMLModel.
|
||||||
|
|
||||||
|
Purpose: catch API breakage in coremltools / ml-stable-diffusion *without*
|
||||||
|
needing a real SD checkpoint, the ANE, or a converted .mlmodelc on disk.
|
||||||
|
Runs in minutes on a hosted macOS-ARM runner (no Apple internal stuff).
|
||||||
|
|
||||||
|
What it asserts:
|
||||||
|
- coremltools.convert still accepts the call shape we use today
|
||||||
|
- the resulting .mlpackage round-trips through CoreMLModel
|
||||||
|
- expected_inputs exposes the input names/shapes we declared
|
||||||
|
- calling the model returns the named output (`noise_pred`)
|
||||||
|
|
||||||
|
Auto-skips on non-Apple-Silicon hosts so Tier 0 CI on Linux ignores it.
|
||||||
|
"""
|
||||||
|
import platform
|
||||||
|
import shutil
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
|
||||||
|
pytestmark = pytest.mark.skipif(
|
||||||
|
platform.system() != "Darwin" or platform.machine() != "arm64",
|
||||||
|
reason="Tier 1 requires macOS on Apple Silicon",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Tiny shapes — large enough to exercise conv2d + linear + addition kernels in
|
||||||
|
# coremltools, small enough that conversion finishes in seconds on CPU.
|
||||||
|
SAMPLE_SHAPE = (1, 4, 8, 8)
|
||||||
|
TIMESTEP_SHAPE = (1,)
|
||||||
|
ENCODER_SHAPE = (1, 64, 1, 4) # matches SD's transposed encoder_hidden_states layout
|
||||||
|
OUT_NAME = "noise_pred"
|
||||||
|
|
||||||
|
|
||||||
|
class TinyUNet(nn.Module):
|
||||||
|
"""Minimal UNet-shaped graph: conv -> add(time+context) -> conv.
|
||||||
|
|
||||||
|
Not a real diffusion model. Just enough op variety to exercise the
|
||||||
|
PyTorch -> MIL frontend in coremltools and confirm we can still wire
|
||||||
|
the inputs/outputs the way ml-stable-diffusion expects.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.conv_in = nn.Conv2d(4, 8, kernel_size=3, padding=1)
|
||||||
|
self.conv_out = nn.Conv2d(8, 4, kernel_size=3, padding=1)
|
||||||
|
self.time_proj = nn.Linear(1, 8)
|
||||||
|
self.text_proj = nn.Linear(64, 8)
|
||||||
|
|
||||||
|
def forward(self, sample, timestep, encoder_hidden_states):
|
||||||
|
h = self.conv_in(sample)
|
||||||
|
t_emb = self.time_proj(timestep.unsqueeze(-1)).view(1, 8, 1, 1)
|
||||||
|
c_emb = self.text_proj(encoder_hidden_states.squeeze(2).mean(-1)).view(1, 8, 1, 1)
|
||||||
|
h = h + t_emb + c_emb
|
||||||
|
return self.conv_out(h)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(scope="module")
|
||||||
|
def tiny_mlpackage(tmp_path_factory):
|
||||||
|
"""Convert TinyUNet once per test session and reuse the .mlpackage."""
|
||||||
|
import coremltools as ct
|
||||||
|
|
||||||
|
torch.manual_seed(0)
|
||||||
|
model = TinyUNet().eval()
|
||||||
|
example = (
|
||||||
|
torch.randn(*SAMPLE_SHAPE),
|
||||||
|
torch.randn(*TIMESTEP_SHAPE),
|
||||||
|
torch.randn(*ENCODER_SHAPE),
|
||||||
|
)
|
||||||
|
traced = torch.jit.trace(model, example)
|
||||||
|
|
||||||
|
mlmodel = ct.convert(
|
||||||
|
traced,
|
||||||
|
inputs=[
|
||||||
|
ct.TensorType(name="sample", shape=SAMPLE_SHAPE, dtype=np.float16),
|
||||||
|
ct.TensorType(name="timestep", shape=TIMESTEP_SHAPE, dtype=np.float16),
|
||||||
|
ct.TensorType(name="encoder_hidden_states", shape=ENCODER_SHAPE, dtype=np.float16),
|
||||||
|
],
|
||||||
|
outputs=[ct.TensorType(name=OUT_NAME, dtype=np.float16)],
|
||||||
|
compute_units=ct.ComputeUnit.CPU_ONLY,
|
||||||
|
compute_precision=ct.precision.FLOAT16,
|
||||||
|
convert_to="mlprogram",
|
||||||
|
minimum_deployment_target=ct.target.macOS13,
|
||||||
|
)
|
||||||
|
|
||||||
|
out_dir = tmp_path_factory.mktemp("tiny_unet")
|
||||||
|
pkg_path = out_dir / "tiny.mlpackage"
|
||||||
|
mlmodel.save(str(pkg_path))
|
||||||
|
yield pkg_path
|
||||||
|
shutil.rmtree(out_dir, ignore_errors=True)
|
||||||
|
|
||||||
|
|
||||||
|
def test_coremltools_convert_round_trips_via_coreml_model(tiny_mlpackage):
|
||||||
|
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
|
||||||
|
|
||||||
|
model = CoreMLModel(str(tiny_mlpackage), "CPU_ONLY", "packages")
|
||||||
|
|
||||||
|
# expected_inputs is the contract our wrappers depend on. Lock the shape
|
||||||
|
# of the dict + a sample entry.
|
||||||
|
expected = dict(model.expected_inputs)
|
||||||
|
assert set(expected.keys()) == {"sample", "timestep", "encoder_hidden_states"}
|
||||||
|
assert tuple(expected["sample"]["shape"]) == SAMPLE_SHAPE
|
||||||
|
assert tuple(expected["timestep"]["shape"]) == TIMESTEP_SHAPE
|
||||||
|
assert tuple(expected["encoder_hidden_states"]["shape"]) == ENCODER_SHAPE
|
||||||
|
|
||||||
|
# Forward pass: drive the model the way CoreMLModelWrapper does.
|
||||||
|
rng = np.random.default_rng(0)
|
||||||
|
inputs = {
|
||||||
|
"sample": rng.standard_normal(SAMPLE_SHAPE).astype(np.float16),
|
||||||
|
"timestep": rng.standard_normal(TIMESTEP_SHAPE).astype(np.float16),
|
||||||
|
"encoder_hidden_states": rng.standard_normal(ENCODER_SHAPE).astype(np.float16),
|
||||||
|
}
|
||||||
|
out = model(**inputs)
|
||||||
|
assert isinstance(out, dict), f"unexpected output type: {type(out)}"
|
||||||
|
assert OUT_NAME in out, f"missing output {OUT_NAME!r}; got {sorted(out)}"
|
||||||
|
assert out[OUT_NAME].shape == SAMPLE_SHAPE, (
|
||||||
|
f"output shape drift: got {out[OUT_NAME].shape}, expected {SAMPLE_SHAPE}"
|
||||||
|
)
|
||||||
@@ -0,0 +1,186 @@
|
|||||||
|
"""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.core.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 @@
|
|||||||
|
"""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 a refactor silently
|
||||||
|
re-shapes them, generation breaks.
|
||||||
|
"""
|
||||||
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from coreml_suite.core.inputs 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 @@
|
|||||||
|
"""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 a refactor
|
||||||
|
cannot silently shift either contract.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from coreml_suite.core.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 a refactor 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,197 @@
|
|||||||
|
"""Characterization tests for the .mlpackage filename composition.
|
||||||
|
|
||||||
|
The filename composition is the pure
|
||||||
|
coreml_suite.core.naming.compose_out_name function. CoreMLConverter.convert
|
||||||
|
calls it; testing the pure function avoids monkey-patching heavy converter
|
||||||
|
internals just to capture the string.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from coreml_suite.core.naming import compose_out_name, lora_names_from_params
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- attention suffixes ----------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"attn_name,suffix",
|
||||||
|
[
|
||||||
|
("SPLIT_EINSUM", "se"),
|
||||||
|
("SPLIT_EINSUM_V2", "se2"),
|
||||||
|
("ORIGINAL", "orig"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_attention_suffix(attn_name, suffix):
|
||||||
|
out = compose_out_name(
|
||||||
|
ckpt_name="dreamshaper_8.safetensors",
|
||||||
|
batch_size=1, width=512, height=512,
|
||||||
|
controlnet_support=False,
|
||||||
|
attention_implementation=attn_name,
|
||||||
|
)
|
||||||
|
assert out == f"dreamshaper_8_1x512x512_{suffix}"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- batch / size ----------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_includes_batch_and_size():
|
||||||
|
out = compose_out_name(
|
||||||
|
ckpt_name="dreamshaper_8.safetensors",
|
||||||
|
batch_size=4, width=768, height=1024,
|
||||||
|
controlnet_support=False,
|
||||||
|
attention_implementation="SPLIT_EINSUM",
|
||||||
|
)
|
||||||
|
assert out == "dreamshaper_8_4x768x1024_se"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- ControlNet ------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_appends_cn_suffix_when_controlnet_support_true():
|
||||||
|
out = compose_out_name(
|
||||||
|
ckpt_name="dreamshaper_8.safetensors",
|
||||||
|
batch_size=1, width=512, height=512,
|
||||||
|
controlnet_support=True,
|
||||||
|
attention_implementation="SPLIT_EINSUM",
|
||||||
|
)
|
||||||
|
assert out == "dreamshaper_8_1x512x512_cn_se"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- ckpt name massage -----------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_drops_extension_at_first_period():
|
||||||
|
out = compose_out_name(
|
||||||
|
ckpt_name="my.checkpoint.v2.safetensors",
|
||||||
|
batch_size=1, width=512, height=512,
|
||||||
|
controlnet_support=False,
|
||||||
|
attention_implementation="SPLIT_EINSUM",
|
||||||
|
)
|
||||||
|
assert out == "my_1x512x512_se"
|
||||||
|
|
||||||
|
|
||||||
|
def test_replaces_spaces_with_underscores():
|
||||||
|
out = compose_out_name(
|
||||||
|
ckpt_name="dream shaper 8.safetensors",
|
||||||
|
batch_size=1, width=512, height=512,
|
||||||
|
controlnet_support=False,
|
||||||
|
attention_implementation="SPLIT_EINSUM",
|
||||||
|
)
|
||||||
|
assert out == "dream_shaper_8_1x512x512_se"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- LoRA suffixes ---------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_single_lora():
|
||||||
|
out = compose_out_name(
|
||||||
|
ckpt_name="dreamshaper_8.safetensors",
|
||||||
|
batch_size=1, width=512, height=512,
|
||||||
|
controlnet_support=False,
|
||||||
|
attention_implementation="SPLIT_EINSUM",
|
||||||
|
lora_names=["epi_noiseoffset.safetensors"],
|
||||||
|
)
|
||||||
|
assert out == "dreamshaper_8_epi_noiseoffset_1x512x512_se"
|
||||||
|
|
||||||
|
|
||||||
|
def test_multiple_loras_sorted():
|
||||||
|
out = compose_out_name(
|
||||||
|
ckpt_name="dreamshaper_8.safetensors",
|
||||||
|
batch_size=1, width=512, height=512,
|
||||||
|
controlnet_support=False,
|
||||||
|
attention_implementation="SPLIT_EINSUM",
|
||||||
|
lora_names=["zoom.safetensors", "alpha.safetensors", "moody.safetensors"],
|
||||||
|
)
|
||||||
|
assert out == "dreamshaper_8_alpha_moody_zoom_1x512x512_se"
|
||||||
|
|
||||||
|
|
||||||
|
def test_lora_plus_controlnet():
|
||||||
|
out = compose_out_name(
|
||||||
|
ckpt_name="dreamshaper_8.safetensors",
|
||||||
|
batch_size=1, width=512, height=512,
|
||||||
|
controlnet_support=True,
|
||||||
|
attention_implementation="SPLIT_EINSUM",
|
||||||
|
lora_names=["a.safetensors"],
|
||||||
|
)
|
||||||
|
assert out == "dreamshaper_8_a_1x512x512_cn_se"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- sdxl combinations -----------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_sdxl_1024_original_gpu():
|
||||||
|
out = compose_out_name(
|
||||||
|
ckpt_name="sd_xl_base_1.0.safetensors",
|
||||||
|
batch_size=1, width=1024, height=1024,
|
||||||
|
controlnet_support=False,
|
||||||
|
attention_implementation="ORIGINAL",
|
||||||
|
)
|
||||||
|
assert out == "sd_xl_base_1_1x1024x1024_orig"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- lora_names_from_params helper ----------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_lora_names_from_params_sorts_by_name():
|
||||||
|
names = lora_names_from_params([
|
||||||
|
("zebra.safetensors", 1.0),
|
||||||
|
("apple.safetensors", 0.5),
|
||||||
|
("mango.safetensors", 0.7),
|
||||||
|
])
|
||||||
|
assert names == ["apple.safetensors", "mango.safetensors", "zebra.safetensors"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_lora_names_from_params_empty_list():
|
||||||
|
assert lora_names_from_params([]) == []
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- quantize_nbits suffix ------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_quantize_nbits_none_appends_nothing():
|
||||||
|
"""'none' is the default and must keep the unquantized filename so
|
||||||
|
existing cached .mlpackages still resolve."""
|
||||||
|
out = compose_out_name(
|
||||||
|
ckpt_name="dreamshaper_8.safetensors",
|
||||||
|
batch_size=1, width=512, height=512,
|
||||||
|
controlnet_support=False,
|
||||||
|
attention_implementation="SPLIT_EINSUM",
|
||||||
|
quantize_nbits="none",
|
||||||
|
)
|
||||||
|
assert out == "dreamshaper_8_1x512x512_se"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("nbits,suffix", [("4", "_q4"), ("6", "_q6"), ("8", "_q8")])
|
||||||
|
def test_quantize_nbits_appends_q_suffix(nbits, suffix):
|
||||||
|
out = compose_out_name(
|
||||||
|
ckpt_name="dreamshaper_8.safetensors",
|
||||||
|
batch_size=1, width=512, height=512,
|
||||||
|
controlnet_support=False,
|
||||||
|
attention_implementation="SPLIT_EINSUM",
|
||||||
|
quantize_nbits=nbits,
|
||||||
|
)
|
||||||
|
assert out == f"dreamshaper_8_1x512x512_se{suffix}"
|
||||||
|
|
||||||
|
|
||||||
|
def test_quantize_nbits_with_controlnet_and_lora():
|
||||||
|
out = compose_out_name(
|
||||||
|
ckpt_name="dreamshaper_8.safetensors",
|
||||||
|
batch_size=1, width=512, height=512,
|
||||||
|
controlnet_support=True,
|
||||||
|
attention_implementation="SPLIT_EINSUM",
|
||||||
|
lora_names=["a.safetensors"],
|
||||||
|
quantize_nbits="6",
|
||||||
|
)
|
||||||
|
assert out == "dreamshaper_8_a_1x512x512_cn_se_q6"
|
||||||
|
|
||||||
|
|
||||||
|
def test_quantize_nbits_invalid_raises():
|
||||||
|
import pytest as _pytest
|
||||||
|
with _pytest.raises(ValueError, match="quantize_nbits"):
|
||||||
|
compose_out_name(
|
||||||
|
ckpt_name="x.safetensors",
|
||||||
|
batch_size=1, width=512, height=512,
|
||||||
|
controlnet_support=False,
|
||||||
|
attention_implementation="SPLIT_EINSUM",
|
||||||
|
quantize_nbits="16", # not in {none, 8, 6, 4}
|
||||||
|
)
|
||||||
@@ -0,0 +1,127 @@
|
|||||||
|
"""Characterization tests for the SDXL options math.
|
||||||
|
|
||||||
|
The SDXL time_ids / text_embeds math lives in
|
||||||
|
coreml_suite.core.sdxl as pure builders. The framework adapter
|
||||||
|
add_sdxl_model_options (in models.py) is exercised separately by the m2
|
||||||
|
golden image test; here we just lock the pure math.
|
||||||
|
"""
|
||||||
|
import inspect
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from coreml_suite.core.sdxl import (
|
||||||
|
build_sdxl_text_embeds,
|
||||||
|
build_sdxl_time_ids,
|
||||||
|
sdxl_model_function_wrapper,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _deterministic_seed():
|
||||||
|
torch.manual_seed(0)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- build_sdxl_time_ids: base (len 6) -------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_time_ids_base_defaults():
|
||||||
|
out = build_sdxl_time_ids({}, {}, is_base=True, is_refiner=False)
|
||||||
|
expected = torch.tensor([[768, 768, 0, 0, 768, 768], [768, 768, 0, 0, 768, 768]])
|
||||||
|
assert out.shape == (2, 6)
|
||||||
|
assert torch.equal(out, expected)
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_time_ids_base_respects_overrides():
|
||||||
|
pos = {"height": 1024, "width": 512, "crop_h": 8, "crop_w": 4,
|
||||||
|
"target_height": 1024, "target_width": 1024}
|
||||||
|
neg = {"height": 256, "width": 256, "crop_h": 0, "crop_w": 0,
|
||||||
|
"target_height": 256, "target_width": 256}
|
||||||
|
out = build_sdxl_time_ids(pos, neg, is_base=True, is_refiner=False)
|
||||||
|
expected = torch.tensor([[1024, 512, 8, 4, 1024, 1024], [256, 256, 0, 0, 256, 256]])
|
||||||
|
assert torch.equal(out, expected)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- build_sdxl_time_ids: refiner (len 5) ----------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_time_ids_refiner_defaults():
|
||||||
|
out = build_sdxl_time_ids({}, {}, is_base=False, is_refiner=True)
|
||||||
|
expected = torch.tensor([[768, 768, 0, 0, 6.0], [768, 768, 0, 0, 2.5]])
|
||||||
|
assert out.shape == (2, 5)
|
||||||
|
assert torch.equal(out, expected)
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_time_ids_refiner_respects_aesthetic_score():
|
||||||
|
pos = {"aesthetic_score": 8.5}
|
||||||
|
neg = {"aesthetic_score": 1.5}
|
||||||
|
out = build_sdxl_time_ids(pos, neg, is_base=False, is_refiner=True)
|
||||||
|
expected = torch.tensor([[768, 768, 0, 0, 8.5], [768, 768, 0, 0, 1.5]])
|
||||||
|
assert torch.equal(out, expected)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- build_sdxl_time_ids: edge case ----------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_time_ids_neither_base_nor_refiner_returns_len4():
|
||||||
|
out = build_sdxl_time_ids({}, {}, is_base=False, is_refiner=False)
|
||||||
|
assert out.shape == (2, 4)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- build_sdxl_text_embeds ------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_text_embeds_concat_pos_then_neg():
|
||||||
|
pos = torch.full((1, 1280), 1.0)
|
||||||
|
neg = torch.full((1, 1280), -1.0)
|
||||||
|
out = build_sdxl_text_embeds(pos, neg)
|
||||||
|
assert out.shape == (2, 1280)
|
||||||
|
assert torch.equal(out[0], pos[0])
|
||||||
|
assert torch.equal(out[1], neg[0])
|
||||||
|
|
||||||
|
|
||||||
|
# ---------- sdxl_model_function_wrapper closure -----------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_wrapper_captures_time_ids_text_embeds_refiner_via_closure():
|
||||||
|
time_ids = torch.zeros(2, 6)
|
||||||
|
text_embeds = torch.zeros(2, 1280)
|
||||||
|
wrapper = sdxl_model_function_wrapper(time_ids, text_embeds, refiner=False)
|
||||||
|
closure = inspect.getclosurevars(wrapper).nonlocals
|
||||||
|
assert closure["time_ids"] is time_ids
|
||||||
|
assert closure["text_embeds"] is text_embeds
|
||||||
|
assert closure["refiner"] is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_wrapper_returns_zero_when_context_missing():
|
||||||
|
"""When c_crossattn is None the wrapper short-circuits to zeros_like(x).
|
||||||
|
Locked here because the refactor mustn't change this default."""
|
||||||
|
wrapper = sdxl_model_function_wrapper(torch.zeros(2, 6), torch.zeros(2, 1280))
|
||||||
|
x = torch.randn(2, 4, 16, 16)
|
||||||
|
out = wrapper(
|
||||||
|
model_function=lambda *a, **kw: pytest.fail("model_function must not run"),
|
||||||
|
params={"input": x, "timestep": torch.zeros(2), "c": {}},
|
||||||
|
)
|
||||||
|
assert torch.equal(out, torch.zeros_like(x))
|
||||||
|
|
||||||
|
|
||||||
|
def test_wrapper_refiner_truncates_context_to_g_clip():
|
||||||
|
"""refiner=True slices c_crossattn[:, :, 768:] before forwarding."""
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
def fake_model(x, t, **c):
|
||||||
|
captured["context_shape"] = c["c_crossattn"].shape
|
||||||
|
captured["time_ids_shape"] = c["time_ids"].shape
|
||||||
|
return x
|
||||||
|
|
||||||
|
wrapper = sdxl_model_function_wrapper(
|
||||||
|
torch.zeros(2, 5), torch.zeros(2, 1280), refiner=True
|
||||||
|
)
|
||||||
|
x = torch.randn(2, 4, 16, 16)
|
||||||
|
context = torch.randn(2, 77, 2048) # 768 + 1280 dims
|
||||||
|
wrapper(
|
||||||
|
model_function=fake_model,
|
||||||
|
params={"input": x, "timestep": torch.zeros(2), "c": {"c_crossattn": context}},
|
||||||
|
)
|
||||||
|
assert captured["context_shape"] == (2, 77, 1280)
|
||||||
|
assert captured["time_ids_shape"] == (2, 5)
|
||||||
+23
-26
@@ -1,19 +1,22 @@
|
|||||||
import pytest
|
"""Smoke tests for the pure batch-chunking helpers in coreml_suite.core.
|
||||||
|
|
||||||
|
Uses torch.device('cpu') instead of comfy.model_management.get_torch_device
|
||||||
|
so Tier 0 runs without ComfyUI.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from comfy.model_management import get_torch_device
|
from coreml_suite.core.controlnet import chunk_control
|
||||||
from coreml_suite.latents import chunk_batch, merge_chunks
|
from coreml_suite.core.inputs import CoreMLInputs
|
||||||
from coreml_suite.controlnet import chunk_control
|
from coreml_suite.core.latents import chunk_batch, merge_chunks
|
||||||
from coreml_suite.models import (
|
|
||||||
CoreMLInputs,
|
|
||||||
)
|
CPU = torch.device("cpu")
|
||||||
from coreml_suite.config import get_model_config
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def expected_inputs():
|
def expected_inputs():
|
||||||
expected = {
|
return {
|
||||||
"sample": {"shape": (2, 4, 64, 64)},
|
"sample": {"shape": (2, 4, 64, 64)},
|
||||||
"timestep": {"shape": (2,)},
|
"timestep": {"shape": (2,)},
|
||||||
"timestep_cond": {"shape": (2, 256)},
|
"timestep_cond": {"shape": (2, 256)},
|
||||||
@@ -21,17 +24,11 @@ def expected_inputs():
|
|||||||
"additional_residual_0": {"shape": (2, 320, 64, 64)},
|
"additional_residual_0": {"shape": (2, 320, 64, 64)},
|
||||||
"additional_residual_1": {"shape": (2, 640, 32, 32)},
|
"additional_residual_1": {"shape": (2, 640, 32, 32)},
|
||||||
}
|
}
|
||||||
return expected
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def model_config():
|
|
||||||
return get_model_config()
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9])
|
@pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9])
|
||||||
def test_batch_chunking(batch_size):
|
def test_batch_chunking(batch_size):
|
||||||
latent_image = torch.randn(batch_size, 4, 64, 64).to(get_torch_device())
|
latent_image = torch.randn(batch_size, 4, 64, 64).to(CPU)
|
||||||
target_shape = (4, 4, 64, 64)
|
target_shape = (4, 4, 64, 64)
|
||||||
|
|
||||||
chunked = chunk_batch(latent_image, target_shape)
|
chunked = chunk_batch(latent_image, target_shape)
|
||||||
@@ -45,7 +42,7 @@ def test_batch_chunking(batch_size):
|
|||||||
|
|
||||||
@pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9])
|
@pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9])
|
||||||
def test_merge_chunks(batch_size):
|
def test_merge_chunks(batch_size):
|
||||||
input_tensor = torch.randn(batch_size, 4, 64, 64).to(get_torch_device())
|
input_tensor = torch.randn(batch_size, 4, 64, 64).to(CPU)
|
||||||
target_shape = (4, 4, 64, 64)
|
target_shape = (4, 4, 64, 64)
|
||||||
chunked = chunk_batch(input_tensor, target_shape)
|
chunked = chunk_batch(input_tensor, target_shape)
|
||||||
|
|
||||||
@@ -57,16 +54,16 @@ def test_merge_chunks(batch_size):
|
|||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def inputs():
|
def inputs():
|
||||||
x = torch.randn(1, 4, 64, 64).to(get_torch_device())
|
x = torch.randn(1, 4, 64, 64).to(CPU)
|
||||||
t = torch.randn([1]).to(get_torch_device())
|
t = torch.randn([1]).to(CPU)
|
||||||
c_crossattn = torch.randn(1, 77, 768).to(get_torch_device())
|
c_crossattn = torch.randn(1, 77, 768).to(CPU)
|
||||||
control = {
|
control = {
|
||||||
"output": [
|
"output": [
|
||||||
torch.randn(1, 320, 64, 64).to(get_torch_device()),
|
torch.randn(1, 320, 64, 64).to(CPU),
|
||||||
torch.randn(1, 640, 32, 32).to(get_torch_device()),
|
torch.randn(1, 640, 32, 32).to(CPU),
|
||||||
],
|
],
|
||||||
}
|
}
|
||||||
timestep_cond = torch.randn(1, 256).to(get_torch_device())
|
timestep_cond = torch.randn(1, 256).to(CPU)
|
||||||
|
|
||||||
return CoreMLInputs(x, t, c_crossattn, control, timestep_cond=timestep_cond)
|
return CoreMLInputs(x, t, c_crossattn, control, timestep_cond=timestep_cond)
|
||||||
|
|
||||||
@@ -86,11 +83,11 @@ def inputs():
|
|||||||
def test_chunking_controlnet(b, target_size, num_chunks):
|
def test_chunking_controlnet(b, target_size, num_chunks):
|
||||||
cn = {
|
cn = {
|
||||||
"output": [
|
"output": [
|
||||||
torch.randn(b, 320, 64, 64).to(get_torch_device()),
|
torch.randn(b, 320, 64, 64).to(CPU),
|
||||||
torch.randn(b, 640, 32, 32).to(get_torch_device()),
|
torch.randn(b, 640, 32, 32).to(CPU),
|
||||||
],
|
],
|
||||||
"middle": [
|
"middle": [
|
||||||
torch.randn(b, 1280, 8, 8).to(get_torch_device()),
|
torch.randn(b, 1280, 8, 8).to(CPU),
|
||||||
],
|
],
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,43 @@
|
|||||||
|
"""Gate: prove the Tier-0 lane is framework-free.
|
||||||
|
|
||||||
|
In a pure `pytest -m unit` run, none of the banned runtime modules
|
||||||
|
(comfy, coremltools, python_coreml_stable_diffusion, folder_paths,
|
||||||
|
nodes, comfy_extras, diffusers, diffusionkit) may be in sys.modules
|
||||||
|
after collection. If they are, a tests/unit/ file is transitively
|
||||||
|
pulling them in and the Tier-0 promise — "runs on Linux with no Mac
|
||||||
|
stack" — is broken.
|
||||||
|
|
||||||
|
When other tiers (m2 / integration) are also collected, comfy is
|
||||||
|
expected in sys.modules (integration imports it deliberately), so the
|
||||||
|
check is skipped in mixed runs — Tier-0 purity is only meaningful when
|
||||||
|
nothing else is loaded.
|
||||||
|
"""
|
||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
BANNED_ROOTS = {
|
||||||
|
"comfy",
|
||||||
|
"comfy_extras",
|
||||||
|
"coremltools",
|
||||||
|
"python_coreml_stable_diffusion",
|
||||||
|
"folder_paths",
|
||||||
|
"nodes",
|
||||||
|
"diffusers",
|
||||||
|
"diffusionkit",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_no_framework_modules_loaded_by_unit_tier(request):
|
||||||
|
markexpr = request.config.option.markexpr
|
||||||
|
if markexpr != "unit":
|
||||||
|
pytest.skip(
|
||||||
|
"purity gate only meaningful in a pure `-m unit` run "
|
||||||
|
f"(got markexpr={markexpr!r}); other tiers are expected to "
|
||||||
|
"import comfy/coremltools."
|
||||||
|
)
|
||||||
|
loaded = {name for name in sys.modules if name.split(".")[0] in BANNED_ROOTS}
|
||||||
|
assert not loaded, (
|
||||||
|
f"Tier-0 leakage: these framework modules are in sys.modules after "
|
||||||
|
f"collecting tests/unit/: {sorted(loaded)}. Pure-core promise broken."
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user