Author SHA1 Message Date
aszc-dev 45be061edb docs: update Stable Diffusion 1.5 links 2026-05-26 20:53:47 +02:00
aszc 20ec450f9b fix: allow custom converter dimensions (#60) 2026-05-26 16:17:32 +02:00
aszc d0cca3c3f4 fix(conversion): load .mlpackage instead of unloadable .mlmodelc (#59)
* fix(conversion): load .mlpackage instead of unloadable .mlmodelc

The native ct.models.MLModel runtime added in the diffusers conversion
path cannot load compiled .mlmodelc directories (no Manifest.json), so
the converter's compiled output failed at load with "A valid manifest
does not exist". Both converters now return the .mlpackage directly and
the loader lists only .mlpackage. Removes the now-dead coremlcompiler
wrappers.

* chore(release): bump version to 2.0.1
2026-05-26 16:04:38 +02:00
aszc 65a2de2fab feat!: modernize toolchain and replace apple/ml-stable-diffusion with native diffusers conversion (#58)
* feat(deps): support ComfyUI's numpy 2 toolchain; make conversion optional

The runtime package now installs and runs under numpy 2 / coremltools 9 /
torch 2.7 — matching current ComfyUI — without apple/ml-stable-diffusion.

- Drop the heavy converter stack (ml-stable-diffusion, diffusers, peft,
  omegaconf, overrides, transformers) from runtime dependencies; require
  numpy>=2.
- Vendor the runtime pieces: a slim CoreMLModel wrapper around coremltools
  and the attention-implementation constants.
- Lazy-import the converters; the Convert nodes raise a clear error when the
  legacy conversion dependencies are absent. Loading and sampling existing
  Core ML models no longer needs them.
- CI: Tier 0 tracks the numpy 2 / torch 2.7 toolchain; drop the Tier 2
  golden-image lane (it converts at runtime, which now requires the legacy
  stack) and its fixtures.

* feat(conversion): replace apple/ml-stable-diffusion with native diffusers path

Reimplement Core ML UNet conversion on top of diffusers instead of the
apple/ml-stable-diffusion git dependency, so the full suite (including
conversion) installs through ComfyUI Manager without extras on the NumPy 2
toolchain.

- Add coreml_suite/conversion package: split-einsum attention processors,
  a conv2d output-shape helper, Transformer2D trace patches, and a UNet
  input-adapter wrapper preserving the historical Core ML I/O contract.
- Drop python_coreml_stable_diffusion and overrides; route SD15, SDXL,
  SDXL refiner, and LCM conversion through diffusers UNet2DConditionModel.
- Declare diffusers, peft, omegaconf, and transformers as runtime deps.
- Add characterization tests asserting split-einsum matches reference
  attention math; extend the synthetic-UNet smoke test for the wrapper.
- Bump to 1.1.0 and set requires-comfyui to a semver constraint (>=0.3.27)
  so the Comfy Registry publish succeeds.

* refactor(conversion)!: native diffusers context layout; drop legacy fallbacks

Address PR review feedback:

- Drop the legacy converter ImportError fallbacks and LEGACY_CONVERTER_MODULES
  guards in nodes.py and lcm/nodes.py. Conversion dependencies are mandatory in
  pyproject, so the indirection is dead code.
- Tier 0 CI resolves its toolchain from pyproject via uv (uv sync + uv run)
  instead of hand-pinned pip installs, removing duplicated version maintenance.
- Document the conversion lineage: credit apple/ml-stable-diffusion as the
  origin, note the implementation has diverged to a native diffusers path, and
  state the intent to iterate independently. Fix stale README links that pointed
  users to apple/ml-stable-diffusion for conversion.
- Drop the unused `sources` argument from CoreMLModel.

BREAKING CHANGE: the converted Core ML UNet now takes encoder_hidden_states in
the native diffusers layout (batch, tokens, hidden) instead of
(batch, hidden, 1, tokens). This removes the boundary transposes in
CoreMLUNetWrapper and CoreMLInputs. Core ML models converted with earlier
versions are incompatible and must be re-converted. Bump to 2.0.0.

* test: widen split-einsum allclose tolerance for cross-platform float drift

The split-einsum attention reorders float32 reductions relative to the
reference, so equality holds only up to rounding. The default allclose atol
(1e-8) is too tight on Linux x86 BLAS and failed Tier 0 CI; use atol=1e-6 to
match the existing chunked-path characterization test.

* ci: run macOS smoke tier on the self-hosted Apple Silicon runner

GitHub-hosted macOS carries a 10x minute multiplier and exhausts the included
Actions minutes too quickly. Move the Tier 1 smoke job onto the self-hosted
Apple Silicon runner ([self-hosted, macOS, ARM64, coreml]) so macOS coverage no
longer consumes hosted minutes. Tier 0 stays on hosted ubuntu (1x).

* ci: fix uv setup for both tiers

astral-sh/setup-uv@v3 was retagged and its old commit garbage-collected, so
codeload 404s when Actions resolves the stale SHA. Bump Tier 0 (ubuntu) to
setup-uv@v7, and drop the action entirely from Tier 1 since the self-hosted
runner already provides uv.

* test(ci): restore golden-image correctness gate on the self-hosted runner

The Tier 2 end-to-end correctness check (real SD1.5 -> Core ML -> image,
gated on SHA/PSNR vs a golden) was dropped during the modernization. With the
breaking 3D-context change, the synthetic smoke and shape/attention
characterization tests no longer cover real-model conversion correctness.

Restore tier2.yml (on the [self-hosted, macOS, ARM64, coreml] runner shared
with Tier 1), the golden-image test, and the e2e workflow. Adapt the pinned
ComfyUI resolution to the requires-comfyui semver tag (vX.Y.Z) instead of a
commit SHA, and re-register the m2 marker. The golden is intentionally not
committed: the first self-hosted run regenerates it under the new 3D contract
and fails for review, per the test's documented bootstrap.

* test(ci): add golden image for SD1.5 seed 42 under the 3D-context contract

Generated by the first Tier 2 self-hosted run after the native diffusers
conversion change. The decoded image is a coherent SD1.5 generation, confirming
the (batch, tokens, hidden) Core ML contract produces correct output
end-to-end. Subsequent runs gate on this golden (SHA-strict, PSNR fallback).

* test(ci): force fresh conversion in Tier 2; drop stale golden

The converter skips conversion when a same-named model already exists, keyed on
conversion parameters but not the conversion code/toolchain. The self-hosted
runner held a pre-existing v1-5 .mlmodelc (4D-context, old toolchain), so the
Tier 2 runs reused it (~8-30s) instead of converting — the gate validated a
stale model, not the new native diffusers path.

Purge the cached Core ML UNets before running so every Tier 2 run does a real
convert -> compile -> sample. Drop the golden generated from the stale cache;
the next run regenerates it from a genuine 3D-contract conversion and fails for
review.

* test(ci): add golden image from a genuine 3D-contract conversion

Regenerated by a Tier 2 run with the model cache purged, so the converter
actually ran (62s, not a cache hit). The fresh model exposes the new 3D
encoder_hidden_states input [1, 77, 768], and its decoded SD1.5 seed-42 image
is byte-identical to the prior baseline — confirming the native diffusers
conversion is behavior-preserving end to end.
2026-05-26 15:46:09 +02:00
aszc 02b6e8ece3 feat: modernize toolchain, refactor core, add tiered CI and opt-in quantization
Modernizes ComfyUI-CoreMLSuite onto Python 3.12 / torch 2.7 / coremltools 9
with a characterization-test safety net. The default conversion path is
unchanged; existing saved workflows produce identical output.

- Toolchain bump (Python 3.12, torch 2.7, coremltools 9, numpy <2) with the
  blocking upstream pins overridden.
- Framework-free logic moved into coreml_suite/core/ (no comfy/coremltools
  imports); old module paths re-export from there.
- Opt-in quantize_nbits dropdown (none|8|6|4) for k-means weight
  palettization; default none is byte-for-byte identical to before.
- Tiered CI: Tier 0 (Linux unit), Tier 1 (macOS-ARM smoke), Tier 2
  (self-hosted Apple Silicon golden-image check on the ANE).
2026-05-25 19:11:49 +02:00
snomiao 7678a07ed5 chore(publish): update GitHub Actions workflow for node publishing
- Added permissions for issue writing
- Updated action version to v1 for publish-node-action
- Added condition to run job only for 'aszc-dev' repository owner
2025-04-01 23:45:31 +02:00
snomiao 43b77e8471 chore(licence-update): Update PyProject Toml - License 2024-08-15 20:37:19 +02:00
aszc-dev c96059ff0b Add basic conversion integration test 2024-07-04 08:44:37 +02:00
aszc-dev 3224d62342 Restructure tests directory 2024-07-04 08:44:37 +02:00
aszc-dev 2fb135df03 Fix set_timestamps for new LCMScheduler implementation 2024-07-04 08:44:37 +02:00
aszc-dev 66e83c2f2f Change syntax to support older Python versions 2024-07-04 08:44:37 +02:00
aszc fb7188e5a2 Update pyproject.toml to test registry workflow 2024-07-03 16:15:37 +02:00
haohaocreates 4096466f8c chore(publish): Add Github Action for Publishing to Comfy Registry 2024-07-03 16:13:36 +02:00
aszc b8c263b763 Update pyproject.toml 2024-07-03 16:08:02 +02:00
haohaocreates 56cff2bd91 chore(pyproject): Add pyproject.toml for Custom Node Registry 2024-07-03 16:08:02 +02:00
Chris Chance 7b3f8fc29e Update ModelSamplingDiscreteLCM to Distilled for latest comfyui 2023-12-01 01:09:14 +01:00
Chris Chance adaecd3f66 Lowered minimum CoreML Size to 256x256 2023-11-28 17:49:43 +01:00
aszc-dev aa60cda09b Add installation using ComfyUI-Manager instructions 2023-11-24 15:10:33 +01:00
aszc-dev 5f7fcd6df3 Add note on SD2.1 to readme 2023-11-24 12:42:33 +01:00
aszc-dev e89cff6d01 Update readme with SDXL info 2023-11-24 12:14:15 +01:00
aszc-dev 9f90083126 Update converter docs and workflows 2023-11-24 12:14:15 +01:00
aszc-dev 5c774ddc5e Remove LCM option from converter for now 2023-11-24 12:14:15 +01:00
aszc-dev b8197c21ef Converting refiner works 2023-11-24 12:14:15 +01:00
aszc-dev 0c78803b25 Base SDXL conversion works 2023-11-24 12:14:15 +01:00
aszc-dev 763ca3961b Handle SDXL config 2023-11-24 12:14:15 +01:00
aszc-dev ae9a9874c5 Add Advanced Sampler node 2023-11-24 12:14:15 +01:00
aszc-dev 67c902f761 Generating SDXL with Core ML Sampler works 2023-11-24 12:14:15 +01:00
aszc-dev bb44b4a35f Link to ComfyUI repo 2023-11-24 12:14:15 +01:00
aszc-dev 46d1124573 Update REAMDE.md (Conversion and LoRA) 2023-11-17 22:55:13 +01:00
aszc-dev ead01c08dd Remove lora.py 2023-11-17 22:55:13 +01:00
aszc-dev b10effc7c2 Add conversion/lora workflows 2023-11-17 22:55:13 +01:00
aszc-dev b1d2e82677 Add peft and omegaconf to requirements 2023-11-17 22:55:13 +01:00
aszc-dev 9f650acb79 Load .yaml config if present 2023-11-17 22:55:13 +01:00
aszc-dev c6d6917827 Setting LoRA model weights works 2023-11-17 22:55:13 +01:00
aszc-dev 63377ebd73 Store lora_params in dict 2023-11-17 22:55:13 +01:00
aszc-dev 42ff10cd43 Add node to load LoRAs 2023-11-17 22:55:13 +01:00
aszc-dev da3a8e13d3 Add logging during conversion 2023-11-17 22:55:13 +01:00
aszc-dev 8092a19173 Enable choosing attention implementation during conversion 2023-11-17 22:55:13 +01:00
aszc-dev 5477e3d71a Remove CLIP loader from nodes 2023-11-17 22:55:13 +01:00
aszc-dev a8d2d6ec46 Move lora related code around, remove clip stuff 2023-11-17 22:55:13 +01:00
aszc-dev 44cffbb8b8 Move load_lora to lora.py 2023-11-17 22:55:13 +01:00
aszc-dev 6907d4910f Remove ckpt loading when loading lora clip 2023-11-17 22:55:13 +01:00
aszc-dev 1930be5c98 Remove CLIP related code 2023-11-17 22:55:13 +01:00
aszc-dev 45be6761d1 Basic conversion + LoRA support works 2023-11-17 22:55:13 +01:00
aszc-dev fc1132a5d5 Fix category for all Core ML nodes 2023-11-17 22:55:13 +01:00
aszc-dev e440f725a4 Specify diffusers and coremltools versions in requirements.txt 2023-11-14 18:43:15 +01:00
aszc-dev f9f25fbeb7 Add LCM info to readme 2023-11-13 13:47:15 +01:00
aszc-dev 4a1359b6b5 Negative optional for LCM 2023-11-13 13:18:13 +01:00
aszc-dev 971e60aa09 Rearrange LCM code 2023-11-11 04:15:59 +01:00
aszc-dev 8bcdeab234 Core ML Sampler supports LCM 2023-11-11 03:11:35 +01:00
aszc-dev c9e403b1d8 WIP: LCM Scheduler refactor 2023-11-11 00:19:16 +01:00
aszc-dev 7492f0b486 Extract lcm sampler from lcm sampling node 2023-11-10 13:24:06 +01:00
aszc-dev c09221945d Remove dead code from LCM Sampler 2023-11-10 03:00:42 +01:00
aszc-dev 6864c233e3 ControlNet works for LCM 2023-11-10 02:07:11 +01:00
aszc-dev 6ccf41e5c9 Refactor LCM sampling 2023-11-09 18:02:37 +01:00
aszc-dev 73aa2d11d3 Download scheduler config from repo 2023-11-09 00:07:31 +01:00
aszc-dev 4c438e1ee6 Leverage Comfy's mechanisms to enable LCM ControlNet support 2023-11-09 00:07:30 +01:00
aszc-dev fa0735746c Refactor model config 2023-11-09 00:04:28 +01:00
aszc-dev c26099b334 Add CoreMLInputs to handle inputs 2023-11-08 22:19:41 +01:00
aszc-dev 27f1a19131 Refactor CoreMLModelWrapper 2023-11-08 21:14:55 +01:00
aszc-dev 701443f59e Wrapped Core ML Model is now diffusion_model attribute of BaseModel 2023-11-08 17:51:27 +01:00
aszc-dev 6d095a67a2 Add diffusers to requirements 2023-11-06 23:30:27 +01:00
aszc-dev bb73e686a0 Add newlines 2023-11-06 23:21:51 +01:00
Robert Dean 967ab7f269 Update requirements.txt
Added overrides decorator
2023-11-06 18:50:14 +01:00
aszc c51d9041a4 Merge pull request #5 from aszc-dev/lcm
LCM Support
2023-11-03 02:04:56 +01:00
aszc-dev eeae4bd6e3 Adjust default values for LCM nodes 2023-11-03 01:29:50 +01:00
aszc-dev 1ebd9e72ae Remove Simple LCM Sampler 2023-11-03 01:29:50 +01:00
aszc-dev c01c60e3c1 Add progress bar and preview to LCM 2023-11-03 01:29:50 +01:00
aszc-dev 44a380ffdf img2img works 2023-11-03 01:29:32 +01:00
aszc-dev e22d8187cd Add more advanced LCM Sampler 2023-11-03 01:28:34 +01:00
aszc-dev 1937f39cca Add support for CN models to LCM 2023-11-03 01:27:48 +01:00
aszc-dev b90591dfd4 Add support for controlnet to LCM converter 2023-11-03 01:27:48 +01:00
aszc-dev 1aa5a19b2a Simplify LCM Sampler 2023-11-03 01:27:48 +01:00
aszc-dev 8a814b7a56 Fix LCM Sampler 2023-11-03 01:27:48 +01:00
aszc-dev 213088241d LCM Converter works 2023-11-03 01:27:48 +01:00
aszc-dev 9d509ad8f4 WIP: LCM 2023-11-03 01:27:48 +01:00
aszc-dev 0092ad5e75 Prepare LCM Model Wrapper 2023-11-03 01:27:48 +01:00
aszc-dev 99a0a9996d Fix cn chunking 2023-11-03 00:31:46 +01:00
aszc-dev db0aea3d9c Fix chunk_inputs 2023-11-01 22:16:47 +01:00
aszc-dev dfdc1bf520 Fix cn chunking 2023-11-01 01:08:14 +01:00
aszc-dev d63df5b62f Remove the controlnet note in readme 2023-10-31 22:05:19 +01:00
aszc-dev 901ea6da16 Simplify no_control 2023-10-31 21:57:32 +01:00
aszc-dev dd438f66cc Fix controlnet residuals chunking 2023-10-31 21:40:11 +01:00
aszc-dev 41797203d7 Improve chunking and padding 2023-10-31 02:49:35 +01:00
aszc-dev 4d83603c98 Chunking works for ControlNet 2023-10-30 18:16:51 +01:00
aszc-dev 8a3e9332e1 Chunk and pad batches 2023-10-30 16:17:18 +01:00
aszc-dev 6319d2aedb Add model adapter for unstable compatibility 2023-10-30 11:49:07 +01:00
aszc-dev d0629b4efc Rearrange stuff 2023-10-30 11:15:42 +01:00
aszc-dev c043e1f9aa Update ControlNet workflow 2023-10-30 01:01:59 +01:00
aszc 133f943472 Merge pull request #2 from aszc-dev/dev
Make Core ML models incompatibile with default nodes
2023-10-30 00:51:12 +01:00
53 changed files with 4255 additions and 667 deletions
+25
View File
@@ -0,0 +1,25 @@
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 }}
+27
View File
@@ -0,0 +1,27 @@
name: Tier 0 — Unit (Linux)
on:
push:
branches: [main]
pull_request:
# Deps are resolved from pyproject.toml via uv, so the toolchain pins live in
# one place. Tier 0 must run without ComfyUI; the in-tree purity gate
# (tests/unit/test_tier0_purity.py) enforces that the suite hasn't started
# leaking framework imports.
jobs:
unit:
runs-on: ubuntu-latest
timeout-minutes: 10
steps:
- uses: actions/checkout@v4
- uses: astral-sh/setup-uv@v7
with:
enable-cache: true
- name: uv sync
run: uv sync --no-install-project
- name: Run Tier 0
run: uv run pytest -m unit tests/ -v
+23
View File
@@ -0,0 +1,23 @@
name: Tier 1 — Smoke (macOS self-hosted)
# macOS smoke tests run on the self-hosted Apple Silicon runner instead of
# GitHub-hosted macOS (10x minute multiplier), which exhausts the included
# Actions minutes too quickly.
on:
push:
branches: [main]
pull_request:
jobs:
smoke:
runs-on: [self-hosted, macOS, ARM64, coreml]
timeout-minutes: 20
steps:
- uses: actions/checkout@v4
# The self-hosted runner provides uv; no setup-uv action needed.
- 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
+134
View File
@@ -0,0 +1,134 @@
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 requires-comfyui version tag + 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
# requires-comfyui is a semver constraint (e.g. ">=0.3.27"); pin the
# gate to the matching ComfyUI release tag (vX.Y.Z).
VERSION="$(sed -nE 's/^requires-comfyui *= *"[^0-9]*([0-9]+\.[0-9]+\.[0-9]+).*/\1/p' pyproject.toml)"
if [ -z "$VERSION" ]; then echo "could not parse requires-comfyui from pyproject.toml"; exit 1; fi
echo "COMFY_MODE=pinned" >> "$GITHUB_ENV"
echo "COMFY_REF=v$VERSION" >> "$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: Purge cached Core ML UNets (force fresh conversion)
# The converter skips when a model of the same name already exists. That
# cache key is conversion *parameters* only, not the conversion code or
# toolchain — so a stale model would let a conversion regression pass.
# Clear it so every Tier 2 run exercises the full convert -> compile ->
# sample path end to end.
run: |
rm -rf "$COMFY_DIR"/models/unet/*.mlpackage "$COMFY_DIR"/models/unet/*.mlmodelc || true
- name: Run Tier 2 (m2 marker)
# Drives the Core ML Converter node, which converts the UNet from the
# checkpoint on every run (cache purged above).
run: uv run --no-sync pytest -m m2 tests/ -v
- name: Stop ComfyUI server
if: always()
run: pkill -f "main.py.*8188" || true
+3 -1
View File
@@ -1,3 +1,5 @@
playground/
experiments/
__pycache__/
models/
.venv/
test_results/
+72 -52
View File
@@ -8,8 +8,8 @@ These models are designed to leverage the Apple Neural Engine (ANE) on Apple Sil
thereby enhancing your workflows and improving performance.
If you're not sure how to obtain these models, you can download them
[here](https://huggingface.co/coreml-community) or convert your own models using
[coremltools](https://github.com/apple/ml-stable-diffusion).
[here](https://huggingface.co/coreml-community) or convert your own checkpoints
directly with the conversion nodes in this suite (see [How to use](#how-to-use)).
In simple terms, think of Core ML models as a tool that can help your ComfyUI work faster and more efficiently.
For instance, during my tests on an M2 Pro 32GB machine,
@@ -81,6 +81,29 @@ These custom nodes come with a host of features, including:
> [!NOTE]
> This repository will continue to be updated with more nodes and features over time.
## Conversion & Acknowledgements
The Core ML conversion pipeline in this repository began as an adaptation of
Apple's [ml-stable-diffusion](https://github.com/apple/ml-stable-diffusion),
which pioneered running Stable Diffusion on the Apple Neural Engine. The
implementation has since diverged and no longer depends on that package:
- UNet conversion runs natively on `diffusers`' `UNet2DConditionModel`.
- The ANE-friendly attention path (`SPLIT_EINSUM`, `SPLIT_EINSUM_V2`) is
reimplemented as standalone `diffusers` attention processors.
- The toolchain tracks current ComfyUI (NumPy 2, Torch 2.7, coremltools 9,
Python 3.12).
The goal is to keep iterating on these methods independently and to explore
support beyond SD1.5.
> [!IMPORTANT]
> **Breaking change in 2.0.0.** The converted Core ML UNet now takes
> `encoder_hidden_states` in the native `diffusers` layout
> `(batch, tokens, hidden)` instead of the previous
> `(batch, hidden, 1, tokens)`. Core ML models converted with earlier versions
> are not compatible with 2.0.0 and must be re-converted.
## Installation
### Using ComfyUI-Manager
@@ -170,8 +193,8 @@ the node name, so if the model already exists, the node will not convert it agai
- **ckpt_name**: The name of the checkpoint to convert. This should be the name of the checkpoint file stored in the
`models/checkpoints` directory.
- **model_version**: Whether the model is based on SD1.5 or SDXL.
- **height**: The desired height of the image generated by the model. The default is 512. Must be a multiple of 8.
- **width**: The desired width of the image generated by the model. The default is 512. Must be a multiple of 8.
- **height**: The desired height of the image generated by the model. The default is 512. Any positive multiple of 8 is accepted.
- **width**: The desired width of the image generated by the model. The default is 512. Any positive multiple of 8 is accepted.
- **batch_size**: The batch size of generated images. If you're planning to generate batches of images, you can try
increasing this value to speed up the generation process. The default is 1.
- **attention_implementation**: The attention implementation used when converting the model. Choose SPLIT_EINSUM or
@@ -284,8 +307,8 @@ can use any CLIP or VAE model as long as it's compatible with Stable Diffusion v
1. **Loading text encoder (CLIP) and VAE models separately**
- This workflow uses CLIP and VAE models available
[here](https://huggingface.co/runwayml/stable-diffusion-v1-5/blob/main/text_encoder/model.safetensors) and
[here](https://huggingface.co/runwayml/stable-diffusion-v1-5/blob/main/vae/diffusion_pytorch_model.safetensors).
[here](https://huggingface.co/stable-diffusion-v1-5/stable-diffusion-v1-5/resolve/main/text_encoder/model.safetensors) and
[here](https://huggingface.co/stable-diffusion-v1-5/stable-diffusion-v1-5/resolve/main/vae/diffusion_pytorch_model.safetensors).
Once downloaded, place the models in the`models/clip` and `models/vae` directories respectively.
- The Core ML UNet model is available
[here](https://huggingface.co/coreml-community/coreml-stable-diffusion-v1-5_cn/blob/main/split_einsum/stable-diffusion-_v1-5_split-einsum_cn.zip).
@@ -293,7 +316,7 @@ can use any CLIP or VAE model as long as it's compatible with Stable Diffusion v
![coreml-unet+clip+vae](./assets/unet+sampler+clip+vae.png?raw=true)
2. **Loading text encoder (CLIP) and VAE models from checkpoint file**
- This workflow loads the CLIP and VAE models from the checkpoint file available
[here](https://huggingface.co/runwayml/stable-diffusion-v1-5/blob/main/v1-5-pruned-emaonly.safetensors).
[here](https://huggingface.co/stable-diffusion-v1-5/stable-diffusion-v1-5/resolve/main/v1-5-pruned-emaonly.safetensors).
Once downloaded, place the model in the`models/checkpoints` directory.
- The Core ML UNet model is available
[here](https://huggingface.co/coreml-community/coreml-stable-diffusion-v1-5_cn/blob/main/split_einsum/stable-diffusion-_v1-5_split-einsum_cn.zip).
@@ -370,62 +393,59 @@ The models used in this workflow are available at the following links:
![sdxl](./assets/sdxl_conversion.png?raw=true)
## 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.
## Limitations
- Core ML models are fixed in terms of their inputs and outputs.
This means you'll need to use latent images of the same size as the input of the model (512x512 is the default for
SD1.5).
However, you can convert the model to a different input size using tools available
in the [apple/ml-stable-diffusion](https://github.com/apple/ml-stable-diffusion) repository.
However, you can re-convert the model to a different input size using the
conversion nodes in this suite (set the desired width and height).
- SD2.1 models are not supported.
[^1]:
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
+4
View File
@@ -0,0 +1,4 @@
"""Top-level conftest: prevent pytest from importing the repo-root
__init__.py (the ComfyUI custom-node entry point pulls in comfy + nodes,
which breaks the Tier-0 'no-framework' promise)."""
collect_ignore = ["__init__.py"]
+18
View File
@@ -0,0 +1,18 @@
# Toolchain ceiling for installing a floating-latest ComfyUI's requirements.txt
# in the Tier 2 nightly canary (.github/workflows/tier2.yml, latest mode).
#
# ComfyUI's requirements.txt requests bare `torch`/`torchvision`/`torchaudio`
# and `numpy>=1.25.0`, which would float past the versions coremltools 9 /
# apple-ml-stable-diffusion have been validated against.
# These constraints cap the resolution so the canary keeps testing the same
# toolchain the suite actually ships.
#
# If upstream ComfyUI ever hard-requires something beyond these bounds, the
# install FAILS — and that failure is the signal we want: it means the host
# outgrew the pinned toolchain and coremltools / ml-stable-diffusion need a
# deliberate bump, not a silent float.
torch>=2.7,<2.8
torchvision>=0.22,<0.23
torchaudio>=2.7,<2.8
numpy>=1.25,<2
coremltools>=9,<10
+5
View File
@@ -0,0 +1,5 @@
ATTENTION_IMPLEMENTATIONS = (
"SPLIT_EINSUM",
"SPLIT_EINSUM_V2",
"ORIGINAL",
)
+1 -8
View File
@@ -1,17 +1,10 @@
from enum import Enum
import torch
from comfy import supported_models_base
from comfy import latent_formats
from comfy.model_detection import convert_config
class ModelVersion(Enum):
SD15 = "sd15"
SDXL = "sdxl"
SDXL_REFINER = "sdxl_refiner"
LCM = "lcm"
from coreml_suite.model_version import ModelVersion
config_map = {
+13 -61
View File
@@ -1,62 +1,14 @@
from itertools import chain
from math import ceil
"""Compatibility shim — re-exports from coreml_suite.core.controlnet."""
from coreml_suite.core.controlnet import (
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
__all__ = [
"chunk_control",
"expand_inputs",
"extract_residual_kwargs",
"no_control",
]
+9
View File
@@ -0,0 +1,9 @@
"""Core ML conversion helpers.
The conversion approach originates from Apple's ml-stable-diffusion
(https://github.com/apple/ml-stable-diffusion). This implementation has since
diverged: it runs natively on diffusers' UNet2DConditionModel with its own
SPLIT_EINSUM / SPLIT_EINSUM_V2 attention processors and no longer depends on
that package. The intent is to keep iterating on these methods independently
while tracking current tooling.
"""
+239
View File
@@ -0,0 +1,239 @@
import logging
import torch
logger = logging.getLogger(__name__)
CHUNK_SIZE = 512
def apply_attention_implementation(unet, attention_implementation):
if attention_implementation == "ORIGINAL":
return unet
if attention_implementation == "SPLIT_EINSUM":
unet.set_attn_processor(SplitEinsumAttnProcessor())
return unet
if attention_implementation == "SPLIT_EINSUM_V2":
unet.set_attn_processor(SplitEinsumV2AttnProcessor())
return unet
raise ValueError(f"Unsupported attention implementation: {attention_implementation}")
class SplitEinsumAttnProcessor:
def __call__(
self,
attn,
hidden_states,
encoder_hidden_states=None,
attention_mask=None,
temb=None,
*args,
**kwargs,
):
return _attention_forward(
attn,
hidden_states,
encoder_hidden_states,
attention_mask,
temb,
split_einsum,
)
class SplitEinsumV2AttnProcessor:
def __call__(
self,
attn,
hidden_states,
encoder_hidden_states=None,
attention_mask=None,
temb=None,
*args,
**kwargs,
):
return _attention_forward(
attn,
hidden_states,
encoder_hidden_states,
attention_mask,
temb,
split_einsum_v2,
)
def _attention_forward(
attn,
hidden_states,
encoder_hidden_states,
attention_mask,
temb,
attention_fn,
):
residual = hidden_states
if attn.spatial_norm is not None:
hidden_states = attn.spatial_norm(hidden_states, temb)
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
else:
batch_size, _, channel = hidden_states.shape
height = None
width = None
batch_size, key_sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
if attention_mask is not None:
attention_mask = attn.prepare_attention_mask(
attention_mask,
key_sequence_length,
batch_size,
)
attention_mask = _prepare_split_einsum_mask(
attention_mask,
batch_size,
attn.heads,
key_sequence_length,
)
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
elif attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
batch_size = query.shape[0]
dim_head = attn.inner_kv_dim // attn.heads
query = _linear_projection_to_bchw(query)
key = _linear_projection_to_bchw(key)
value = _linear_projection_to_bchw(value)
hidden_states = attention_fn(
query,
key,
value,
attention_mask,
attn.heads,
dim_head,
)
hidden_states = hidden_states.squeeze(2).transpose(1, 2)
hidden_states = hidden_states.reshape(batch_size, -1, attn.inner_dim)
hidden_states = attn.to_out[0](hidden_states)
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(
batch_size,
channel,
height,
width,
)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
def split_einsum(q, k, v, mask, heads, dim_head):
q_heads = _split_heads(q, heads, dim_head)
k = k.transpose(1, 3)
k_heads = [
k[:, :, :, head_idx * dim_head : (head_idx + 1) * dim_head]
for head_idx in range(heads)
]
v_heads = _split_heads(v, heads, dim_head)
weights = [
torch.einsum("bchq,bkhc->bkhq", query, key) * (dim_head**-0.5)
for query, key in zip(q_heads, k_heads)
]
if mask is not None:
weights = [weight + mask for weight in weights]
weights = [weight.softmax(dim=1) for weight in weights]
outputs = [
torch.einsum("bkhq,bchk->bchq", weight, value)
for weight, value in zip(weights, v_heads)
]
return torch.cat(outputs, dim=1)
def split_einsum_v2(q, k, v, mask, heads, dim_head):
query_length = q.size(3)
num_chunks = query_length // CHUNK_SIZE
if num_chunks == 0:
logger.info(
"SPLIT_EINSUM_V2 query sequence is shorter than %s; using SPLIT_EINSUM.",
CHUNK_SIZE,
)
return split_einsum(q, k, v, mask, heads, dim_head)
q_heads = _split_heads(q, heads, dim_head)
q_chunks = [
[
head[..., chunk_idx * CHUNK_SIZE : (chunk_idx + 1) * CHUNK_SIZE]
for chunk_idx in range(num_chunks)
]
for head in q_heads
]
k = k.transpose(1, 3)
k_heads = [
k[:, :, :, head_idx * dim_head : (head_idx + 1) * dim_head]
for head_idx in range(heads)
]
v_heads = _split_heads(v, heads, dim_head)
head_outputs = []
for query_chunks, key, value in zip(q_chunks, k_heads, v_heads):
chunk_outputs = []
for query_chunk in query_chunks:
weights = torch.einsum("bchq,bkhc->bkhq", query_chunk, key)
weights = weights * (dim_head**-0.5)
if mask is not None:
weights = weights + mask
weights = weights.softmax(dim=1)
chunk_outputs.append(torch.einsum("bkhq,bchk->bchq", weights, value))
head_outputs.append(torch.cat(chunk_outputs, dim=3))
return torch.cat(head_outputs, dim=1)
def _split_heads(x, heads, dim_head):
return [
x[:, head_idx * dim_head : (head_idx + 1) * dim_head, :, :]
for head_idx in range(heads)
]
def _linear_projection_to_bchw(x):
return x.transpose(1, 2).unsqueeze(2)
def _prepare_split_einsum_mask(mask, batch_size, heads, key_sequence_length):
if mask.ndim == 2:
mask = mask[:, None, :]
if mask.shape[0] == batch_size * heads:
mask = mask.reshape(batch_size, heads, -1, key_sequence_length)
mask = mask[:, 0]
if mask.ndim == 3:
mask = mask[:, :, None, None]
return mask
+20
View File
@@ -0,0 +1,20 @@
def conv2d_output_shape(height, width, conv):
"""Return the spatial output shape for a torch.nn.Conv2d-like module."""
kernel_h, kernel_w = _pair(conv.kernel_size)
stride_h, stride_w = _pair(conv.stride)
pad_h, pad_w = _pair(conv.padding)
dilation_h, dilation_w = _pair(conv.dilation)
out_h = _conv_output_dim(height, kernel_h, stride_h, pad_h, dilation_h)
out_w = _conv_output_dim(width, kernel_w, stride_w, pad_w, dilation_w)
return out_h, out_w
def _conv_output_dim(size, kernel, stride, padding, dilation):
return ((size + (2 * padding) - (dilation * (kernel - 1)) - 1) // stride) + 1
def _pair(value):
if isinstance(value, tuple):
return value
return value, value
+61
View File
@@ -0,0 +1,61 @@
from types import MethodType
from diffusers.models.transformers.transformer_2d import Transformer2DModel
def prepare_unet_for_coreml_trace(unet):
for module in unet.modules():
if isinstance(module, Transformer2DModel):
module._operate_on_continuous_inputs = MethodType(
_operate_on_continuous_inputs,
module,
)
module._get_output_for_continuous_inputs = MethodType(
_get_output_for_continuous_inputs,
module,
)
return unet
def _operate_on_continuous_inputs(self, hidden_states):
hidden_states = self.norm(hidden_states)
if not self.use_linear_projection:
hidden_states = self.proj_in(hidden_states)
inner_dim = self.inner_dim
hidden_states = hidden_states.flatten(2).transpose(1, 2)
else:
inner_dim = hidden_states.shape[1]
hidden_states = hidden_states.flatten(2).transpose(1, 2)
hidden_states = self.proj_in(hidden_states)
return hidden_states, inner_dim
def _get_output_for_continuous_inputs(
self,
hidden_states,
residual,
batch_size,
height,
width,
inner_dim,
):
if not self.use_linear_projection:
hidden_states = hidden_states.transpose(1, 2).reshape(
batch_size,
inner_dim,
height,
width,
)
hidden_states = self.proj_out(hidden_states)
else:
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.transpose(1, 2).reshape(
batch_size,
inner_dim,
height,
width,
)
return hidden_states + residual
+54
View File
@@ -0,0 +1,54 @@
import torch
class CoreMLUNetWrapper(torch.nn.Module):
"""Adapt diffusers UNet inputs to CoreMLSuite's stable Core ML contract."""
def __init__(self, unet, model_version):
super().__init__()
self.unet = unet
self.model_version = model_version
def forward(self, sample, timestep, encoder_hidden_states, *extra_inputs):
input_index = 0
timestep_cond = None
if self._is_lcm:
timestep_cond = extra_inputs[input_index]
input_index += 1
added_cond_kwargs = None
if self._is_sdxl:
time_ids = extra_inputs[input_index]
text_embeds = extra_inputs[input_index + 1]
input_index += 2
added_cond_kwargs = {
"time_ids": time_ids,
"text_embeds": text_embeds,
}
additional_residuals = extra_inputs[input_index:]
down_residuals = None
mid_residual = None
if additional_residuals:
down_residuals = tuple(additional_residuals[:-1])
mid_residual = additional_residuals[-1]
outputs = self.unet(
sample,
timestep,
encoder_hidden_states=encoder_hidden_states,
timestep_cond=timestep_cond,
added_cond_kwargs=added_cond_kwargs,
down_block_additional_residuals=down_residuals,
mid_block_additional_residual=mid_residual,
return_dict=False,
)
return outputs[0]
@property
def _is_lcm(self):
return self.model_version.name == "LCM"
@property
def _is_sdxl(self):
return self.model_version.name in {"SDXL", "SDXL_REFINER"}
+79 -119
View File
@@ -1,72 +1,38 @@
import gc
import os
import shutil
import time
from typing import Union
import coremltools as ct
import numpy as np
import python_coreml_stable_diffusion.unet
import torch
from diffusers import (
StableDiffusionPipeline,
LatentConsistencyModelPipeline,
StableDiffusionXLPipeline,
)
from python_coreml_stable_diffusion.unet import (
UNet2DConditionModel,
UNet2DConditionModelXL,
AttentionImplementations,
)
from diffusers import UNet2DConditionModel
from coreml_suite.config import ModelVersion
from coreml_suite.lcm.unet import UNet2DConditionModelLCM
from coreml_suite.attention import ATTENTION_IMPLEMENTATIONS
from coreml_suite.conversion.attention import apply_attention_implementation
from coreml_suite.conversion.shapes import conv2d_output_shape
from coreml_suite.conversion.trace import prepare_unet_for_coreml_trace
from coreml_suite.conversion.unet import CoreMLUNetWrapper
from coreml_suite.logger import logger
from folder_paths import get_folder_paths
from coreml_suite.model_version import ModelVersion
DEFAULT_TRACE_TIMESTEP = 999.0
TEXT_TOKEN_SEQUENCE_LENGTH = 77
class StableDiffusionLCMPipeline(LatentConsistencyModelPipeline):
pass
MODEL_TYPE_TO_UNET_CLS = {
ModelVersion.SD15: UNet2DConditionModel,
ModelVersion.SDXL: UNet2DConditionModelXL,
ModelVersion.LCM: UNet2DConditionModelLCM,
}
MODEL_TYPE_TO_PIPE_CLS = {
ModelVersion.SD15: StableDiffusionPipeline,
ModelVersion.SDXL: StableDiffusionXLPipeline,
ModelVersion.LCM: StableDiffusionLCMPipeline,
}
def get_unet(model_type: ModelVersion, ref_pipe):
ref_unet = ref_pipe.unet
unet_cls = MODEL_TYPE_TO_UNET_CLS[model_type]
cml_unet = unet_cls.from_config(ref_unet.config).eval()
cml_unet.load_state_dict(ref_unet.state_dict(), strict=False)
return cml_unet
def get_encoder_hidden_states_shape(ref_pipe, batch_size):
text_encoder = (
ref_pipe.text_encoder_2
if hasattr(ref_pipe, "text_encoder_2")
else ref_pipe.text_encoder
def get_unet(model_version: ModelVersion, ref_unet, attention_implementation):
ref_unet = prepare_unet_for_coreml_trace(ref_unet)
unet = apply_attention_implementation(
ref_unet.eval(),
attention_implementation,
)
return CoreMLUNetWrapper(unet, model_version)
text_token_sequence_length = text_encoder.config.max_position_embeddings
hidden_size = (text_encoder.config.hidden_size,)
def get_encoder_hidden_states_shape(ref_unet, batch_size):
encoder_hidden_states_shape = (
batch_size,
ref_pipe.unet.config.cross_attention_dim or hidden_size,
1,
text_token_sequence_length,
TEXT_TOKEN_SEQUENCE_LENGTH,
ref_unet.config.cross_attention_dim,
)
return encoder_hidden_states_shape
@@ -122,38 +88,21 @@ def convert_to_coreml(
def get_out_path(submodule_name, model_name):
from folder_paths import get_folder_paths
fname = f"{model_name}_{submodule_name}.mlpackage"
unet_path = get_folder_paths(submodule_name)[0]
out_path = os.path.join(unet_path, fname)
return out_path
def compile_coreml_model(source_model_path, output_dir, final_name):
"""Compiles Core ML models using the coremlcompiler utility from Xcode toolchain"""
target_path = os.path.join(output_dir, f"{final_name}.mlmodelc")
if os.path.exists(target_path):
logger.warning(f"Found existing compiled model at {target_path}! Skipping..")
return target_path
logger.info(f"Compiling {source_model_path}")
source_model_name = os.path.basename(os.path.splitext(source_model_path)[0])
os.system(f"xcrun coremlcompiler compile {source_model_path} {output_dir}")
compiled_output = os.path.join(output_dir, f"{source_model_name}.mlmodelc")
shutil.move(compiled_output, target_path)
return target_path
def get_sample_input(batch_size, encoder_hidden_states_shape, sample_shape, scheduler):
def get_sample_input(batch_size, encoder_hidden_states_shape, sample_shape):
sample_unet_inputs = dict(
[
("sample", torch.rand(*sample_shape)),
(
"timestep",
torch.tensor([scheduler.timesteps[0].item()] * batch_size).to(
torch.float32
),
torch.tensor([DEFAULT_TRACE_TIMESTEP] * batch_size).to(torch.float32),
),
("encoder_hidden_states", torch.rand(*encoder_hidden_states_shape)),
]
@@ -166,7 +115,7 @@ def lcm_inputs(sample_unet_inputs):
return {"timestep_cond": torch.randn(batch_size, 256).to(torch.float32)}
def sdxl_inputs(sample_unet_inputs, ref_pipe):
def sdxl_inputs(sample_unet_inputs, ref_unet, model_version):
sample_shape = sample_unet_inputs["sample"].shape
batch_size = sample_shape[0]
h = sample_shape[2] * 8
@@ -174,10 +123,7 @@ def sdxl_inputs(sample_unet_inputs, ref_pipe):
original_size = (h, w)
crops_coords_top_left = (0, 0)
is_refiner = (
hasattr(ref_pipe.config, "requires_aesthetics_score")
and ref_pipe.config.requires_aesthetics_score
)
is_refiner = model_version == ModelVersion.SDXL_REFINER
if is_refiner:
aesthetic_score = (6.0,)
@@ -187,7 +133,7 @@ def sdxl_inputs(sample_unet_inputs, ref_pipe):
time_ids_list = list(original_size + crops_coords_top_left + target_size)
time_ids = torch.tensor(time_ids_list).repeat(batch_size, 1).to(torch.int64)
text_embeds_shape = (batch_size, ref_pipe.text_encoder_2.config.hidden_size)
text_embeds_shape = (batch_size, get_sdxl_text_embeds_dim(ref_unet, len(time_ids_list)))
return {
"time_ids": time_ids,
@@ -195,21 +141,25 @@ def sdxl_inputs(sample_unet_inputs, ref_pipe):
}
def get_sdxl_text_embeds_dim(ref_unet, time_ids_dim):
projection_dim = ref_unet.config.projection_class_embeddings_input_dim
time_embed_dim = ref_unet.config.addition_time_embed_dim
return projection_dim - (time_ids_dim * time_embed_dim)
def get_inputs_spec(inputs):
inputs_spec = {k: (v.shape, v.dtype) for k, v in inputs.items()}
return inputs_spec
def add_cnet_support(sample_shape, reference_unet):
from python_coreml_stable_diffusion.unet import calculate_conv2d_output_shape
additional_residuals_shapes = []
batch_size = sample_shape[0]
h, w = sample_shape[2:]
# conv_in
out_h, out_w = calculate_conv2d_output_shape(
out_h, out_w = conv2d_output_shape(
h,
w,
reference_unet.conv_in,
@@ -226,9 +176,7 @@ def add_cnet_support(sample_shape, reference_unet):
]
if hasattr(down_block, "downsamplers") and down_block.downsamplers is not None:
for downsampler in down_block.downsamplers:
out_h, out_w = calculate_conv2d_output_shape(
out_h, out_w, downsampler.conv
)
out_h, out_w = conv2d_output_shape(out_h, out_w, downsampler.conv)
additional_residuals_shapes.append(
(
batch_size,
@@ -252,15 +200,16 @@ def add_cnet_support(sample_shape, reference_unet):
def convert_unet(
ref_pipe,
ref_unet,
model_version: ModelVersion,
unet_out_path: str,
batch_size: int = 1,
sample_size: tuple[int, int] = (64, 64),
controlnet_support: bool = False,
attention_implementation: str = ATTENTION_IMPLEMENTATIONS[0],
quantize_nbits: str = "none",
):
coreml_unet = get_unet(model_version, ref_pipe)
ref_unet = ref_pipe.unet
coreml_unet = get_unet(model_version, ref_unet, attention_implementation)
sample_shape = (
batch_size, # B
@@ -269,20 +218,17 @@ def convert_unet(
sample_size[1], # W
)
encoder_hidden_states_shape = get_encoder_hidden_states_shape(ref_pipe, batch_size)
scheduler = ref_pipe.scheduler
scheduler.set_timesteps(50)
encoder_hidden_states_shape = get_encoder_hidden_states_shape(ref_unet, batch_size)
sample_inputs = get_sample_input(
batch_size, encoder_hidden_states_shape, sample_shape, scheduler
batch_size, encoder_hidden_states_shape, sample_shape
)
if model_version == ModelVersion.LCM:
sample_inputs |= lcm_inputs(sample_inputs)
if model_version == ModelVersion.SDXL:
sample_inputs |= sdxl_inputs(sample_inputs, ref_pipe)
if model_version in {ModelVersion.SDXL, ModelVersion.SDXL_REFINER}:
sample_inputs |= sdxl_inputs(sample_inputs, ref_unet, model_version)
if controlnet_support:
sample_inputs |= add_cnet_support(sample_shape, ref_unet)
@@ -305,6 +251,24 @@ 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}")
@@ -316,47 +280,43 @@ def convert(
batch_size: int = 1,
sample_size: tuple[int, int] = (64, 64),
controlnet_support: bool = False,
lora_weights: list[tuple[Union[str, os.PathLike], float]] = None,
attn_impl: str = AttentionImplementations.SPLIT_EINSUM.name,
lora_weights: list[tuple[str | os.PathLike, float]] = None,
attn_impl: str = ATTENTION_IMPLEMENTATIONS[0],
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..")
return
python_coreml_stable_diffusion.unet.ATTENTION_IMPLEMENTATION_IN_EFFECT = (
AttentionImplementations(attn_impl)
)
ref_pipe = get_pipeline(ckpt_path, config_path, model_version)
if attn_impl not in ATTENTION_IMPLEMENTATIONS:
raise ValueError(
f"Unsupported attention implementation {attn_impl!r}. "
f"Expected one of {ATTENTION_IMPLEMENTATIONS}."
)
ref_unet = load_unet(ckpt_path, config_path)
for i, lora_weight in enumerate(lora_weights or []):
lora_path, strength = lora_weight
adapter_name = f"lora_{i}"
ref_pipe.load_lora_weights(lora_path, adapter_name=adapter_name)
ref_pipe.set_adapters([adapter_name], adapter_weights=[strength])
ref_pipe.fuse_lora()
ref_unet.load_lora_adapter(lora_path, adapter_name=adapter_name)
ref_unet.set_adapters([adapter_name], weights=[strength])
ref_unet.fuse_lora()
convert_unet(
ref_pipe,
ref_unet,
model_version,
unet_out_path,
batch_size,
sample_size,
controlnet_support,
attention_implementation=attn_impl,
quantize_nbits=quantize_nbits,
)
def get_pipeline(ckpt_path, config_path, model_version):
pipe_cls = MODEL_TYPE_TO_PIPE_CLS[model_version]
ref_pipe = pipe_cls.from_single_file(ckpt_path, original_config_file=config_path)
return ref_pipe
def compile_model(out_path, out_name, submodule_name):
# Compile the model
target_path = compile_coreml_model(
out_path, get_folder_paths(submodule_name)[0], f"{out_name}_{submodule_name}"
def load_unet(ckpt_path, config_path):
return UNet2DConditionModel.from_single_file(
ckpt_path,
original_config=config_path,
)
logger.info(f"Compiled {out_path} to {target_path}")
return target_path
+10
View File
@@ -0,0 +1,10 @@
"""Framework-free pure-logic core of ComfyUI-CoreMLSuite.
Modules under this package must NOT import `comfy`, `coremltools`,
`python_coreml_stable_diffusion`, `folder_paths`, `nodes`, or any other
ComfyUI / Apple runtime. Only `numpy` and `torch` are allowed.
The thin adapters in `coreml_suite.{latents,controlnet,models}` keep the
old public import paths working so `coreml_suite/nodes.py` and downstream
ComfyUI workflows are unchanged.
"""
+67
View File
@@ -0,0 +1,67 @@
"""Pure helpers around the ControlNet residual inputs of the Core ML UNet.
Re-exported by coreml_suite.controlnet. Characterization tests cover
shapes, dtype (fp16), and zero-fill fallback.
"""
from itertools import chain
from math import ceil
import numpy as np
import torch
from coreml_suite.core.latents import chunk_batch
def expand_inputs(inputs):
expanded = inputs.copy()
for k, v in inputs.items():
if isinstance(v, np.ndarray):
expanded[k] = np.concatenate([v] * 2) if v.shape[0] == 1 else v
elif isinstance(v, torch.Tensor):
expanded[k] = torch.cat([v] * 2) if v.shape[0] == 1 else v
elif isinstance(v, list):
expanded[k] = v * 2 if len(v) == 1 else v
elif isinstance(v, dict):
expand_inputs(v)
return expanded
def extract_residual_kwargs(expected_inputs, control):
if "additional_residual_0" not in expected_inputs.keys():
return {}
if control is None:
return no_control(expected_inputs)
residual_kwargs = {
"additional_residual_{}".format(i): r.cpu().numpy().astype(np.float16)
for i, r in enumerate(chain(control["output"], control["middle"]))
}
return residual_kwargs
def no_control(expected_inputs):
shapes_dict = {
k: v["shape"] for k, v in expected_inputs.items() if k.startswith("additional")
}
residual_kwargs = {
k: torch.zeros(*shape).cpu().numpy().astype(dtype=np.float16)
for k, shape in shapes_dict.items()
}
return residual_kwargs
def chunk_control(cn, target_size):
if cn is None:
return [None] * target_size
num_chunks = ceil(cn["output"][0].shape[0] / target_size)
out = [{"output": [], "middle": []} for _ in range(num_chunks)]
for k, v in cn.items():
for i, x in enumerate(v):
chunks = chunk_batch(x, (target_size, *x.shape[1:]))
for j, chunk in enumerate(chunks):
out[j][k].append(chunk)
return out
+111
View File
@@ -0,0 +1,111 @@
"""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)
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"]
context_shape = expected_inputs["encoder_hidden_states"]["shape"]
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,
)
]
+42
View File
@@ -0,0 +1,42 @@
"""Pure batch-chunking helpers for Core ML's fixed-shape UNet inputs.
Re-exported by coreml_suite.latents. Characterization tests cover the
contract (padding-zero regions, truncation in merge_chunks,
identity-passthrough when shape already matches).
"""
import torch
def chunk_batch(input_tensor, target_shape):
if input_tensor.shape == target_shape:
return [input_tensor]
batch_size = input_tensor.shape[0]
target_batch_size = target_shape[0]
num_chunks = batch_size // target_batch_size
if num_chunks == 0:
padding = torch.zeros(target_batch_size - batch_size, *target_shape[1:]).to(
input_tensor.device
)
return [torch.cat((input_tensor, padding), dim=0)]
mod = batch_size % target_batch_size
if mod != 0:
chunks = list(torch.chunk(input_tensor[:-mod], num_chunks))
padding = torch.zeros(target_batch_size - mod, *target_shape[1:]).to(
input_tensor.device
)
padded = torch.cat((input_tensor[-mod:], padding), dim=0)
chunks.append(padded)
return chunks
chunks = list(torch.chunk(input_tensor, num_chunks))
return chunks
def merge_chunks(chunks, orig_shape):
merged = torch.cat(chunks, dim=0)
if merged.shape == orig_shape:
return merged
return merged[: orig_shape[0]]
+68
View File
@@ -0,0 +1,68 @@
"""Pure out_name composition for the Core ML UNet artifact.
Extracted from CoreMLConverter.convert so the filename contract
can be tested + reused without instantiating the node. The string is the
cache key: every workflow that references a converted .mlpackage depends
on it staying byte-for-byte identical.
"""
from typing import Iterable, Tuple
ATTN_SUFFIX = {
"SPLIT_EINSUM": "se",
"SPLIT_EINSUM_V2": "se2",
"ORIGINAL": "orig",
}
# Palettization bits. "none" = no quantization (default; keeps the
# unquantized filename intact so existing workflows still resolve their
# cached .mlpackage). Numeric values append a `_q<bits>` suffix.
QUANT_NBITS_VALUES = ("none", "8", "6", "4")
def compose_out_name(
*,
ckpt_name: str,
batch_size: int,
width: int,
height: int,
controlnet_support: bool,
attention_implementation: str,
lora_names: Iterable[str] = (),
quantize_nbits: str = "none",
) -> str:
"""Build the .mlpackage stem from convert() parameters.
Locked behaviour (characterization tests):
- first '.' in ckpt_name wins (`a.b.c.safetensors` -> `a`)
- spaces collapse to underscores
- LoRA names are taken stem-only, sorted, joined with '_' and
prefixed with '_' when present (caller is expected to pass a
sorted list; we sort defensively)
- controlnet adds `_cn`
- attn suffix is `_se` | `_se2` | `_orig`
Quantization:
- quantize_nbits "none" (default) appends nothing — existing
unquantized .mlpackages keep the old filename
- "4" / "6" / "8" appends `_q<bits>` after the attn suffix
"""
if quantize_nbits not in QUANT_NBITS_VALUES:
raise ValueError(
f"quantize_nbits={quantize_nbits!r} not in {QUANT_NBITS_VALUES}"
)
stem = ckpt_name.split(".")[0]
sorted_names = sorted(lora_names)
lora_str = "_" + "_".join(name.split(".")[0] for name in sorted_names) if sorted_names else ""
cn_suffix = "_cn" if controlnet_support else ""
attn_suffix = "_" + ATTN_SUFFIX[attention_implementation]
quant_suffix = f"_q{quantize_nbits}" if quantize_nbits != "none" else ""
out_name = (
f"{stem}{lora_str}_{batch_size}x{width}x{height}"
f"{cn_suffix}{attn_suffix}{quant_suffix}"
)
return out_name.replace(" ", "_")
def lora_names_from_params(lora_params: Iterable[Tuple[str, float]]) -> list[str]:
"""Mirror the sort applied inside CoreMLConverter.convert."""
return [name for name, _ in sorted(lora_params, key=lambda pair: pair[0])]
+91
View File
@@ -0,0 +1,91 @@
"""Pure SDXL detection + time_ids/text_embeds assembly.
The framework-coupled adapter `add_sdxl_model_options` lives in models.py
and delegates the math here. Characterization tests cover base (len 6) vs
refiner (len 5) and the closure free-vars produced by
`sdxl_model_function_wrapper`.
"""
import torch
def is_sdxl(coreml_model):
return (
"time_ids" in coreml_model.expected_inputs
and "text_embeds" in coreml_model.expected_inputs
)
def is_sdxl_base(coreml_model):
return (
is_sdxl(coreml_model)
and coreml_model.expected_inputs["time_ids"]["shape"][1] == 6
)
def is_sdxl_refiner(coreml_model):
return (
is_sdxl(coreml_model)
and coreml_model.expected_inputs["time_ids"]["shape"][1] == 5
)
def build_sdxl_time_ids(pos_dict, neg_dict, *, is_base: bool, is_refiner: bool):
"""Compose the (2, N) time_ids tensor for the SDXL Core ML UNet.
- base: N=6 -> [h, w, crop_h, crop_w, target_h, target_w]
- refiner: N=5 -> [h, w, crop_h, crop_w, aesthetic_score]
- neither: N=4 -> [h, w, crop_h, crop_w] (edge case kept for parity)
"""
pos_time_ids = [
pos_dict.get("height", 768),
pos_dict.get("width", 768),
pos_dict.get("crop_h", 0),
pos_dict.get("crop_w", 0),
]
neg_time_ids = [
neg_dict.get("height", 768),
neg_dict.get("width", 768),
neg_dict.get("crop_h", 0),
neg_dict.get("crop_w", 0),
]
if is_base:
pos_time_ids += [
pos_dict.get("target_height", 768),
pos_dict.get("target_width", 768),
]
neg_time_ids += [
neg_dict.get("target_height", 768),
neg_dict.get("target_width", 768),
]
if is_refiner:
pos_time_ids += [pos_dict.get("aesthetic_score", 6)]
neg_time_ids += [neg_dict.get("aesthetic_score", 2.5)]
return torch.tensor([pos_time_ids, neg_time_ids])
def build_sdxl_text_embeds(pos_pooled, neg_pooled):
"""Concat pos then neg along the batch dim. Locked contract."""
return torch.cat((pos_pooled, neg_pooled))
def sdxl_model_function_wrapper(time_ids, text_embeds, refiner=False):
def wrapper(model_function, params):
x = params["input"]
t = params["timestep"]
c = params["c"]
context = c.get("c_crossattn")
if context is None:
return torch.zeros_like(x)
if refiner and context is not None:
# converted refiner accepts only g clip
c["c_crossattn"] = context[:, :, 768:]
return model_function(x, t, **c, time_ids=time_ids, text_embeds=text_embeds)
return wrapper
+42
View File
@@ -0,0 +1,42 @@
import time
import coremltools as ct
from coreml_suite.logger import logger
class CoreMLModel:
"""Small runtime wrapper around coremltools.models.MLModel.
This keeps the inference path independent from apple/ml-stable-diffusion's
CoreMLModel wrapper while preserving the contract used by the sampler code:
``expected_inputs`` and callable prediction.
"""
def __init__(self, model_path, compute_unit):
self.model_path = model_path
self.compute_unit = self._compute_unit(compute_unit)
logger.info(f"Loading {model_path} to {self.compute_unit.name}")
start = time.time()
self.model = ct.models.MLModel(model_path, compute_units=self.compute_unit)
logger.info(f"Loading {model_path} took {time.time() - start:.1f} seconds")
self.expected_inputs = self._expected_inputs()
def __call__(self, **kwargs):
return self.model.predict(kwargs)
@staticmethod
def _compute_unit(compute_unit):
if isinstance(compute_unit, ct.ComputeUnit):
return compute_unit
return ct.ComputeUnit[compute_unit]
def _expected_inputs(self):
return {
feature.name: {
"shape": tuple(feature.type.multiArrayType.shape),
}
for feature in self.model.get_spec().description.input
}
+3 -35
View File
@@ -1,36 +1,4 @@
import torch
"""Compatibility shim — re-exports from coreml_suite.core.latents."""
from coreml_suite.core.latents import chunk_batch, merge_chunks
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]]
__all__ = ["chunk_batch", "merge_chunks"]
+17 -55
View File
@@ -1,5 +1,4 @@
import os
import shutil
import logging
import time
import gc
@@ -9,24 +8,20 @@ import torch
from diffusers import UNet2DConditionModel, LCMScheduler
from diffusers.loaders import LoraLoaderMixin
from comfy.model_management import get_torch_device
from coreml_suite.lcm.unet import UNet2DConditionModelLCM
from coreml_suite.conversion.attention import apply_attention_implementation
from coreml_suite.conversion.shapes import conv2d_output_shape
from coreml_suite.conversion.unet import CoreMLUNetWrapper
from coreml_suite.model_version import ModelVersion
from transformers import CLIPTextModel
import coremltools as ct
from folder_paths import get_folder_paths
logging.basicConfig()
logger = logging.getLogger(__name__)
logger.setLevel(logging.DEBUG)
MODEL_VERSION = "SimianLuo/LCM_Dreamshaper_v7"
MODEL_NAME = MODEL_VERSION.split("/")[-1] + "_4k"
import python_coreml_stable_diffusion.unet as unet
unet.ATTENTION_IMPLEMENTATION_IN_EFFECT = unet.AttentionImplementations.SPLIT_EINSUM
TEXT_TOKEN_SEQUENCE_LENGTH = 77
def get_unets():
@@ -37,31 +32,27 @@ def get_unets():
low_cpu_mem_usage=False,
)
cml_unet = UNet2DConditionModelLCM.from_config(ref_unet.config).eval()
cml_unet.load_state_dict(ref_unet.state_dict(), strict=False)
cml_unet = CoreMLUNetWrapper(
apply_attention_implementation(ref_unet.eval(), "SPLIT_EINSUM"),
ModelVersion.LCM,
)
return cml_unet, ref_unet
def get_encoder_hidden_states_shape(unet_config, batch_size):
text_encoder = CLIPTextModel.from_pretrained(
MODEL_VERSION, subfolder="text_encoder"
)
text_token_sequence_length = text_encoder.config.max_position_embeddings
hidden_size = (text_encoder.config.hidden_size,)
encoder_hidden_states_shape = (
batch_size,
unet_config.cross_attention_dim or hidden_size,
1,
text_token_sequence_length,
TEXT_TOKEN_SEQUENCE_LENGTH,
unet_config.cross_attention_dim,
)
return encoder_hidden_states_shape
def get_scheduler():
from comfy.model_management import get_torch_device
scheduler = LCMScheduler.from_pretrained(MODEL_VERSION, subfolder="scheduler")
scheduler.set_timesteps(50, get_torch_device(), 50)
return scheduler
@@ -117,29 +108,14 @@ def convert_to_coreml(
def get_out_path(submodule_name, model_name):
from folder_paths import get_folder_paths
fname = f"{model_name}_{submodule_name}.mlpackage"
unet_path = get_folder_paths(submodule_name)[0]
out_path = os.path.join(unet_path, fname)
return out_path
def compile_coreml_model(source_model_path, output_dir, final_name):
"""Compiles Core ML models using the coremlcompiler utility from Xcode toolchain"""
target_path = os.path.join(output_dir, f"{final_name}.mlmodelc")
if os.path.exists(target_path):
logger.warning(f"Found existing compiled model at {target_path}! Skipping..")
return target_path
logger.info(f"Compiling {source_model_path}")
source_model_name = os.path.basename(os.path.splitext(source_model_path)[0])
os.system(f"xcrun coremlcompiler compile {source_model_path} {output_dir}")
compiled_output = os.path.join(output_dir, f"{source_model_name}.mlmodelc")
shutil.move(compiled_output, target_path)
return target_path
def get_sample_input(batch_size, encoder_hidden_states_shape, sample_shape, scheduler):
sample_unet_inputs = dict(
[
@@ -165,15 +141,13 @@ def get_unet_inputs_spec(sample_unet_inputs):
def add_cnet_support(sample_shape, reference_unet):
from python_coreml_stable_diffusion.unet import calculate_conv2d_output_shape
additional_residuals_shapes = []
batch_size = sample_shape[0]
h, w = sample_shape[2:]
# conv_in
out_h, out_w = calculate_conv2d_output_shape(
out_h, out_w = conv2d_output_shape(
h,
w,
reference_unet.conv_in,
@@ -190,9 +164,7 @@ def add_cnet_support(sample_shape, reference_unet):
]
if hasattr(down_block, "downsamplers") and down_block.downsamplers is not None:
for downsampler in down_block.downsamplers:
out_h, out_w = calculate_conv2d_output_shape(
out_h, out_w, downsampler.conv
)
out_h, out_w = conv2d_output_shape(out_h, out_w, downsampler.conv)
additional_residuals_shapes.append(
(
batch_size,
@@ -272,15 +244,6 @@ def convert(
logger.info(f"Saved unet into {out_path}")
def compile_model(out_path, out_name):
# Compile the model
target_path = compile_coreml_model(
out_path, get_folder_paths("unet")[0], f"{out_name}_unet"
)
logger.info(f"Compiled {out_path} to {target_path}")
return target_path
if __name__ == "__main__":
h = 512
w = 512
@@ -294,4 +257,3 @@ if __name__ == "__main__":
out_path = get_out_path("unet", f"{out_name}")
if not os.path.exists(out_path):
convert(out_path=out_path, sample_size=sample_size, batch_size=batch_size)
compile_model(out_path=out_path, out_name=out_name)
+4 -4
View File
@@ -1,10 +1,9 @@
import os
from coremltools import ComputeUnit
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
from coreml_suite import COREML_NODE
from coreml_suite.lcm import converter as lcm_converter
from coreml_suite.coreml_model import CoreMLModel
class COREML_CONVERT_LCM(COREML_NODE):
@@ -48,6 +47,8 @@ class COREML_CONVERT_LCM(COREML_NODE):
The converted model is also saved to "models/unet" directory and
can be loaded with the "LCMCoreMLLoaderUNet" node.
"""
from coreml_suite.lcm import converter as lcm_converter
h = height
w = width
sample_size = (h // 8, w // 8)
@@ -65,6 +66,5 @@ class COREML_CONVERT_LCM(COREML_NODE):
batch_size=batch_size,
controlnet_support=controlnet_support,
)
target_path = lcm_converter.compile_model(out_path=out_path, out_name=out_name)
return (CoreMLModel(target_path, compute_unit, "compiled"),)
return (CoreMLModel(out_path, compute_unit),)
+2 -3
View File
@@ -1,5 +1,5 @@
from overrides import overrides
from python_coreml_stable_diffusion.unet import UNet2DConditionModel, TimestepEmbedding
from diffusers import UNet2DConditionModel
from diffusers.models.embeddings import TimestepEmbedding
class UNet2DConditionModelLCM(UNet2DConditionModel):
@@ -17,7 +17,6 @@ class UNet2DConditionModelLCM(UNet2DConditionModel):
)
self.time_embedding = time_embedding
@overrides(check_signature=False)
def forward(
self,
sample,
+8
View File
@@ -0,0 +1,8 @@
from enum import Enum
class ModelVersion(Enum):
SD15 = "sd15"
SDXL = "sdxl"
SDXL_REFINER = "sdxl_refiner"
LCM = "lcm"
+40 -188
View File
@@ -1,15 +1,44 @@
import numpy as np
"""Framework-coupled glue between Core ML UNets and ComfyUI's sampler stack.
Pure math (CoreMLInputs, SDXL detection, time_ids/text_embeds assembly,
sdxl_model_function_wrapper) lives in coreml_suite.core.*.
This module is what touches comfy.*: model_base, ModelPatcher, the
diffusion_model wrapper, and the maintainer-facing add_sdxl_model_options
adapter.
"""
import torch
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.controlnet import extract_residual_kwargs, chunk_control
from coreml_suite.latents import chunk_batch, merge_chunks
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.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):
@@ -68,204 +97,27 @@ 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]
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_base = model_patcher.model.diffusion_model.is_sdxl_base
is_refiner = model_patcher.model.diffusion_model.is_sdxl_refiner
if is_refiner:
pos_time_ids += [
pos_dict.get("aesthetic_score", 6),
]
neg_time_ids += [
neg_dict.get("aesthetic_score", 2.5),
]
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"]
)
time_ids = torch.tensor([pos_time_ids, neg_time_ids])
text_embeds = torch.cat((pos_pooled, neg_pooled))
model_options = {
mp.model_options |= {
"model_function_wrapper": sdxl_model_function_wrapper(
time_ids, text_embeds, is_refiner
),
}
mp.model_options |= model_options
return mp
+33 -36
View File
@@ -1,15 +1,19 @@
import os
from coremltools import ComputeUnit
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
from python_coreml_stable_diffusion.unet import AttentionImplementations
import folder_paths
from coreml_suite import COREML_NODE
from coreml_suite import converter
from coreml_suite.config import ModelVersion
from coreml_suite.attention import ATTENTION_IMPLEMENTATIONS
from coreml_suite.coreml_model import CoreMLModel
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 coreml_suite.model_version import ModelVersion
from nodes import KSampler, LoraLoader, KSamplerAdvanced
from coreml_suite.models import (
@@ -163,7 +167,7 @@ class CoreMLLoader(COREML_NODE):
@classmethod
def coreml_filenames(cls):
extensions = (".mlmodelc", ".mlpackage")
extensions = (".mlpackage",)
all_paths = folder_paths.get_filename_list_(cls.PACKAGE_DIRNAME)[1]
coreml_paths = folder_paths.filter_files_extensions(all_paths, extensions)
@@ -174,9 +178,7 @@ class CoreMLLoader(COREML_NODE):
coreml_path = self.coreml_filenames()[coreml_name]
sources = "compiled" if coreml_name.endswith(".mlmodelc") else "packages"
return (CoreMLModel(coreml_path, compute_unit, sources),)
return (CoreMLModel(coreml_path, compute_unit),)
class CoreMLLoaderUNet(CoreMLLoader):
@@ -223,15 +225,11 @@ class CoreMLConverter(COREML_NODE):
ModelVersion.SDXL.name,
],
),
"height": ("INT", {"default": 512, "min": 256, "max": 2048, "step": 8}),
"width": ("INT", {"default": 512, "min": 256, "max": 2048, "step": 8}),
"height": ("INT", {"default": 512, "min": 8, "step": 8}),
"width": ("INT", {"default": 512, "min": 8, "step": 8}),
"batch_size": ("INT", {"default": 1, "min": 1, "max": 64}),
"attention_implementation": (
[
AttentionImplementations.SPLIT_EINSUM.name,
AttentionImplementations.SPLIT_EINSUM_V2.name,
AttentionImplementations.ORIGINAL.name,
],
list(ATTENTION_IMPLEMENTATIONS),
),
"compute_unit": (
[
@@ -244,6 +242,12 @@ 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",),
},
}
@@ -262,6 +266,7 @@ class CoreMLConverter(COREML_NODE):
attention_implementation,
compute_unit,
controlnet_support,
quantize_nbits="none",
lora_params=None,
):
"""Converts a LCM model to Core ML.
@@ -288,24 +293,17 @@ class CoreMLConverter(COREML_NODE):
h = height
w = width
sample_size = (h // 8, w // 8)
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 ""
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,
)
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}")
@@ -317,6 +315,8 @@ class CoreMLConverter(COREML_NODE):
for lora_param in lora_params:
logger.info(f" {lora_param[0]} - strength: {lora_param[1]}")
from coreml_suite import converter
unet_out_path = converter.get_out_path("unet", f"{out_name}")
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
@@ -335,12 +335,9 @@ 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"
)
return (CoreMLModel(unet_target_path, compute_unit, "compiled"),)
return (CoreMLModel(unet_out_path, compute_unit),)
@staticmethod
def lora_path(lora_name):
+56
View File
@@ -0,0 +1,56 @@
[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 = "2.0.2"
license = "MIT"
requires-python = ">=3.12,<3.13"
packages = [{ include = "coreml_suite" }]
dependencies = [
"torch>=2.7,<2.8",
"coremltools>=9,<10",
"numpy>=2,<3",
"diffusers>=0.30",
"peft>=0.13",
"omegaconf>=2.3",
"transformers>=4.44",
]
[project.urls]
Repository = "https://github.com/aszc-dev/ComfyUI-CoreMLSuite"
[tool.comfy]
PublisherId = "aszc-dev"
DisplayName = "ComfyUI-CoreMLSuite"
Icon = ""
requires-comfyui = ">=0.3.27"
[dependency-groups]
dev = [
"pillow>=12.2.0",
"psutil>=7.2.2",
"pytest>=9.0.3",
]
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.pytest.ini_options]
markers = [
"unit: framework-free unit test (Tier 0)",
"smoke: macOS-ARM smoke test on a synthetic micro-model (Tier 1)",
"m2: requires Apple Silicon + Neural Engine (Tier 2)",
]
testpaths = ["tests"]
addopts = ["--import-mode=importlib", "--confcutdir=tests"]
+6 -5
View File
@@ -1,6 +1,7 @@
git+https://github.com/apple/ml-stable-diffusion.git
coremltools>=7.1
overrides
diffusers>=0.22
peft>=0.6.2
torch>=2.7,<2.8
coremltools>=9,<10
numpy>=2,<3
diffusers>=0.30
peft>=0.13
omegaconf>=2.3
transformers>=4.44
View File
+61
View File
@@ -0,0 +1,61 @@
"""Pytest bootstrap for ComfyUI-CoreMLSuite tests.
- Adds the ComfyUI checkout to sys.path so the framework-coupled modules
that transitively import `comfy.*` resolve when pytest is invoked from
this package's root.
- Auto-applies tier markers based on the directory a test lives in, so
individual files don't have to repeat @pytest.mark.unit / .smoke.
"""
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 smoke), skip the other
# directories at collection time. Tier-0 cannot afford to import tests/smoke
# files because they pull in coremltools which Linux CI won't have.
_TIER_DIRS = {
"unit": ("/tests/unit/",),
"m2": ("/tests/m2/", "/tests/integration/"),
"smoke": ("/tests/smoke/",),
}
def pytest_ignore_collect(collection_path, config):
expr = config.option.markexpr
if expr not in _TIER_DIRS:
return None
allowed = _TIER_DIRS[expr]
rel = str(collection_path).replace("\\", "/")
if "/tests/" not in rel:
return None
# Always allow tests/ root + the tier's own dirs.
if rel.endswith("/tests"):
return None
if any(frag in rel + "/" for frag in allowed):
return None
return True
def pytest_collection_modifyitems(config, items):
for item in items:
path = str(item.fspath).replace("\\", "/")
for fragment, marker in _TIER_BY_DIR.items():
if f"/{fragment}/" in path:
item.add_marker(getattr(pytest.mark, marker))
break
@@ -1,72 +0,0 @@
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
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 448 KiB

+1
View File
@@ -0,0 +1 @@
e89344e544d4edfbd3ebe9a1c78dadb2729f53549666052b74ac7308f326f4fc
+170
View File
@@ -0,0 +1,170 @@
"""[M2-ANE] golden-image anchor.
Runs the e2e SD1.5 + CoreML workflow against a local ComfyUI server, fetches
the generated PNG, and asserts both:
- byte-identical SHA256 against the stored golden, OR
- PSNR >= GOLDEN_PSNR_MIN_DB against the stored golden PNG.
The hash is the strict gate (a refactor that doesn't touch the math
should hit it). PSNR is the soft gate that tolerates the drift a
toolchain bump injects through different MIL graphs / kernel selection
/ fp accumulation order — anything below the threshold is treated as a
regression.
The 20 dB default absorbs Apple Neural Engine run-to-run nondeterminism:
the same model and seed can drift several dB between runs as the 20
sampling steps amplify tiny per-step UNet differences (kernel selection /
fp accumulation order). Same-scene ANE outputs have been observed at
~23 dB, so 20 leaves margin while still catching gross regressions — a
broken image lands far lower. Bump it up for pure-refactor PRs that must
not change math; down for toolchain bumps.
Skips entirely on non-Apple-Silicon hosts or when the server / converted
model is missing, so the unit lane on Linux still passes.
The first run with no golden writes one and fails so it's reviewed before
being committed.
"""
import hashlib
import json
import os
import platform
import shutil
import time
import urllib.error
import urllib.request
from pathlib import Path
import numpy as np
import pytest
from PIL import Image
REPO_ROOT = Path(__file__).resolve().parents[2]
COMFY_DIR = Path(os.environ.get("COMFY_DIR", REPO_ROOT.parents[1])).resolve()
COMFY_HOST = os.environ.get("COMFY_HOST", "localhost")
COMFY_PORT = int(os.environ.get("COMFY_PORT", "8188"))
COMFY_URL = f"http://{COMFY_HOST}:{COMFY_PORT}"
CKPT_NAME = os.environ.get("CKPT_NAME", "v1-5-pruned-emaonly.safetensors")
WORKFLOW_PATH = (
REPO_ROOT / "tests" / "integration" / "workflows" / "e2e-1.5-basic-conversion.json"
)
GOLDEN_DIR = Path(__file__).parent / "goldens"
GOLDEN_HASH_PATH = GOLDEN_DIR / "sd15_seed42.sha256"
GOLDEN_PNG_PATH = GOLDEN_DIR / "sd15_seed42.png"
GOLDEN_PSNR_MIN_DB = float(os.environ.get("GOLDEN_PSNR_MIN_DB", "20"))
SEED = 42
def _server_reachable() -> bool:
try:
with urllib.request.urlopen(f"{COMFY_URL}/prompt", timeout=3) as r:
return r.status == 200
except (urllib.error.URLError, urllib.error.HTTPError, ConnectionError):
return False
@pytest.fixture(scope="module")
def comfy_server():
if platform.machine() != "arm64":
pytest.skip("requires Apple Silicon")
if not _server_reachable():
pytest.skip(f"ComfyUI server not reachable at {COMFY_URL}")
return COMFY_URL
def _http_post_json(path: str, payload: dict) -> dict:
data = json.dumps(payload).encode("utf-8")
req = urllib.request.Request(
f"{COMFY_URL}{path}", data=data,
headers={"Content-Type": "application/json"}, method="POST",
)
with urllib.request.urlopen(req, timeout=300) as r:
return json.loads(r.read().decode())
def _http_get_json(path: str, timeout: int = 300) -> dict:
"""ComfyUI runs UNet inference on its single asyncio loop, so GET /prompt
blocks while the queued prompt is executing. Use a generous timeout."""
with urllib.request.urlopen(f"{COMFY_URL}{path}", timeout=timeout) as r:
return json.loads(r.read().decode())
def _drain_queue(timeout_s: int = 600) -> None:
deadline = time.time() + timeout_s
while time.time() < deadline:
try:
q = _http_get_json("/prompt")
except (urllib.error.URLError, TimeoutError):
# Transient block while server executes; retry until our overall
# deadline expires.
continue
if q.get("exec_info", {}).get("queue_remaining", -1) == 0:
return
time.sleep(2)
raise TimeoutError(f"queue did not drain within {timeout_s}s")
def _post_workflow_and_collect_png() -> bytes:
workflow = json.loads(WORKFLOW_PATH.read_text())
for nid in ("4", "10"):
if nid in workflow:
workflow[nid]["inputs"]["ckpt_name"] = CKPT_NAME
for nid in ("3", "11"):
if nid in workflow and "seed" in workflow[nid].get("inputs", {}):
workflow[nid]["inputs"]["seed"] = SEED
# Drop the MPS reference branch — only the Core ML pipeline is needed here.
for nid in ("3", "8", "9"):
workflow.pop(nid, None)
_http_post_json("/prompt", {"prompt": workflow})
_drain_queue()
comfy_out = COMFY_DIR / "output"
matches = sorted(comfy_out.glob("E2E-1.5-CoreML_*.png"), reverse=True)
if not matches:
raise FileNotFoundError(f"no Core ML image under {comfy_out}")
return matches[0].read_bytes()
def _psnr(a: np.ndarray, b: np.ndarray) -> float:
mse = float(np.mean((a.astype(np.float64) - b.astype(np.float64)) ** 2))
if mse == 0:
return 100.0
return 20.0 * float(np.log10(255.0 / np.sqrt(mse)))
def test_sd15_seed42_image_matches_golden(comfy_server):
GOLDEN_DIR.mkdir(parents=True, exist_ok=True)
png_bytes = _post_workflow_and_collect_png()
sha = hashlib.sha256(png_bytes).hexdigest()
if not GOLDEN_HASH_PATH.exists() or not GOLDEN_PNG_PATH.exists():
GOLDEN_HASH_PATH.write_text(sha + "\n")
# Persist the PNG too for visual diffing + PSNR.
tmp_path = Path(__file__).parent / "_latest_generated.png"
tmp_path.write_bytes(png_bytes)
shutil.copy2(tmp_path, GOLDEN_PNG_PATH)
pytest.fail(
f"No golden present; wrote {GOLDEN_HASH_PATH.name} and "
f"{GOLDEN_PNG_PATH.name}. Review the image and re-run."
)
expected_hash = GOLDEN_HASH_PATH.read_text().strip()
if sha == expected_hash:
return
# Hash drift: fall back to PSNR to distinguish a refactor-safe rounding
# change from a real regression.
a = np.array(Image.open(GOLDEN_PNG_PATH).convert("RGB"))
b_path = Path(__file__).parent / "_latest_generated.png"
b_path.write_bytes(png_bytes)
b = np.array(Image.open(b_path).convert("RGB"))
if a.shape != b.shape:
pytest.fail(f"shape mismatch: golden={a.shape} actual={b.shape}")
psnr_db = _psnr(a, b)
assert psnr_db >= GOLDEN_PSNR_MIN_DB, (
f"hash drifted (got {sha[:12]}.., expected {expected_hash[:12]}..) and "
f"PSNR {psnr_db:.2f} dB < {GOLDEN_PSNR_MIN_DB} dB threshold; "
f"diff PNG at {b_path}"
)
View File
@@ -0,0 +1,41 @@
import platform
import pytest
import torch
from diffusers.models.attention_processor import Attention, AttnProcessor
from coreml_suite.conversion.attention import (
SplitEinsumAttnProcessor,
SplitEinsumV2AttnProcessor,
)
pytestmark = pytest.mark.skipif(
platform.system() != "Darwin" or platform.machine() != "arm64",
reason="Tier 1 requires macOS on Apple Silicon",
)
@pytest.mark.parametrize(
"processor",
[
SplitEinsumAttnProcessor(),
SplitEinsumV2AttnProcessor(),
],
)
def test_split_einsum_processor_matches_diffusers_attention(processor):
torch.manual_seed(0)
reference = Attention(query_dim=32, heads=4, dim_head=8, dropout=0.0)
reference.set_processor(AttnProcessor())
candidate = Attention(query_dim=32, heads=4, dim_head=8, dropout=0.0)
candidate.load_state_dict(reference.state_dict())
candidate.set_processor(processor)
hidden_states = torch.randn(2, 17, 32)
encoder_hidden_states = torch.randn(2, 11, 32)
expected = reference(hidden_states, encoder_hidden_states=encoder_hidden_states)
actual = candidate(hidden_states, encoder_hidden_states=encoder_hidden_states)
assert torch.allclose(actual, expected, atol=1e-5)
+138
View File
@@ -0,0 +1,138 @@
"""Tier 1 smoke: convert a synthetic micro-UNet through coremltools and load
it back with CoreMLSuite's runtime CoreMLModel wrapper.
Purpose: catch API breakage in coremltools *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 CoreMLSuite's 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
from types import SimpleNamespace
import numpy as np
import pytest
import torch
import torch.nn as nn
from coreml_suite.conversion.unet import CoreMLUNetWrapper
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, 4, 64) # native diffusers encoder_hidden_states (batch, tokens, hidden)
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 CoreMLSuite's runtime 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,
timestep_cond=None,
added_cond_kwargs=None,
down_block_additional_residuals=None,
mid_block_additional_residual=None,
return_dict=True,
):
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.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 = CoreMLUNetWrapper(
TinyUNet().eval(),
SimpleNamespace(name="SD15"),
)
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 coreml_suite.coreml_model import CoreMLModel
model = CoreMLModel(str(tiny_mlpackage), "CPU_ONLY")
# expected_inputs is the contract our wrappers depend on. Lock the shape
# of the dict + a sample entry.
expected = dict(model.expected_inputs)
assert set(expected.keys()) == {"sample", "timestep", "encoder_hidden_states"}
assert tuple(expected["sample"]["shape"]) == SAMPLE_SHAPE
assert tuple(expected["timestep"]["shape"]) == TIMESTEP_SHAPE
assert tuple(expected["encoder_hidden_states"]["shape"]) == ENCODER_SHAPE
# Forward pass: drive the model the way CoreMLModelWrapper does.
rng = np.random.default_rng(0)
inputs = {
"sample": rng.standard_normal(SAMPLE_SHAPE).astype(np.float16),
"timestep": rng.standard_normal(TIMESTEP_SHAPE).astype(np.float16),
"encoder_hidden_states": rng.standard_normal(ENCODER_SHAPE).astype(np.float16),
}
out = model(**inputs)
assert isinstance(out, dict), f"unexpected output type: {type(out)}"
assert OUT_NAME in out, f"missing output {OUT_NAME!r}; got {sorted(out)}"
assert out[OUT_NAME].shape == SAMPLE_SHAPE, (
f"output shape drift: got {out[OUT_NAME].shape}, expected {SAMPLE_SHAPE}"
)
@@ -0,0 +1,186 @@
"""Characterization tests for coreml_suite.controlnet.
Locks shapes + dtypes + zero-fill behavior of expand_inputs / no_control /
extract_residual_kwargs / chunk_control. These pure helpers feed the Core ML
UNet's additional_residual_N inputs; any drift here silently breaks
ControlNet-based workflows.
"""
import numpy as np
import pytest
import torch
from coreml_suite.core.controlnet import (
chunk_control,
expand_inputs,
extract_residual_kwargs,
no_control,
)
@pytest.fixture(autouse=True)
def _deterministic_seed():
torch.manual_seed(0)
np.random.seed(0)
SD15_RESIDUAL_SPEC = {
"additional_residual_0": {"shape": (2, 320, 64, 64)},
"additional_residual_1": {"shape": (2, 640, 32, 32)},
"additional_residual_2": {"shape": (2, 1280, 8, 8)},
}
NON_RESIDUAL_SPEC = {
"sample": {"shape": (2, 4, 64, 64)},
"encoder_hidden_states": {"shape": (2, 77, 768)},
}
# ---------- expand_inputs ----------------------------------------------------
def test_expand_inputs_doubles_singleton_numpy():
inputs = {"a": np.ones((1, 4), dtype=np.float32)}
out = expand_inputs(inputs)
assert out["a"].shape == (2, 4)
assert np.array_equal(out["a"], np.ones((2, 4)))
def test_expand_inputs_doubles_singleton_torch():
inputs = {"a": torch.ones(1, 4)}
out = expand_inputs(inputs)
assert out["a"].shape == (2, 4)
assert torch.equal(out["a"], torch.ones(2, 4))
def test_expand_inputs_doubles_singleton_list():
inputs = {"a": [42]}
out = expand_inputs(inputs)
assert out["a"] == [42, 42]
def test_expand_inputs_skips_already_batched():
"""batch > 1 inputs are returned unchanged (same object identity)."""
arr = np.ones((2, 4), dtype=np.float32)
tensor = torch.ones(3, 4)
lst = [1, 2]
out = expand_inputs({"a": arr, "b": tensor, "c": lst})
assert out["a"] is arr
assert out["b"] is tensor
assert out["c"] is lst
def test_expand_inputs_preserves_unknown_value_types():
# Strings/None pass through untouched — locks current permissive contract.
inputs = {"s": "hello", "none": None, "int": 7}
out = expand_inputs(inputs)
assert out == {"s": "hello", "none": None, "int": 7}
# ---------- no_control -------------------------------------------------------
def test_no_control_returns_zero_fp16_for_residuals():
out = no_control({**SD15_RESIDUAL_SPEC, **NON_RESIDUAL_SPEC})
# Only additional_residual_* keys are produced.
assert set(out.keys()) == set(SD15_RESIDUAL_SPEC.keys())
for key, spec in SD15_RESIDUAL_SPEC.items():
arr = out[key]
assert arr.shape == spec["shape"]
assert arr.dtype == np.float16
assert np.all(arr == 0)
def test_no_control_returns_empty_when_no_residuals():
out = no_control(NON_RESIDUAL_SPEC)
assert out == {}
# ---------- extract_residual_kwargs -----------------------------------------
def test_extract_residual_kwargs_empty_when_model_has_no_residual_inputs():
out = extract_residual_kwargs(NON_RESIDUAL_SPEC, control={"output": [], "middle": []})
assert out == {}
def test_extract_residual_kwargs_none_control_returns_no_control_shapes():
out = extract_residual_kwargs(SD15_RESIDUAL_SPEC, control=None)
assert set(out.keys()) == set(SD15_RESIDUAL_SPEC.keys())
for key, spec in SD15_RESIDUAL_SPEC.items():
assert out[key].shape == spec["shape"]
assert out[key].dtype == np.float16
assert np.all(out[key] == 0)
def test_extract_residual_kwargs_flattens_output_then_middle_and_casts_fp16():
"""output residuals come first (indexed 0..N-1), then middle residuals
(indexed N..M-1). Values come out of CPU as fp16 numpy arrays."""
control = {
"output": [torch.ones(2, 320, 64, 64) * 0.5, torch.ones(2, 640, 32, 32) * 2.0],
"middle": [torch.ones(2, 1280, 8, 8) * -1.0],
}
out = extract_residual_kwargs(SD15_RESIDUAL_SPEC, control)
assert set(out.keys()) == {"additional_residual_0", "additional_residual_1", "additional_residual_2"}
assert out["additional_residual_0"].shape == (2, 320, 64, 64)
assert out["additional_residual_1"].shape == (2, 640, 32, 32)
assert out["additional_residual_2"].shape == (2, 1280, 8, 8)
for arr in out.values():
assert arr.dtype == np.float16
# Locked order: index 0 == first output residual (0.5), index 2 == middle (-1.0).
assert np.allclose(out["additional_residual_0"], 0.5)
assert np.allclose(out["additional_residual_1"], 2.0)
assert np.allclose(out["additional_residual_2"], -1.0)
# ---------- chunk_control ----------------------------------------------------
def test_chunk_control_none_returns_list_of_nones_with_length_target():
"""`no_control` path: when there's no control, you get [None] * target_size
(NOT [None, None] regardless of target — this is the contract today)."""
assert chunk_control(None, 1) == [None]
assert chunk_control(None, 2) == [None, None]
assert chunk_control(None, 4) == [None, None, None, None]
@pytest.mark.parametrize(
"batch,target,expected_chunks",
[(1, 2, 1), (2, 2, 1), (3, 2, 2), (4, 2, 2), (5, 3, 2), (9, 4, 3)],
)
def test_chunk_control_shapes_after_chunking(batch, target, expected_chunks):
cn = {
"output": [
torch.randn(batch, 320, 64, 64),
torch.randn(batch, 640, 32, 32),
],
"middle": [torch.randn(batch, 1280, 8, 8)],
}
chunks = chunk_control(cn, target)
assert len(chunks) == expected_chunks
for c in chunks:
assert c["output"][0].shape == (target, 320, 64, 64)
assert c["output"][1].shape == (target, 640, 32, 32)
assert c["middle"][0].shape == (target, 1280, 8, 8)
def test_chunk_control_preserves_keys_order():
"""Output dicts contain exactly {"output", "middle"} in that order."""
cn = {
"output": [torch.zeros(2, 4, 4, 4)],
"middle": [torch.zeros(2, 4, 4, 4)],
}
chunks = chunk_control(cn, 2)
assert list(chunks[0].keys()) == ["output", "middle"]
def test_chunk_control_zero_pads_remainder():
"""A batch=3, target=2 split puts the third row alongside a zero row."""
cn = {
"output": [torch.arange(3 * 4).reshape(3, 1, 2, 2).float()],
"middle": [torch.arange(3 * 4).reshape(3, 1, 2, 2).float()],
}
chunks = chunk_control(cn, 2)
assert len(chunks) == 2
last_out = chunks[-1]["output"][0]
# First row is the original third row; second row is padding zeros.
assert torch.equal(last_out[0], cn["output"][0][2])
assert torch.equal(last_out[1], torch.zeros(1, 2, 2))
+228
View File
@@ -0,0 +1,228 @@
"""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, 77, 768)},
}
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, 77, 2048)},
"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, 77, 1280)},
"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 keeps Comfy's native (b, seq, dim) layout.
assert out["encoder_hidden_states"].shape == (1, 77, 768)
assert out["encoder_hidden_states"].dtype == np.float16
assert out["timestep"].shape == (1,)
assert out["timestep"].dtype == np.float16
def test_coreml_kwargs_sd15_with_controlnet_emits_residuals():
inputs = _sd15_inputs(batch=1, with_control=True)
out = inputs.coreml_kwargs(SD15_WITH_CN)
assert "additional_residual_0" in out
assert "additional_residual_1" in out
assert out["additional_residual_0"].shape == (1, 320, 64, 64)
assert out["additional_residual_1"].shape == (1, 640, 32, 32)
def test_coreml_kwargs_sd15_without_controlnet_zero_fills_residuals():
inputs = _sd15_inputs(batch=1, with_control=False)
out = inputs.coreml_kwargs(SD15_WITH_CN)
assert np.all(out["additional_residual_0"] == 0)
assert np.all(out["additional_residual_1"] == 0)
def test_coreml_kwargs_lcm_adds_timestep_cond():
inputs = _sd15_inputs(batch=1, with_ts_cond=True)
out = inputs.coreml_kwargs(LCM_EXPECTED)
assert "timestep_cond" in out
assert out["timestep_cond"].shape == (1, 256)
assert out["timestep_cond"].dtype == np.float16
def test_coreml_kwargs_lcm_skips_timestep_cond_when_not_provided():
"""timestep_cond is only forwarded when the input supplied one — even if
the model's expected_inputs lists it."""
inputs = _sd15_inputs(batch=1, with_ts_cond=False)
out = inputs.coreml_kwargs(LCM_EXPECTED)
assert "timestep_cond" not in out
def test_coreml_kwargs_sdxl_base_emits_time_ids_and_text_embeds():
out = _sdxl_inputs(batch=1, refiner=False).coreml_kwargs(SDXL_BASE_EXPECTED)
assert out["time_ids"].shape == (1, 6)
assert out["text_embeds"].shape == (1, 1280)
assert out["time_ids"].dtype == np.float16
assert out["text_embeds"].dtype == np.float16
def test_coreml_kwargs_sdxl_refiner_uses_len5_time_ids():
out = _sdxl_inputs(batch=1, refiner=True).coreml_kwargs(SDXL_REFINER_EXPECTED)
assert out["time_ids"].shape == (1, 5)
# ---------- chunks ----------------------------------------------------------
def test_chunks_sd15_pad_to_batch2_returns_one_chunk():
chunked = _sd15_inputs(batch=1).chunks(SD15_EXPECTED)
assert len(chunked) == 1
c = chunked[0]
assert c.x.shape == (2, 4, 64, 64)
assert c.t.shape == (2,)
# context shape: (b, seq, dim) padded along batch dim.
assert c.context.shape == (2, 77, 768)
assert c.control is None
assert c.ts_cond is None
assert c.time_ids is None
assert c.text_embeds is None
def test_chunks_sd15_with_controlnet_chunks_residuals_too():
chunked = _sd15_inputs(batch=1, with_control=True).chunks(SD15_EXPECTED)
assert len(chunked) == 1
cn = chunked[0].control
assert cn is not None
assert cn["output"][0].shape == (2, 320, 64, 64)
assert cn["output"][1].shape == (2, 640, 32, 32)
def test_chunks_lcm_carries_timestep_cond_per_chunk():
chunked = _sd15_inputs(batch=1, with_ts_cond=True).chunks(LCM_EXPECTED)
assert len(chunked) == 1
assert chunked[0].ts_cond is not None
assert chunked[0].ts_cond.shape == (2, 256)
def test_chunks_sdxl_base_propagates_time_ids_and_text_embeds():
chunked = _sdxl_inputs(batch=1, refiner=False).chunks(SDXL_BASE_EXPECTED)
assert len(chunked) == 1
c = chunked[0]
assert c.time_ids is not None and c.time_ids.shape == (2, 6)
assert c.text_embeds is not None and c.text_embeds.shape == (2, 1280)
def test_chunks_sdxl_refiner_uses_len5_time_ids():
chunked = _sdxl_inputs(batch=1, refiner=True).chunks(SDXL_REFINER_EXPECTED)
assert chunked[0].time_ids.shape == (2, 5)
def test_chunks_sdxl_synthesizes_zero_time_ids_when_caller_omits():
"""If the model expects time_ids but caller passed nothing, the suite
fabricates a zero-filled tensor. Lock that fallback."""
x = torch.randn(1, 4, 128, 128)
t = torch.full((1,), 999.0)
context = torch.randn(1, 77, 2048)
inputs = CoreMLInputs(x, t, context, control=None)
chunked = inputs.chunks(SDXL_BASE_EXPECTED)
assert chunked[0].time_ids.shape == (2, 6)
assert torch.equal(chunked[0].time_ids, torch.zeros(2, 6))
assert chunked[0].text_embeds.shape == (2, 1280)
assert torch.equal(chunked[0].text_embeds, torch.zeros(2, 1280))
def test_chunks_splits_batch_into_multiple_target2_chunks():
"""batch=5 with target_batch=2 -> 3 chunks (last padded)."""
chunked = _sd15_inputs(batch=5).chunks(SD15_EXPECTED)
assert len(chunked) == 3
for c in chunked:
assert c.x.shape == (2, 4, 64, 64)
assert c.context.shape == (2, 77, 768)
# Last chunk's second batch row is the zero-pad.
assert torch.equal(chunked[-1].x[1], torch.zeros(4, 64, 64))
def test_chunks_timestep_is_broadcast_from_first_value():
"""t is rebuilt from t[0] across all chunks: locks current behavior that
discards any per-row timestep variation."""
x = torch.randn(2, 4, 64, 64)
t = torch.tensor([42.0, 99.0]) # the second value will be lost
context = torch.randn(2, 77, 768)
inputs = CoreMLInputs(x, t, context, control=None)
chunked = inputs.chunks(SD15_EXPECTED)
assert chunked[0].t.shape == (2,)
assert torch.equal(chunked[0].t, torch.full((2,), 42.0))
+118
View File
@@ -0,0 +1,118 @@
"""Characterization tests for coreml_suite.latents.
Locks the *current* behavior of chunk_batch / merge_chunks — including the
zero-pad regions and the truncation in merge — so a refactor
cannot silently shift either contract.
"""
import pytest
import torch
from coreml_suite.core.latents import chunk_batch, merge_chunks
@pytest.fixture(autouse=True)
def _deterministic_seed():
torch.manual_seed(0)
def _const_tensor(batch, *rest):
return torch.arange(batch * 4 * 8 * 8, dtype=torch.float32).reshape(batch, 4, 8, 8)
# ---------- chunk_batch ------------------------------------------------------
def test_chunk_batch_passthrough_when_shape_matches():
x = _const_tensor(2)
out = chunk_batch(x, (2, 4, 8, 8))
assert len(out) == 1
# passthrough: the same object identity is returned (no copy).
assert out[0] is x
def test_chunk_batch_pads_single_chunk_when_input_smaller():
"""batch=1, target=2 -> one padded chunk; the second row is exact zero."""
x = _const_tensor(1)
out = chunk_batch(x, (2, 4, 8, 8))
assert len(out) == 1
assert out[0].shape == (2, 4, 8, 8)
assert torch.equal(out[0][0], x[0])
assert torch.equal(out[0][1], torch.zeros(4, 8, 8))
def test_chunk_batch_splits_exact_multiple():
"""batch=4, target=2 -> two chunks, no padding."""
x = _const_tensor(4)
out = chunk_batch(x, (2, 4, 8, 8))
assert len(out) == 2
assert out[0].shape == (2, 4, 8, 8)
assert out[1].shape == (2, 4, 8, 8)
assert torch.equal(out[0], x[:2])
assert torch.equal(out[1], x[2:])
def test_chunk_batch_pads_remainder_chunk():
"""batch=5, target=2 -> chunks=[x[0:2], x[2:4]] then [x[4], 0]."""
x = _const_tensor(5)
out = chunk_batch(x, (2, 4, 8, 8))
assert len(out) == 3
assert torch.equal(out[0], x[0:2])
assert torch.equal(out[1], x[2:4])
last = out[-1]
assert last.shape == (2, 4, 8, 8)
assert torch.equal(last[0], x[4])
# The remainder row is zero-padded; lock that exact contract.
assert torch.equal(last[1], torch.zeros(4, 8, 8))
assert last[1].sum() == 0
@pytest.mark.parametrize(
"batch_size,target,expected_chunks",
[
(1, 4, 1),
(3, 2, 2),
(5, 3, 2),
(9, 4, 3),
],
)
def test_chunk_batch_pad_region_is_zero(batch_size, target, expected_chunks):
x = _const_tensor(batch_size)
out = chunk_batch(x, (target, 4, 8, 8))
assert len(out) == expected_chunks
mod = batch_size % target
if mod == 0 and batch_size >= target:
return
last = out[-1]
pad_rows = target - (mod if (mod != 0 and batch_size >= target) else batch_size)
pad_region = last[-pad_rows:]
assert torch.equal(pad_region, torch.zeros_like(pad_region))
# ---------- merge_chunks -----------------------------------------------------
def test_merge_chunks_exact_concat():
x = _const_tensor(4)
chunks = chunk_batch(x, (2, 4, 8, 8))
merged = merge_chunks(chunks, x.shape)
assert merged.shape == x.shape
assert torch.equal(merged, x)
def test_merge_chunks_truncates_padding():
"""Round-trip with a padded last chunk drops the pad rows."""
x = _const_tensor(5)
chunks = chunk_batch(x, (2, 4, 8, 8))
merged = merge_chunks(chunks, x.shape)
assert merged.shape == x.shape
assert torch.equal(merged, x)
def test_merge_chunks_singleton_returns_equal_copy_when_shape_matches():
"""A singleton chunk list still goes through torch.cat, so we get a new
tensor equal to the input — locked here because a refactor might be tempted
to short-circuit and accidentally return the same object."""
x = _const_tensor(2)
out = merge_chunks([x], x.shape)
assert torch.equal(out, x)
assert out is not x
@@ -0,0 +1,197 @@
"""Characterization tests for the .mlpackage filename composition.
The filename composition is the pure
coreml_suite.core.naming.compose_out_name function. CoreMLConverter.convert
calls it; testing the pure function avoids monkey-patching heavy converter
internals just to capture the string.
"""
import pytest
from coreml_suite.core.naming import compose_out_name, lora_names_from_params
# ---------- attention suffixes ----------------------------------------------
@pytest.mark.parametrize(
"attn_name,suffix",
[
("SPLIT_EINSUM", "se"),
("SPLIT_EINSUM_V2", "se2"),
("ORIGINAL", "orig"),
],
)
def test_attention_suffix(attn_name, suffix):
out = compose_out_name(
ckpt_name="dreamshaper_8.safetensors",
batch_size=1, width=512, height=512,
controlnet_support=False,
attention_implementation=attn_name,
)
assert out == f"dreamshaper_8_1x512x512_{suffix}"
# ---------- batch / size ----------------------------------------------------
def test_includes_batch_and_size():
out = compose_out_name(
ckpt_name="dreamshaper_8.safetensors",
batch_size=4, width=768, height=1024,
controlnet_support=False,
attention_implementation="SPLIT_EINSUM",
)
assert out == "dreamshaper_8_4x768x1024_se"
# ---------- ControlNet ------------------------------------------------------
def test_appends_cn_suffix_when_controlnet_support_true():
out = compose_out_name(
ckpt_name="dreamshaper_8.safetensors",
batch_size=1, width=512, height=512,
controlnet_support=True,
attention_implementation="SPLIT_EINSUM",
)
assert out == "dreamshaper_8_1x512x512_cn_se"
# ---------- ckpt name massage -----------------------------------------------
def test_drops_extension_at_first_period():
out = compose_out_name(
ckpt_name="my.checkpoint.v2.safetensors",
batch_size=1, width=512, height=512,
controlnet_support=False,
attention_implementation="SPLIT_EINSUM",
)
assert out == "my_1x512x512_se"
def test_replaces_spaces_with_underscores():
out = compose_out_name(
ckpt_name="dream shaper 8.safetensors",
batch_size=1, width=512, height=512,
controlnet_support=False,
attention_implementation="SPLIT_EINSUM",
)
assert out == "dream_shaper_8_1x512x512_se"
# ---------- LoRA suffixes ---------------------------------------------------
def test_single_lora():
out = compose_out_name(
ckpt_name="dreamshaper_8.safetensors",
batch_size=1, width=512, height=512,
controlnet_support=False,
attention_implementation="SPLIT_EINSUM",
lora_names=["epi_noiseoffset.safetensors"],
)
assert out == "dreamshaper_8_epi_noiseoffset_1x512x512_se"
def test_multiple_loras_sorted():
out = compose_out_name(
ckpt_name="dreamshaper_8.safetensors",
batch_size=1, width=512, height=512,
controlnet_support=False,
attention_implementation="SPLIT_EINSUM",
lora_names=["zoom.safetensors", "alpha.safetensors", "moody.safetensors"],
)
assert out == "dreamshaper_8_alpha_moody_zoom_1x512x512_se"
def test_lora_plus_controlnet():
out = compose_out_name(
ckpt_name="dreamshaper_8.safetensors",
batch_size=1, width=512, height=512,
controlnet_support=True,
attention_implementation="SPLIT_EINSUM",
lora_names=["a.safetensors"],
)
assert out == "dreamshaper_8_a_1x512x512_cn_se"
# ---------- sdxl combinations -----------------------------------------------
def test_sdxl_1024_original_gpu():
out = compose_out_name(
ckpt_name="sd_xl_base_1.0.safetensors",
batch_size=1, width=1024, height=1024,
controlnet_support=False,
attention_implementation="ORIGINAL",
)
assert out == "sd_xl_base_1_1x1024x1024_orig"
# ---------- lora_names_from_params helper ----------------------------------
def test_lora_names_from_params_sorts_by_name():
names = lora_names_from_params([
("zebra.safetensors", 1.0),
("apple.safetensors", 0.5),
("mango.safetensors", 0.7),
])
assert names == ["apple.safetensors", "mango.safetensors", "zebra.safetensors"]
def test_lora_names_from_params_empty_list():
assert lora_names_from_params([]) == []
# ---------- quantize_nbits suffix ------------------------------------------
def test_quantize_nbits_none_appends_nothing():
"""'none' is the default and must keep the unquantized filename so
existing cached .mlpackages still resolve."""
out = compose_out_name(
ckpt_name="dreamshaper_8.safetensors",
batch_size=1, width=512, height=512,
controlnet_support=False,
attention_implementation="SPLIT_EINSUM",
quantize_nbits="none",
)
assert out == "dreamshaper_8_1x512x512_se"
@pytest.mark.parametrize("nbits,suffix", [("4", "_q4"), ("6", "_q6"), ("8", "_q8")])
def test_quantize_nbits_appends_q_suffix(nbits, suffix):
out = compose_out_name(
ckpt_name="dreamshaper_8.safetensors",
batch_size=1, width=512, height=512,
controlnet_support=False,
attention_implementation="SPLIT_EINSUM",
quantize_nbits=nbits,
)
assert out == f"dreamshaper_8_1x512x512_se{suffix}"
def test_quantize_nbits_with_controlnet_and_lora():
out = compose_out_name(
ckpt_name="dreamshaper_8.safetensors",
batch_size=1, width=512, height=512,
controlnet_support=True,
attention_implementation="SPLIT_EINSUM",
lora_names=["a.safetensors"],
quantize_nbits="6",
)
assert out == "dreamshaper_8_a_1x512x512_cn_se_q6"
def test_quantize_nbits_invalid_raises():
import pytest as _pytest
with _pytest.raises(ValueError, match="quantize_nbits"):
compose_out_name(
ckpt_name="x.safetensors",
batch_size=1, width=512, height=512,
controlnet_support=False,
attention_implementation="SPLIT_EINSUM",
quantize_nbits="16", # not in {none, 8, 6, 4}
)
@@ -0,0 +1,127 @@
"""Characterization tests for the SDXL options math.
The SDXL time_ids / text_embeds math lives in
coreml_suite.core.sdxl as pure builders. The framework adapter
add_sdxl_model_options lives in models.py; 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)
+24 -27
View File
@@ -1,37 +1,34 @@
import pytest
"""Smoke tests for the pure batch-chunking helpers in coreml_suite.core.
Uses torch.device('cpu') instead of comfy.model_management.get_torch_device
so Tier 0 runs without ComfyUI.
"""
import pytest
import torch
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
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")
@pytest.fixture
def expected_inputs():
expected = {
return {
"sample": {"shape": (2, 4, 64, 64)},
"timestep": {"shape": (2,)},
"timestep_cond": {"shape": (2, 256)},
"encoder_hidden_states": {"shape": (2, 768, 1, 77)},
"encoder_hidden_states": {"shape": (2, 77, 768)},
"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(get_torch_device())
latent_image = torch.randn(batch_size, 4, 64, 64).to(CPU)
target_shape = (4, 4, 64, 64)
chunked = chunk_batch(latent_image, target_shape)
@@ -45,7 +42,7 @@ def test_batch_chunking(batch_size):
@pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9])
def test_merge_chunks(batch_size):
input_tensor = torch.randn(batch_size, 4, 64, 64).to(get_torch_device())
input_tensor = torch.randn(batch_size, 4, 64, 64).to(CPU)
target_shape = (4, 4, 64, 64)
chunked = chunk_batch(input_tensor, target_shape)
@@ -57,16 +54,16 @@ def test_merge_chunks(batch_size):
@pytest.fixture
def inputs():
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())
x = torch.randn(1, 4, 64, 64).to(CPU)
t = torch.randn([1]).to(CPU)
c_crossattn = torch.randn(1, 77, 768).to(CPU)
control = {
"output": [
torch.randn(1, 320, 64, 64).to(get_torch_device()),
torch.randn(1, 640, 32, 32).to(get_torch_device()),
torch.randn(1, 320, 64, 64).to(CPU),
torch.randn(1, 640, 32, 32).to(CPU),
],
}
timestep_cond = torch.randn(1, 256).to(get_torch_device())
timestep_cond = torch.randn(1, 256).to(CPU)
return CoreMLInputs(x, t, c_crossattn, control, timestep_cond=timestep_cond)
@@ -86,11 +83,11 @@ def inputs():
def test_chunking_controlnet(b, target_size, num_chunks):
cn = {
"output": [
torch.randn(b, 320, 64, 64).to(get_torch_device()),
torch.randn(b, 640, 32, 32).to(get_torch_device()),
torch.randn(b, 320, 64, 64).to(CPU),
torch.randn(b, 640, 32, 32).to(CPU),
],
"middle": [
torch.randn(b, 1280, 8, 8).to(get_torch_device()),
torch.randn(b, 1280, 8, 8).to(CPU),
],
}
+183
View File
@@ -0,0 +1,183 @@
from types import SimpleNamespace
import torch
from coreml_suite.conversion.attention import (
SplitEinsumAttnProcessor,
SplitEinsumV2AttnProcessor,
apply_attention_implementation,
split_einsum,
split_einsum_v2,
)
from coreml_suite.conversion.shapes import conv2d_output_shape
from coreml_suite.conversion.unet import CoreMLUNetWrapper
class RecordingUNet(torch.nn.Module):
def __init__(self):
super().__init__()
self.call = None
def forward(
self,
sample,
timestep,
encoder_hidden_states,
timestep_cond=None,
added_cond_kwargs=None,
down_block_additional_residuals=None,
mid_block_additional_residual=None,
return_dict=True,
**kwargs,
):
self.call = {
"sample": sample,
"timestep": timestep,
"encoder_hidden_states": encoder_hidden_states,
"timestep_cond": timestep_cond,
"added_cond_kwargs": added_cond_kwargs,
"down_block_additional_residuals": down_block_additional_residuals,
"mid_block_additional_residual": mid_block_additional_residual,
"return_dict": return_dict,
}
return (sample + 1,)
def test_conv2d_output_shape_matches_torch_conv2d_contract():
conv = torch.nn.Conv2d(
4,
8,
kernel_size=(3, 5),
stride=(2, 3),
padding=(1, 2),
dilation=(1, 2),
)
assert conv2d_output_shape(17, 19, conv) == (9, 5)
def test_unet_wrapper_passes_context_through_for_sd15():
unet = RecordingUNet()
wrapper = CoreMLUNetWrapper(unet, SimpleNamespace(name="SD15"))
sample = torch.randn(2, 4, 8, 8)
timestep = torch.randn(2)
context = torch.randn(2, 77, 768)
out = wrapper(sample, timestep, context)
assert torch.equal(out, sample + 1)
assert unet.call["encoder_hidden_states"] is context
assert unet.call["return_dict"] is False
def test_unet_wrapper_routes_lcm_sdxl_and_controlnet_inputs():
unet = RecordingUNet()
wrapper = CoreMLUNetWrapper(unet, SimpleNamespace(name="LCM"))
sample = torch.randn(1, 4, 8, 8)
timestep = torch.randn(1)
context = torch.randn(1, 77, 768)
timestep_cond = torch.randn(1, 256)
down_residual = torch.randn(1, 320, 8, 8)
mid_residual = torch.randn(1, 1280, 1, 1)
wrapper(sample, timestep, context, timestep_cond, down_residual, mid_residual)
assert unet.call["timestep_cond"] is timestep_cond
assert len(unet.call["down_block_additional_residuals"]) == 1
assert unet.call["down_block_additional_residuals"][0] is down_residual
assert unet.call["mid_block_additional_residual"] is mid_residual
def test_unet_wrapper_routes_sdxl_added_conditioning():
unet = RecordingUNet()
wrapper = CoreMLUNetWrapper(unet, SimpleNamespace(name="SDXL"))
sample = torch.randn(1, 4, 8, 8)
timestep = torch.randn(1)
context = torch.randn(1, 77, 2048)
time_ids = torch.randn(1, 6)
text_embeds = torch.randn(1, 1280)
wrapper(sample, timestep, context, time_ids, text_embeds)
assert unet.call["added_cond_kwargs"]["time_ids"] is time_ids
assert unet.call["added_cond_kwargs"]["text_embeds"] is text_embeds
def test_split_einsum_matches_original_attention_math():
torch.manual_seed(0)
batch = 2
heads = 3
dim_head = 4
sequence = 16
channels = heads * dim_head
q = torch.randn(batch, channels, 1, sequence)
k = torch.randn(batch, channels, 1, sequence)
v = torch.randn(batch, channels, 1, sequence)
expected = _original_attention(q, k, v, None, heads, dim_head)
# split-einsum reorders the float32 reductions vs the reference, so equality
# only holds up to rounding; the drift exceeds allclose's default atol on
# some BLAS backends (e.g. Linux x86 CI).
assert torch.allclose(split_einsum(q, k, v, None, heads, dim_head), expected, atol=1e-6)
assert torch.allclose(split_einsum_v2(q, k, v, None, heads, dim_head), expected, atol=1e-6)
def test_split_einsum_v2_chunked_path_matches_original_attention_math():
torch.manual_seed(0)
batch = 1
heads = 2
dim_head = 2
sequence = 512
channels = heads * dim_head
q = torch.randn(batch, channels, 1, sequence)
k = torch.randn(batch, channels, 1, sequence)
v = torch.randn(batch, channels, 1, sequence)
expected = _original_attention(q, k, v, None, heads, dim_head)
assert torch.allclose(
split_einsum_v2(q, k, v, None, heads, dim_head),
expected,
atol=1e-6,
)
def test_apply_attention_implementation_sets_split_processors():
unet = RecordingProcessorUNet()
assert apply_attention_implementation(unet, "ORIGINAL") is unet
assert unet.processor is None
apply_attention_implementation(unet, "SPLIT_EINSUM")
assert isinstance(unet.processor, SplitEinsumAttnProcessor)
apply_attention_implementation(unet, "SPLIT_EINSUM_V2")
assert isinstance(unet.processor, SplitEinsumV2AttnProcessor)
class RecordingProcessorUNet:
def __init__(self):
self.processor = None
def set_attn_processor(self, processor):
self.processor = processor
def _original_attention(q, k, v, mask, heads, dim_head):
batch = q.size(0)
mh_q = q.view(batch, heads, dim_head, -1)
mh_k = k.view(batch, heads, dim_head, -1)
mh_v = v.view(batch, heads, dim_head, -1)
weights = torch.einsum("bhcq,bhck->bhqk", mh_q, mh_k)
weights = weights * (dim_head**-0.5)
if mask is not None:
weights = weights + mask
weights = weights.softmax(dim=3)
attn = torch.einsum("bhqk,bhck->bhcq", weights, mh_v)
return attn.contiguous().view(batch, heads * dim_head, 1, -1)
+43
View File
@@ -0,0 +1,43 @@
"""Gate: prove the Tier-0 lane is framework-free.
In a pure `pytest -m unit` run, none of the banned runtime modules
(comfy, coremltools, python_coreml_stable_diffusion, folder_paths,
nodes, comfy_extras, diffusers, diffusionkit) may be in sys.modules
after collection. If they are, a tests/unit/ file is transitively
pulling them in and the Tier-0 promise — "runs on Linux with no Mac
stack" — is broken.
When other tiers are also collected, framework modules may be imported
deliberately (e.g. smoke pulls in coremltools), so the check is skipped
unless the run is purely `-m unit` — 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."
)
Generated
+1350
View File
File diff suppressed because it is too large Load Diff