Improve chunking and padding
This commit is contained in:
+18
-13
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user