Files
aszc-dev-ComfyUI-CoreMLSuite/tests/unit/test_chunks.py
T
aszc-dev ef2a18cff3 chore(phase1): pin baseline toolchain and add bench harness scaffold
Phase 1 of the modernization plan: freeze the currently-working environment
so later refactors have a measured reference point.

- Pin python-coreml-stable-diffusion to commit e5d960c4 (the one already
  installed in the maintainer's apple_env), plus torch==2.0.1, coremltools==8.2
  and numpy<1.25 to match the only env that loads ComfyUI successfully
  (Comfy's checkpoint-safe-loading branch in utils.py is gated on torch>=2.4,
  so newer torch + numpy 1.23 breaks at import).
- Mirror the same pins in requirements.txt and commit uv.lock for
  reproducible installs.
- Add requires-comfyui pinning ComfyUI to ab541335 (the validated commit).
- Fix tests/unit/test_chunks.py fixture: get_model_config() now takes a
  ModelVersion argument; pass ModelVersion.SD15 (the previously-broken test
  was the only Phase 1 production-code change required).
- Add the Phase 1 baseline harness: bench/run.py (direct Core ML UNet
  latency, deterministic), bench/scripts/convert_sd15.py (one-command
  conversion bypassing the node graph), bench/scripts/smoke_image.py (POSTs
  the existing e2e workflow to a local ComfyUI server and saves the Core ML
  image), bench/env/capture.sh (env snapshot), bench/prompts.json (fixed
  prompt set).
- Ignore apple_env/, comfy_env/, and bench/scripts/*.log.

Tests: 20/20 unit pass (test_chunks + test_controlnet).
2026-05-22 15:06:24 +02:00

126 lines
3.5 KiB
Python

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 ModelVersion, get_model_config
@pytest.fixture
def expected_inputs():
expected = {
"sample": {"shape": (2, 4, 64, 64)},
"timestep": {"shape": (2,)},
"timestep_cond": {"shape": (2, 256)},
"encoder_hidden_states": {"shape": (2, 768, 1, 77)},
"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(ModelVersion.SD15)
@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())
target_shape = (4, 4, 64, 64)
chunked = chunk_batch(latent_image, target_shape)
for chunk in chunked:
assert chunk.shape == target_shape
if batch_size % target_shape[0] != 0:
assert chunked[-1][batch_size % target_shape[0] :].sum() == 0
@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())
target_shape = (4, 4, 64, 64)
chunked = chunk_batch(input_tensor, target_shape)
merged = merge_chunks(chunked, input_tensor.shape)
assert merged.shape == input_tensor.shape
assert torch.equal(input_tensor, merged)
@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())
control = {
"output": [
torch.randn(1, 320, 64, 64).to(get_torch_device()),
torch.randn(1, 640, 32, 32).to(get_torch_device()),
],
}
timestep_cond = torch.randn(1, 256).to(get_torch_device())
return CoreMLInputs(x, t, c_crossattn, control, timestep_cond=timestep_cond)
@pytest.mark.parametrize(
"b, target_size, num_chunks",
[
(1, 2, 1),
(1, 1, 1),
(2, 2, 1),
(3, 2, 2),
(4, 2, 2),
(5, 3, 2),
(9, 4, 3),
],
)
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()),
],
"middle": [
torch.randn(b, 1280, 8, 8).to(get_torch_device()),
],
}
chunked = chunk_control(cn, target_size)
assert len(chunked) == num_chunks
for chunk in chunked:
assert chunk["output"][0].shape == (target_size, 320, 64, 64)
assert chunk["output"][1].shape == (target_size, 640, 32, 32)
assert chunk["middle"][0].shape == (target_size, 1280, 8, 8)
def test_chunking_no_control():
cn = None
target_size = 2
chunked = chunk_control(cn, target_size)
assert chunked == [None, None]
def test_chunking_inputs(expected_inputs, inputs):
chunked = inputs.chunks(expected_inputs)
assert len(chunked) == 1
assert chunked[0].x.shape == (2, 4, 64, 64)
assert chunked[0].t.shape == (2,)
assert chunked[0].context.shape == (2, 77, 768)
assert chunked[0].control["output"][0].shape == (2, 320, 64, 64)
assert chunked[0].control["output"][1].shape == (2, 640, 32, 32)
assert chunked[0].ts_cond.shape == (2, 256)