diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 64eba09..1a9ef60 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -39,7 +39,7 @@ class CoreMLModelWrapper(BaseModel): control=None, transformer_options={}, ): - chunked_in = self._chunk_inputs(x, t, c_crossattn, control) + chunked_in = self.chunk_inputs(x, t, c_crossattn, control) chunked_out = [ self._apply_model( x, t, c_concat, c_crossattn, c_adm, control, transformer_options @@ -83,19 +83,19 @@ class CoreMLModelWrapper(BaseModel): np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"] return torch.from_numpy(np_out).to(x.device) - @property - def expected_inputs(self): - return self.diffusion_model.expected_inputs - - def _chunk_inputs(self, x, t, c_crossattn, control): - sample_shape = self.diffusion_model.expected_inputs["sample"]["shape"] - timestep_shape = self.diffusion_model.expected_inputs["timestep"]["shape"] - hidden_shape = self.diffusion_model.expected_inputs["encoder_hidden_states"][ - "shape" - ] + def chunk_inputs(self, x, t, c_crossattn, control): + sample_shape = self.expected_inputs["sample"]["shape"] + timestep_shape = self.expected_inputs["timestep"]["shape"] + hidden_shape = self.expected_inputs["encoder_hidden_states"]["shape"] context_shape = (hidden_shape[0], hidden_shape[3], hidden_shape[1]) + chunked_x = chunk_batch(x, sample_shape) ts = list(torch.full((len(chunked_x), timestep_shape[0]), t[0])) chunked_context = chunk_batch(c_crossattn, context_shape) - chunked_control = chunk_control(control, x.shape[0]) + chunked_control = chunk_control(control, sample_shape[0]) + return chunked_x, ts, chunked_context, chunked_control + + @property + def expected_inputs(self): + return self.diffusion_model.expected_inputs diff --git a/tests/test_chunks.py b/tests/test_chunks.py index a1ef1b2..abf8231 100644 --- a/tests/test_chunks.py +++ b/tests/test_chunks.py @@ -1,3 +1,5 @@ +from unittest import mock + import pytest import torch @@ -5,6 +7,25 @@ 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 CoreMLModelWrapper, get_model_config + + +@pytest.fixture +def coreml_model(): + model = mock.Mock() + model.expected_inputs = { + "sample": {"shape": (2, 4, 64, 64)}, + "timestep": {"shape": (2,)}, + "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 model + + +@pytest.fixture +def model_config(): + return get_model_config() @pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9]) @@ -36,6 +57,7 @@ def test_merge_chunks(batch_size): @pytest.mark.parametrize( "b, target_size, num_chunks", [ + (1, 2, 1), (1, 1, 1), (2, 2, 1), (3, 2, 2), @@ -71,3 +93,31 @@ def test_chunking_no_control(): chunked = chunk_control(cn, target_size) assert chunked == [None, None] + + +def test_chunking_inputs(coreml_model, model_config): + model = CoreMLModelWrapper(model_config, coreml_model) + 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()), + ], + } + + chunked_x, ts, chunked_context, chunked_control = model.chunk_inputs( + x, t, c_crossattn, control + ) + + assert len(chunked_x) == 1 + assert len(ts) == 1 + assert len(chunked_context) == 1 + assert len(chunked_control) == 1 + + assert chunked_x[0].shape == (2, 4, 64, 64) + assert ts[0].shape == (2,) + assert chunked_context[0].shape == (2, 77, 768) + assert chunked_control[0]["output"][0].shape == (2, 320, 64, 64) + assert chunked_control[0]["output"][1].shape == (2, 640, 32, 32)