From 41797203d798a0a9ae02e7da8bc6ed041d695d76 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Tue, 31 Oct 2023 02:49:35 +0100 Subject: [PATCH] Improve chunking and padding --- coreml_suite/latents.py | 31 ++++++++++++++++++------------- coreml_suite/models.py | 26 +++++++++++++++----------- coreml_suite/nodes.py | 6 ------ tests/test_chunks.py | 14 ++++++++++---- 4 files changed, 43 insertions(+), 34 deletions(-) diff --git a/coreml_suite/latents.py b/coreml_suite/latents.py index ce42ac5..db823dc 100644 --- a/coreml_suite/latents.py +++ b/coreml_suite/latents.py @@ -1,29 +1,34 @@ import torch +from comfy.model_management import get_torch_device -def chunk_batch(latent_image, target_shape): - if latent_image.shape == target_shape: - return [latent_image] - batch_size = latent_image.shape[0] +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:]) - return [torch.cat((latent_image, padding), dim=0)] + padding = torch.zeros(target_batch_size - batch_size, *target_shape[1:]).to( + get_torch_device() + ) + return [torch.cat((input_tensor, padding), dim=0)] mod = batch_size % target_batch_size if mod != 0: - chunks = list(torch.chunk(latent_image[:-mod], num_chunks)) - padding = torch.zeros(target_batch_size - mod, *target_shape[1:]) - padded = torch.cat((latent_image[-mod:], padding), dim=0) + chunks = list(torch.chunk(input_tensor[:-mod], num_chunks)) + padding = torch.zeros(target_batch_size - mod, *target_shape[1:]).to( + get_torch_device() + ) + padded = torch.cat((input_tensor[-mod:], padding), dim=0) chunks.append(padded) - return [chunk for chunk in chunks] + return chunks - chunks = list(torch.chunk(latent_image, num_chunks)) - - return [chunk for chunk in chunks] + chunks = list(torch.chunk(input_tensor, num_chunks)) + return chunks def merge_chunks(chunks, orig_shape): diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 07e7fb6..4b7623a 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -1,7 +1,6 @@ import numpy as np import torch - from comfy import supported_models_base from comfy.latent_formats import SD15 from comfy.model_base import BaseModel @@ -40,20 +39,12 @@ class CoreMLModelWrapper(BaseModel): control=None, transformer_options={}, ): - sample_shape = self.diffusion_model.expected_inputs["sample"]["shape"] - - 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_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 ) - for x, t, c_crossattn, control in zip( - chunked_x, ts, chunked_context, chunked_control - ) + for x, t, c_crossattn, control in zip(*chunked_in) ] merged_out = merge_chunks(chunked_out, x.shape) @@ -95,3 +86,16 @@ class CoreMLModelWrapper(BaseModel): @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" + ] + 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, len(chunked_x)) + return chunked_x, ts, chunked_context, chunked_control diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index f4b598e..f73c79e 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -49,12 +49,6 @@ class CoreMLSampler(KSampler): expected = coreml_model.expected_inputs["sample"]["shape"] latent_image = {"samples": torch.zeros(expected[0] // 2, *expected[1:])} - batch_size = latent_image["samples"].shape[0] - if batch_size != coreml_model.expected_inputs["sample"]["shape"][0]: - logger.warning( - "Batch size is different from expected input size. Chunking and/or padding will be applied." - ) - return super().sample( model, seed, diff --git a/tests/test_chunks.py b/tests/test_chunks.py index 5a50dea..09898b6 100644 --- a/tests/test_chunks.py +++ b/tests/test_chunks.py @@ -2,13 +2,14 @@ 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 @pytest.mark.parametrize("batch_size", [2, 4, 5, 9]) def test_batch_chunking(batch_size): - latent_image = torch.randn(batch_size, 4, 64, 64) + 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) @@ -22,7 +23,7 @@ def test_batch_chunking(batch_size): @pytest.mark.parametrize("batch_size", [2, 4, 5, 9]) def test_merge_chunks(batch_size): - input_tensor = torch.randn(batch_size, 4, 64, 64) + 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) @@ -34,8 +35,13 @@ def test_merge_chunks(batch_size): 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)], + "output": [ + torch.randn(4, 4, 64, 64).to(get_torch_device()), + torch.randn(4, 4, 128, 128).to(get_torch_device()), + ], + "middle": [ + torch.randn(4, 4, 256, 256).to(get_torch_device()), + ], } target_batch_size = 2 num_chunks = cn["output"][0].shape[0] // target_batch_size