diff --git a/coreml_suite/controlnet.py b/coreml_suite/controlnet.py index 8c37bfc..9cce3d6 100644 --- a/coreml_suite/controlnet.py +++ b/coreml_suite/controlnet.py @@ -60,3 +60,20 @@ def no_control(model): for i in range(len(residuals_names)) } return residual_kwargs + + +def chunk_control(cn, num_chunks): + if cn is None: + return [None] * num_chunks + + chunked = [] + + chunk_size = len(cn["output"][0]) // num_chunks + + for i in range(0, len(cn["output"][0]), chunk_size): + chunk = {} + chunk["output"] = [x[i : i + chunk_size] for x in cn["output"]] + chunk["middle"] = [x[i : i + chunk_size] for x in cn["middle"]] + chunked.append(chunk) + + return chunked diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 37a6e88..07e7fb6 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -6,7 +6,7 @@ from comfy import supported_models_base from comfy.latent_formats import SD15 from comfy.model_base import BaseModel -from coreml_suite.controlnet import expand_inputs, extract_residual_kwargs +from coreml_suite.controlnet import extract_residual_kwargs, chunk_control from coreml_suite.latents import chunk_batch, merge_chunks @@ -45,12 +45,15 @@ class CoreMLModelWrapper(BaseModel): chunked_x = chunk_batch(x, sample_shape) ts = t.chunk(len(chunked_x), dim=0) chunked_context = c_crossattn.chunk(len(chunked_x), dim=0) + chunked_control = chunk_control(control, len(chunked_x)) chunked_out = [ self._apply_model( x, t, c_concat, c_crossattn, c_adm, control, transformer_options ) - for x, t, c_crossattn in zip(chunked_x, ts, chunked_context) + for x, t, c_crossattn, control in zip( + chunked_x, ts, chunked_context, chunked_control + ) ] merged_out = merge_chunks(chunked_out, x.shape) diff --git a/tests/test_chunks.py b/tests/test_chunks.py index f3ec378..5a50dea 100644 --- a/tests/test_chunks.py +++ b/tests/test_chunks.py @@ -3,6 +3,7 @@ import pytest import torch from coreml_suite.latents import chunk_batch, merge_chunks +from coreml_suite.controlnet import chunk_control @pytest.mark.parametrize("batch_size", [2, 4, 5, 9]) @@ -29,3 +30,28 @@ def test_merge_chunks(batch_size): assert merged.shape == input_tensor.shape assert torch.equal(input_tensor, merged) + + +def test_chunking_controlnet(): + cn = { + "output": [torch.randn(4, 4, 64, 64), torch.randn(4, 4, 128, 128)], + "middle": [torch.randn(4, 4, 256, 256)], + } + target_batch_size = 2 + num_chunks = cn["output"][0].shape[0] // target_batch_size + + chunked = chunk_control(cn, num_chunks) + + for chunk in chunked: + assert chunk["output"][0].shape == (target_batch_size, 4, 64, 64) + assert chunk["output"][1].shape == (target_batch_size, 4, 128, 128) + assert chunk["middle"][0].shape == (target_batch_size, 4, 256, 256) + + +def test_chunking_no_control(): + cn = None + num_chunks = 2 + + chunked = chunk_control(cn, num_chunks) + + assert chunked == [None, None]