From 8a3e9332e1a82439c3c18b0de842f9a2a3817a61 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Mon, 30 Oct 2023 16:17:18 +0100 Subject: [PATCH] Chunk and pad batches --- coreml_suite/latents.py | 43 +++++++++++++++++++++++++++-------------- coreml_suite/models.py | 39 +++++++++++++++++++++++++++++++++---- coreml_suite/nodes.py | 41 +++++++++++++++++++++++---------------- tests/test_chunks.py | 31 +++++++++++++++++++++++++++++ tests/test_latents.py | 19 ------------------ 5 files changed, 118 insertions(+), 55 deletions(-) create mode 100644 tests/test_chunks.py delete mode 100644 tests/test_latents.py diff --git a/coreml_suite/latents.py b/coreml_suite/latents.py index 469074a..ce42ac5 100644 --- a/coreml_suite/latents.py +++ b/coreml_suite/latents.py @@ -1,20 +1,33 @@ import torch -from torchvision.transforms.functional import resize - -from coreml_suite.logger import logger -def reshape_latent_image(latent_image, target_shape): - if latent_image is None: - logger.warning("No latent image provided, using zeros.") - return {"samples": torch.zeros(target_shape)} +def chunk_batch(latent_image, target_shape): + if latent_image.shape == target_shape: + return [latent_image] - if latent_image["samples"].shape == target_shape: - return latent_image + batch_size = latent_image.shape[0] + target_batch_size = target_shape[0] - logger.warning( - "Latent image shape does not match model input shape," - " resizing to match models expected input shape." - ) - resized = resize(latent_image["samples"], target_shape[-2:]) - return {"samples": resized} + 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)] + + 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.append(padded) + return [chunk for chunk in chunks] + + chunks = list(torch.chunk(latent_image, num_chunks)) + + return [chunk for chunk in 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]] diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 8540f8d..37a6e88 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -7,6 +7,7 @@ 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.latents import chunk_batch, merge_chunks def get_model_config(): @@ -38,6 +39,36 @@ class CoreMLModelWrapper(BaseModel): c_adm=None, 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_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) + ] + + merged_out = merge_chunks(chunked_out, x.shape) + return merged_out + + def get_dtype(self): + # Hardcoding torch-compatible dtype (used for memory allocation) + return torch.float16 + + def _apply_model( + self, + x, + t, + c_concat=None, + c_crossattn=None, + c_adm=None, + control=None, + transformer_options={}, ): sample = x.cpu().numpy().astype(np.float16) @@ -53,11 +84,11 @@ class CoreMLModelWrapper(BaseModel): } residual_kwargs = extract_residual_kwargs(self.diffusion_model, control) model_input_kwargs |= residual_kwargs - model_input_kwargs = expand_inputs(model_input_kwargs) + # model_input_kwargs = expand_inputs(model_input_kwargs) np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"] return torch.from_numpy(np_out).to(x.device) - def get_dtype(self): - # Hardcoding torch-compatible dtype (used for memory allocation) - return torch.float16 + @property + def expected_inputs(self): + return self.diffusion_model.expected_inputs diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index d1e58bb..f4b598e 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -1,5 +1,6 @@ import os +import torch from coremltools import ComputeUnit from python_coreml_stable_diffusion.coreml_model import CoreMLModel @@ -7,7 +8,6 @@ import folder_paths from comfy.model_management import get_torch_device from comfy.model_patcher import ModelPatcher from coreml_suite.logger import logger -from coreml_suite.latents import reshape_latent_image from nodes import KSampler from coreml_suite.models import CoreMLModelWrapper, get_model_config @@ -28,26 +28,33 @@ class CoreMLSampler(KSampler): CATEGORY = "Core ML Suite" def sample( - self, - coreml_model, - seed, - steps, - cfg, - sampler_name, - scheduler, - positive, - negative, - latent_image=None, - denoise=1.0, + self, + coreml_model, + seed, + steps, + cfg, + sampler_name, + scheduler, + positive, + negative, + latent_image=None, + denoise=1.0, ): - sample_shape = coreml_model.expected_inputs["sample"]["shape"] - latent_image = reshape_latent_image(latent_image, sample_shape) - latent_image["samples"] = latent_image["samples"][0:1] - model_config = get_model_config() - wrapped_model = CoreMLModelAdapter(model_config, coreml_model) + wrapped_model = CoreMLModelWrapper(model_config, coreml_model) model = ModelPatcher(wrapped_model, get_torch_device(), None) + if latent_image is None: + logger.warning("No latent image provided, using empty tensor.") + 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 new file mode 100644 index 0000000..f3ec378 --- /dev/null +++ b/tests/test_chunks.py @@ -0,0 +1,31 @@ +import pytest + +import torch + +from coreml_suite.latents import chunk_batch, merge_chunks + + +@pytest.mark.parametrize("batch_size", [2, 4, 5, 9]) +def test_batch_chunking(batch_size): + latent_image = torch.randn(batch_size, 4, 64, 64) + 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", [2, 4, 5, 9]) +def test_merge_chunks(batch_size): + input_tensor = torch.randn(batch_size, 4, 64, 64) + 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) diff --git a/tests/test_latents.py b/tests/test_latents.py deleted file mode 100644 index dd47294..0000000 --- a/tests/test_latents.py +++ /dev/null @@ -1,19 +0,0 @@ -import pytest - -import torch - -from coreml_suite.latents import reshape_latent_image - - -def test_fix_latents_no_latent_image(): - reshaped = reshape_latent_image(None, (2, 4, 64, 64)) - assert reshaped["samples"].shape == (2, 4, 64, 64) - - -@pytest.mark.parametrize( - "latent_shape", [(2, 4, 64, 64), (2, 4, 128, 128), (2, 4, 32, 32), (2, 4, 128, 64)] -) -def test_reshape_latents(latent_shape): - latent_image = {"samples": torch.zeros(latent_shape)} - reshaped = reshape_latent_image(latent_image, (2, 4, 64, 64)) - assert reshaped["samples"].shape == (2, 4, 64, 64)