Improve chunking and padding

This commit is contained in:
aszc-dev
2023-10-31 02:49:35 +01:00
parent 4d83603c98
commit 41797203d7
4 changed files with 43 additions and 34 deletions
+18 -13
View File
@@ -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):
+15 -11
View File
@@ -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
-6
View File
@@ -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,
+10 -4
View File
@@ -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