Chunking works for ControlNet
This commit is contained in:
@@ -60,3 +60,20 @@ def no_control(model):
|
|||||||
for i in range(len(residuals_names))
|
for i in range(len(residuals_names))
|
||||||
}
|
}
|
||||||
return residual_kwargs
|
return residual_kwargs
|
||||||
|
|
||||||
|
|
||||||
|
def chunk_control(cn, num_chunks):
|
||||||
|
if cn is None:
|
||||||
|
return [None] * num_chunks
|
||||||
|
|
||||||
|
chunked = []
|
||||||
|
|
||||||
|
chunk_size = len(cn["output"][0]) // 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)
|
||||||
|
|
||||||
|
return chunked
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ from comfy import supported_models_base
|
|||||||
from comfy.latent_formats import SD15
|
from comfy.latent_formats import SD15
|
||||||
from comfy.model_base import BaseModel
|
from comfy.model_base import BaseModel
|
||||||
|
|
||||||
from coreml_suite.controlnet import expand_inputs, extract_residual_kwargs
|
from coreml_suite.controlnet import extract_residual_kwargs, chunk_control
|
||||||
from coreml_suite.latents import chunk_batch, merge_chunks
|
from coreml_suite.latents import chunk_batch, merge_chunks
|
||||||
|
|
||||||
|
|
||||||
@@ -45,12 +45,15 @@ class CoreMLModelWrapper(BaseModel):
|
|||||||
chunked_x = chunk_batch(x, sample_shape)
|
chunked_x = chunk_batch(x, sample_shape)
|
||||||
ts = t.chunk(len(chunked_x), dim=0)
|
ts = t.chunk(len(chunked_x), dim=0)
|
||||||
chunked_context = c_crossattn.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_out = [
|
chunked_out = [
|
||||||
self._apply_model(
|
self._apply_model(
|
||||||
x, t, c_concat, c_crossattn, c_adm, control, transformer_options
|
x, t, c_concat, c_crossattn, c_adm, control, transformer_options
|
||||||
)
|
)
|
||||||
for x, t, c_crossattn in zip(chunked_x, ts, chunked_context)
|
for x, t, c_crossattn, control in zip(
|
||||||
|
chunked_x, ts, chunked_context, chunked_control
|
||||||
|
)
|
||||||
]
|
]
|
||||||
|
|
||||||
merged_out = merge_chunks(chunked_out, x.shape)
|
merged_out = merge_chunks(chunked_out, x.shape)
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ import pytest
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from coreml_suite.latents import chunk_batch, merge_chunks
|
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", [2, 4, 5, 9])
|
||||||
@@ -29,3 +30,28 @@ def test_merge_chunks(batch_size):
|
|||||||
|
|
||||||
assert merged.shape == input_tensor.shape
|
assert merged.shape == input_tensor.shape
|
||||||
assert torch.equal(input_tensor, merged)
|
assert torch.equal(input_tensor, merged)
|
||||||
|
|
||||||
|
|
||||||
|
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)],
|
||||||
|
}
|
||||||
|
target_batch_size = 2
|
||||||
|
num_chunks = cn["output"][0].shape[0] // target_batch_size
|
||||||
|
|
||||||
|
chunked = chunk_control(cn, 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)
|
||||||
|
|
||||||
|
|
||||||
|
def test_chunking_no_control():
|
||||||
|
cn = None
|
||||||
|
num_chunks = 2
|
||||||
|
|
||||||
|
chunked = chunk_control(cn, num_chunks)
|
||||||
|
|
||||||
|
assert chunked == [None, None]
|
||||||
|
|||||||
Reference in New Issue
Block a user