Fix controlnet residuals chunking
This commit is contained in:
+12
-10
@@ -1,8 +1,10 @@
|
||||
from itertools import chain
|
||||
from math import ceil
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from coreml_suite.latents import chunk_batch
|
||||
from coreml_suite.logger import logger
|
||||
|
||||
|
||||
@@ -62,18 +64,18 @@ def no_control(model):
|
||||
return residual_kwargs
|
||||
|
||||
|
||||
def chunk_control(cn, num_chunks):
|
||||
def chunk_control(cn, target_size):
|
||||
if cn is None:
|
||||
return [None] * num_chunks
|
||||
return [None] * target_size
|
||||
|
||||
chunked = []
|
||||
num_chunks = ceil(cn["output"][0].shape[0] / target_size)
|
||||
|
||||
chunk_size = len(cn["output"][0]) // num_chunks
|
||||
out = [{"output": [], "middle": []} for _ in range(num_chunks)]
|
||||
|
||||
for i in range(0, len(cn["output"][0]), chunk_size):
|
||||
chunk = {}
|
||||
chunk["output"] = [x[i : i + chunk_size] for x in cn["output"]]
|
||||
chunk["middle"] = [x[i : i + chunk_size] for x in cn["middle"]]
|
||||
chunked.append(chunk)
|
||||
for k, v in cn.items():
|
||||
for i, x in enumerate(v):
|
||||
chunks = chunk_batch(x, (target_size, *x.shape[1:]))
|
||||
for j, chunk in enumerate(chunks):
|
||||
out[j][k].append(chunk)
|
||||
|
||||
return chunked
|
||||
return out
|
||||
|
||||
@@ -97,5 +97,5 @@ class CoreMLModelWrapper(BaseModel):
|
||||
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))
|
||||
chunked_control = chunk_control(control, sample_shape[0])
|
||||
return chunked_x, ts, chunked_context, chunked_control
|
||||
|
||||
+24
-14
@@ -7,7 +7,7 @@ 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])
|
||||
@pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9])
|
||||
def test_batch_chunking(batch_size):
|
||||
latent_image = torch.randn(batch_size, 4, 64, 64).to(get_torch_device())
|
||||
target_shape = (4, 4, 64, 64)
|
||||
@@ -21,7 +21,7 @@ def test_batch_chunking(batch_size):
|
||||
assert chunked[-1][batch_size % target_shape[0] :].sum() == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize("batch_size", [2, 4, 5, 9])
|
||||
@pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9])
|
||||
def test_merge_chunks(batch_size):
|
||||
input_tensor = torch.randn(batch_size, 4, 64, 64).to(get_torch_device())
|
||||
target_shape = (4, 4, 64, 64)
|
||||
@@ -33,31 +33,41 @@ def test_merge_chunks(batch_size):
|
||||
assert torch.equal(input_tensor, merged)
|
||||
|
||||
|
||||
def test_chunking_controlnet():
|
||||
@pytest.mark.parametrize(
|
||||
"b, target_size, num_chunks",
|
||||
[
|
||||
(1, 1, 1),
|
||||
(2, 2, 1),
|
||||
(3, 2, 2),
|
||||
(4, 2, 2),
|
||||
(5, 3, 2),
|
||||
(9, 4, 3),
|
||||
],
|
||||
)
|
||||
def test_chunking_controlnet(b, target_size, num_chunks):
|
||||
cn = {
|
||||
"output": [
|
||||
torch.randn(4, 4, 64, 64).to(get_torch_device()),
|
||||
torch.randn(4, 4, 128, 128).to(get_torch_device()),
|
||||
torch.randn(b, 320, 64, 64).to(get_torch_device()),
|
||||
torch.randn(b, 640, 32, 32).to(get_torch_device()),
|
||||
],
|
||||
"middle": [
|
||||
torch.randn(4, 4, 256, 256).to(get_torch_device()),
|
||||
torch.randn(b, 1280, 8, 8).to(get_torch_device()),
|
||||
],
|
||||
}
|
||||
target_batch_size = 2
|
||||
num_chunks = cn["output"][0].shape[0] // target_batch_size
|
||||
|
||||
chunked = chunk_control(cn, num_chunks)
|
||||
chunked = chunk_control(cn, target_size)
|
||||
|
||||
assert len(chunked) == num_chunks
|
||||
for chunk in chunked:
|
||||
assert chunk["output"][0].shape == (target_batch_size, 4, 64, 64)
|
||||
assert chunk["output"][1].shape == (target_batch_size, 4, 128, 128)
|
||||
assert chunk["middle"][0].shape == (target_batch_size, 4, 256, 256)
|
||||
assert chunk["output"][0].shape == (target_size, 320, 64, 64)
|
||||
assert chunk["output"][1].shape == (target_size, 640, 32, 32)
|
||||
assert chunk["middle"][0].shape == (target_size, 1280, 8, 8)
|
||||
|
||||
|
||||
def test_chunking_no_control():
|
||||
cn = None
|
||||
num_chunks = 2
|
||||
target_size = 2
|
||||
|
||||
chunked = chunk_control(cn, num_chunks)
|
||||
chunked = chunk_control(cn, target_size)
|
||||
|
||||
assert chunked == [None, None]
|
||||
|
||||
Reference in New Issue
Block a user