Compare commits
78
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d4009e7b69 | ||
|
|
b20507e9af | ||
|
|
7ea420aef1 | ||
|
|
d4bda0e740 | ||
|
|
21ab05fd4f | ||
|
|
73b13e0d23 | ||
|
|
991ab6a40a | ||
|
|
1320e5cd9c | ||
|
|
ebdc700177 | ||
|
|
d752f41b66 | ||
|
|
ff9a5c91d2 | ||
|
|
d17a93323f | ||
|
|
9fb310700e | ||
|
|
f95a439d62 | ||
|
|
b044efe201 | ||
|
|
3f666ac0ea | ||
|
|
acecd10aee | ||
|
|
2f6597b8c0 | ||
|
|
a03a56a58e | ||
|
|
9b7e5a6cd8 | ||
|
|
600023382c | ||
|
|
7caf1cea1d | ||
|
|
c6229e5c5f | ||
|
|
5de5722474 | ||
|
|
a3cf825d79 | ||
|
|
75057e4ed2 | ||
|
|
a1d81faf68 | ||
|
|
1ff260fc36 | ||
|
|
83ad02748f | ||
|
|
4c9195bbc0 | ||
|
|
7643211d8d | ||
|
|
6be88fee2e | ||
|
|
1a82e4b48f | ||
|
|
7a5d040b61 | ||
|
|
b4313d731e | ||
|
|
9489503cbe | ||
|
|
edc8e39c83 | ||
|
|
de8915eb6c | ||
|
|
93ebaf4d5d | ||
|
|
1bc728d0ea | ||
|
|
33829c292f | ||
|
|
7b1c3c7ba7 | ||
|
|
8c9fbacb45 | ||
|
|
997c6a78ff | ||
|
|
d9be9c13e2 | ||
|
|
31a6ac6d2f | ||
|
|
11772e4e69 | ||
|
|
d4f3ed6fa9 | ||
|
|
1d450cca3c | ||
|
|
b2102592cd | ||
|
|
3d7473903b | ||
|
|
b12cd83041 | ||
|
|
ab567e48af | ||
|
|
12f667190f | ||
|
|
d612d1ffef | ||
|
|
d4666d3615 | ||
|
|
42c6a66a7c | ||
|
|
f4a1eb974b | ||
|
|
c09bbeabe2 | ||
|
|
43f8d330a0 | ||
|
|
4e32ca8dbc | ||
|
|
8f639eb2a0 | ||
|
|
bfe22d8d06 | ||
|
|
92080ae196 | ||
|
|
a638a79f81 | ||
|
|
d6f7188f7e | ||
|
|
ca715599c1 | ||
|
|
ef78f8596f | ||
|
|
3d6f8b7dcd | ||
|
|
83f49f0937 | ||
|
|
183d0b2707 | ||
|
|
ddbeb36d52 | ||
|
|
2ee81bd41d | ||
|
|
a1c66249e2 | ||
|
|
01fafd70f3 | ||
|
|
447b25c774 | ||
|
|
c3038501eb | ||
|
|
b326b3d3b9 |
@@ -1,25 +0,0 @@
|
||||
name: Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'aszc-dev' }}
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
@@ -1,33 +0,0 @@
|
||||
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
|
||||
@@ -1,31 +0,0 @@
|
||||
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
|
||||
@@ -1,124 +0,0 @@
|
||||
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
|
||||
+1
-4
@@ -1,6 +1,3 @@
|
||||
playground/
|
||||
experiments/
|
||||
__pycache__/
|
||||
models/
|
||||
.venv/
|
||||
test_results/
|
||||
tests/m2/_latest_generated.png
|
||||
|
||||
@@ -370,47 +370,6 @@ 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
|
||||
|
||||
- Core ML models are fixed in terms of their inputs and outputs.
|
||||
@@ -424,6 +383,49 @@ verifies this on every Tier 2 run.
|
||||
Unless [EnumeratedShapes](https://apple.github.io/coremltools/docs-guides/source/flexible-inputs.html#select-from-predetermined-shapes)
|
||||
is used during conversion. Needs more testing.
|
||||
|
||||
## FAQ
|
||||
|
||||
### Hardware and Performance
|
||||
|
||||
#### What's the difference between MPS, GPU, and ANE?
|
||||
- **MPS (Metal Performance Shaders)**: Apple's framework for GPU acceleration. It's what PyTorch uses by default on Apple Silicon.
|
||||
- **GPU**: The graphics processing unit on your Apple Silicon chip.
|
||||
- **ANE (Apple Neural Engine)**: A specialized hardware accelerator for machine learning tasks.
|
||||
|
||||
#### Which compute unit should I choose?
|
||||
- **CPU_AND_ANE**: Best for models converted with `--attention-implementation SPLIT_EINSUM`. This is the default and recommended option for most users.
|
||||
- **CPU_AND_GPU**: Best for models converted with `--attention-implementation ORIGINAL`. Use this if you experience issues with ANE.
|
||||
- **CPU_ONLY**: Use this as a fallback if you experience issues with both ANE and GPU.
|
||||
|
||||
#### Do I need `PYTORCH_ENABLE_MPS_FALLBACK=1`?
|
||||
While our Core ML nodes don't use this environment variable directly, it may still be relevant for other parts of ComfyUI that use PyTorch with MPS backend. The setting of this variable is a user preference and depends on your specific needs and workflow requirements.
|
||||
|
||||
### Model Conversion and Compatibility
|
||||
|
||||
#### Is there a performance penalty when using the Core ML Adapter?
|
||||
Yes, there might be a slight performance penalty compared to using directly converted models. However, the adapter provides more flexibility and compatibility with standard ComfyUI nodes.
|
||||
|
||||
#### Does the Core ML Adapter support SDXL?
|
||||
Currently, SDXL support in the Core ML Adapter is limited. While it may work with some models, it's not officially supported and may cause issues.
|
||||
|
||||
#### Are `mlmodelc` and `mlpackage` formats safe?
|
||||
Yes, both formats are safe to use. However, we recommend:
|
||||
1. Always downloading original `.safetensors` files from trusted sources
|
||||
2. Converting them yourself using our tools
|
||||
3. Using the converted `.mlmodelc` files for better performance
|
||||
|
||||
#### Do Core ML models produce identical results to their safetensors counterparts?
|
||||
While the results should be very similar, there might be slight differences due to:
|
||||
- Different numerical precision
|
||||
- Hardware-specific optimizations
|
||||
- Different attention implementations
|
||||
|
||||
#### Should I convert models every time I queue a generation?
|
||||
No! The conversion only happens once when you first use the converter node. After that, you should use the `CoreMLUnetLoader` to load the already converted model.
|
||||
|
||||
#### Will SDXL ever be supported on ANE?
|
||||
Currently, there are technical limitations preventing SDXL from running efficiently on ANE. We recommend using `CPU_AND_GPU` or `CPU_ONLY` for SDXL models.
|
||||
|
||||
## Support
|
||||
|
||||
I'm here to help! If you have any questions or suggestions, don't hesitate to open an issue and I'll do my best
|
||||
|
||||
@@ -1,4 +0,0 @@
|
||||
"""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"]
|
||||
@@ -1,18 +0,0 @@
|
||||
# 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
|
||||
+61
-13
@@ -1,14 +1,62 @@
|
||||
"""Compatibility shim — re-exports from coreml_suite.core.controlnet."""
|
||||
from coreml_suite.core.controlnet import (
|
||||
chunk_control,
|
||||
expand_inputs,
|
||||
extract_residual_kwargs,
|
||||
no_control,
|
||||
)
|
||||
from itertools import chain
|
||||
from math import ceil
|
||||
|
||||
__all__ = [
|
||||
"chunk_control",
|
||||
"expand_inputs",
|
||||
"extract_residual_kwargs",
|
||||
"no_control",
|
||||
]
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from coreml_suite.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
|
||||
|
||||
@@ -258,7 +258,6 @@ def convert_unet(
|
||||
batch_size: int = 1,
|
||||
sample_size: tuple[int, int] = (64, 64),
|
||||
controlnet_support: bool = False,
|
||||
quantize_nbits: str = "none",
|
||||
):
|
||||
coreml_unet = get_unet(model_version, ref_pipe)
|
||||
ref_unet = ref_pipe.unet
|
||||
@@ -306,24 +305,6 @@ def convert_unet(
|
||||
del traced_unet
|
||||
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)
|
||||
logger.info(f"Saved unet into {unet_out_path}")
|
||||
|
||||
@@ -338,7 +319,6 @@ def convert(
|
||||
lora_weights: list[tuple[Union[str, os.PathLike], float]] = None,
|
||||
attn_impl: str = AttentionImplementations.SPLIT_EINSUM.name,
|
||||
config_path: str = None,
|
||||
quantize_nbits: str = "none",
|
||||
):
|
||||
if os.path.exists(unet_out_path):
|
||||
logger.info(f"Found existing model at {unet_out_path}! Skipping..")
|
||||
@@ -364,7 +344,6 @@ def convert(
|
||||
batch_size,
|
||||
sample_size,
|
||||
controlnet_support,
|
||||
quantize_nbits=quantize_nbits,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,10 +0,0 @@
|
||||
"""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.
|
||||
"""
|
||||
@@ -1,67 +0,0 @@
|
||||
"""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
|
||||
@@ -1,113 +0,0 @@
|
||||
"""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,
|
||||
)
|
||||
]
|
||||
@@ -1,42 +0,0 @@
|
||||
"""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]]
|
||||
@@ -1,68 +0,0 @@
|
||||
"""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])]
|
||||
@@ -1,91 +0,0 @@
|
||||
"""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
|
||||
+35
-3
@@ -1,4 +1,36 @@
|
||||
"""Compatibility shim — re-exports from coreml_suite.core.latents."""
|
||||
from coreml_suite.core.latents import chunk_batch, merge_chunks
|
||||
import torch
|
||||
|
||||
__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]]
|
||||
|
||||
+188
-40
@@ -1,44 +1,15 @@
|
||||
"""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 numpy as np
|
||||
import torch
|
||||
|
||||
from comfy import model_base
|
||||
from comfy.model_management import get_torch_device
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from coreml_suite.config import get_model_config, ModelVersion
|
||||
from coreml_suite.core.inputs import CoreMLInputs
|
||||
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.controlnet import extract_residual_kwargs, chunk_control
|
||||
from coreml_suite.latents import chunk_batch, merge_chunks
|
||||
from coreml_suite.lcm.utils import is_lcm
|
||||
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:
|
||||
def __init__(self, coreml_model):
|
||||
@@ -97,27 +68,204 @@ class CoreMLModelWrapperLCM(CoreMLModelWrapper):
|
||||
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):
|
||||
mp = model_patcher.clone()
|
||||
|
||||
pos_dict = positive[0][1]
|
||||
neg_dict = negative[0][1]
|
||||
|
||||
is_base = model_patcher.model.diffusion_model.is_sdxl_base
|
||||
pos_pooled = pos_dict["pooled_output"]
|
||||
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
|
||||
if is_refiner:
|
||||
pos_time_ids += [
|
||||
pos_dict.get("aesthetic_score", 6),
|
||||
]
|
||||
|
||||
time_ids = build_sdxl_time_ids(
|
||||
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"]
|
||||
)
|
||||
neg_time_ids += [
|
||||
neg_dict.get("aesthetic_score", 2.5),
|
||||
]
|
||||
|
||||
mp.model_options |= {
|
||||
time_ids = torch.tensor([pos_time_ids, neg_time_ids])
|
||||
text_embeds = torch.cat((pos_pooled, neg_pooled))
|
||||
|
||||
model_options = {
|
||||
"model_function_wrapper": sdxl_model_function_wrapper(
|
||||
time_ids, text_embeds, is_refiner
|
||||
),
|
||||
}
|
||||
mp.model_options |= model_options
|
||||
|
||||
return mp
|
||||
|
||||
|
||||
|
||||
+16
-22
@@ -8,11 +8,6 @@ import folder_paths
|
||||
from coreml_suite import COREML_NODE
|
||||
from coreml_suite import converter
|
||||
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.logger import logger
|
||||
from nodes import KSampler, LoraLoader, KSamplerAdvanced
|
||||
@@ -249,12 +244,6 @@ class CoreMLConverter(COREML_NODE):
|
||||
"controlnet_support": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"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",),
|
||||
},
|
||||
}
|
||||
@@ -273,7 +262,6 @@ class CoreMLConverter(COREML_NODE):
|
||||
attention_implementation,
|
||||
compute_unit,
|
||||
controlnet_support,
|
||||
quantize_nbits="none",
|
||||
lora_params=None,
|
||||
):
|
||||
"""Converts a LCM model to Core ML.
|
||||
@@ -300,17 +288,24 @@ class CoreMLConverter(COREML_NODE):
|
||||
h = height
|
||||
w = width
|
||||
sample_size = (h // 8, w // 8)
|
||||
out_name = compose_out_name(
|
||||
ckpt_name=ckpt_name,
|
||||
batch_size=batch_size,
|
||||
width=w,
|
||||
height=h,
|
||||
controlnet_support=controlnet_support,
|
||||
attention_implementation=attention_implementation,
|
||||
lora_names=lora_names_from_params(lora_params),
|
||||
quantize_nbits=quantize_nbits,
|
||||
batch_size = batch_size
|
||||
cn_support_str = "_cn" if controlnet_support else ""
|
||||
lora_str = (
|
||||
"_" + "_".join(lora_param[0].split(".")[0] for lora_param in lora_params)
|
||||
if lora_params
|
||||
else ""
|
||||
)
|
||||
|
||||
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"Batch size: {batch_size}")
|
||||
logger.info(f"Width: {w}, Height: {h}")
|
||||
@@ -340,7 +335,6 @@ class CoreMLConverter(COREML_NODE):
|
||||
lora_weights=lora_weights,
|
||||
attn_impl=attention_implementation,
|
||||
config_path=config_path,
|
||||
quantize_nbits=quantize_nbits,
|
||||
)
|
||||
unet_target_path = converter.compile_model(
|
||||
out_path=unet_out_path, out_name=out_name, submodule_name="unet"
|
||||
|
||||
@@ -1,94 +0,0 @@
|
||||
[project]
|
||||
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."
|
||||
version = "1.0.1"
|
||||
license = "MIT"
|
||||
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]
|
||||
Repository = "https://github.com/aszc-dev/ComfyUI-CoreMLSuite"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "aszc-dev"
|
||||
DisplayName = "ComfyUI-CoreMLSuite"
|
||||
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"]
|
||||
+2
-4
@@ -1,7 +1,5 @@
|
||||
git+https://github.com/apple/ml-stable-diffusion.git@e5d960c41a6a4ab200b8db379194127607b1c590
|
||||
torch>=2.7,<2.8
|
||||
coremltools==8.2
|
||||
numpy>=2,<3
|
||||
git+https://github.com/apple/ml-stable-diffusion.git
|
||||
coremltools>=7.1
|
||||
overrides
|
||||
diffusers>=0.22
|
||||
peft>=0.6.2
|
||||
|
||||
@@ -1,61 +0,0 @@
|
||||
"""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
|
||||
@@ -0,0 +1,72 @@
|
||||
import json
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
|
||||
from folder_paths import get_save_image_path, get_output_directory
|
||||
|
||||
IMAGE_PREFIX = "E2E-1.5-CoreML"
|
||||
|
||||
|
||||
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))
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def output_image_repository():
|
||||
repo = OutputImageRepository(IMAGE_PREFIX)
|
||||
yield repo
|
||||
repo.delete_images()
|
||||
|
||||
|
||||
def test_basic_conversion_1_5(output_image_repository):
|
||||
with open("integration/workflows/e2e-1.5-basic-conversion.json") as f:
|
||||
prompt = json.load(f)
|
||||
queue_prompt(prompt)
|
||||
|
||||
full_output_folder, images = output_image_repository.list_images()
|
||||
assert len(images) == 2
|
||||
assert all(image.startswith(IMAGE_PREFIX) for image in images)
|
||||
assert all(image.endswith(".png") for image in images)
|
||||
assert all(
|
||||
os.path.isfile(os.path.join(full_output_folder, image)) for image in images
|
||||
)
|
||||
|
||||
image1 = Image.open(os.path.join(full_output_folder, images[0]))
|
||||
image2 = Image.open(os.path.join(full_output_folder, images[1]))
|
||||
assert psnr(np.array(image1), np.array(image2)) > 30
|
||||
assert psnr(np.array(image2), np.array(image1)) > 30
|
||||
|
||||
|
||||
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
|
||||
@@ -179,4 +179,4 @@
|
||||
"title": "Save Image"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 448 KiB |
@@ -1 +0,0 @@
|
||||
e89344e544d4edfbd3ebe9a1c78dadb2729f53549666052b74ac7308f326f4fc
|
||||
@@ -1,170 +0,0 @@
|
||||
"""[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}"
|
||||
)
|
||||
@@ -1,122 +0,0 @@
|
||||
"""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}"
|
||||
)
|
||||
@@ -1,186 +0,0 @@
|
||||
"""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))
|
||||
@@ -1,228 +0,0 @@
|
||||
"""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))
|
||||
@@ -1,118 +0,0 @@
|
||||
"""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
|
||||
@@ -1,197 +0,0 @@
|
||||
"""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}
|
||||
)
|
||||
@@ -1,127 +0,0 @@
|
||||
"""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)
|
||||
+26
-23
@@ -1,22 +1,19 @@
|
||||
"""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
|
||||
|
||||
from coreml_suite.core.controlnet import chunk_control
|
||||
from coreml_suite.core.inputs import CoreMLInputs
|
||||
from coreml_suite.core.latents import chunk_batch, merge_chunks
|
||||
|
||||
|
||||
CPU = torch.device("cpu")
|
||||
from comfy.model_management import get_torch_device
|
||||
from coreml_suite.latents import chunk_batch, merge_chunks
|
||||
from coreml_suite.controlnet import chunk_control
|
||||
from coreml_suite.models import (
|
||||
CoreMLInputs,
|
||||
)
|
||||
from coreml_suite.config import get_model_config
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def expected_inputs():
|
||||
return {
|
||||
expected = {
|
||||
"sample": {"shape": (2, 4, 64, 64)},
|
||||
"timestep": {"shape": (2,)},
|
||||
"timestep_cond": {"shape": (2, 256)},
|
||||
@@ -24,11 +21,17 @@ def expected_inputs():
|
||||
"additional_residual_0": {"shape": (2, 320, 64, 64)},
|
||||
"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])
|
||||
def test_batch_chunking(batch_size):
|
||||
latent_image = torch.randn(batch_size, 4, 64, 64).to(CPU)
|
||||
latent_image = torch.randn(batch_size, 4, 64, 64).to(get_torch_device())
|
||||
target_shape = (4, 4, 64, 64)
|
||||
|
||||
chunked = chunk_batch(latent_image, target_shape)
|
||||
@@ -42,7 +45,7 @@ def test_batch_chunking(batch_size):
|
||||
|
||||
@pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9])
|
||||
def test_merge_chunks(batch_size):
|
||||
input_tensor = torch.randn(batch_size, 4, 64, 64).to(CPU)
|
||||
input_tensor = torch.randn(batch_size, 4, 64, 64).to(get_torch_device())
|
||||
target_shape = (4, 4, 64, 64)
|
||||
chunked = chunk_batch(input_tensor, target_shape)
|
||||
|
||||
@@ -54,16 +57,16 @@ def test_merge_chunks(batch_size):
|
||||
|
||||
@pytest.fixture
|
||||
def inputs():
|
||||
x = torch.randn(1, 4, 64, 64).to(CPU)
|
||||
t = torch.randn([1]).to(CPU)
|
||||
c_crossattn = torch.randn(1, 77, 768).to(CPU)
|
||||
x = torch.randn(1, 4, 64, 64).to(get_torch_device())
|
||||
t = torch.randn([1]).to(get_torch_device())
|
||||
c_crossattn = torch.randn(1, 77, 768).to(get_torch_device())
|
||||
control = {
|
||||
"output": [
|
||||
torch.randn(1, 320, 64, 64).to(CPU),
|
||||
torch.randn(1, 640, 32, 32).to(CPU),
|
||||
torch.randn(1, 320, 64, 64).to(get_torch_device()),
|
||||
torch.randn(1, 640, 32, 32).to(get_torch_device()),
|
||||
],
|
||||
}
|
||||
timestep_cond = torch.randn(1, 256).to(CPU)
|
||||
timestep_cond = torch.randn(1, 256).to(get_torch_device())
|
||||
|
||||
return CoreMLInputs(x, t, c_crossattn, control, timestep_cond=timestep_cond)
|
||||
|
||||
@@ -83,11 +86,11 @@ def inputs():
|
||||
def test_chunking_controlnet(b, target_size, num_chunks):
|
||||
cn = {
|
||||
"output": [
|
||||
torch.randn(b, 320, 64, 64).to(CPU),
|
||||
torch.randn(b, 640, 32, 32).to(CPU),
|
||||
torch.randn(b, 320, 64, 64).to(get_torch_device()),
|
||||
torch.randn(b, 640, 32, 32).to(get_torch_device()),
|
||||
],
|
||||
"middle": [
|
||||
torch.randn(b, 1280, 8, 8).to(CPU),
|
||||
torch.randn(b, 1280, 8, 8).to(get_torch_device()),
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@@ -1,43 +0,0 @@
|
||||
"""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