Fix chunk_inputs

This commit is contained in:
aszc-dev
2023-11-01 22:16:47 +01:00
parent dfdc1bf520
commit db0aea3d9c
2 changed files with 62 additions and 12 deletions
+12 -12
View File
@@ -39,7 +39,7 @@ class CoreMLModelWrapper(BaseModel):
control=None,
transformer_options={},
):
chunked_in = self._chunk_inputs(x, t, c_crossattn, control)
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
@@ -83,19 +83,19 @@ class CoreMLModelWrapper(BaseModel):
np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"]
return torch.from_numpy(np_out).to(x.device)
@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"
]
def chunk_inputs(self, x, t, c_crossattn, control):
sample_shape = self.expected_inputs["sample"]["shape"]
timestep_shape = self.expected_inputs["timestep"]["shape"]
hidden_shape = self.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, x.shape[0])
chunked_control = chunk_control(control, sample_shape[0])
return chunked_x, ts, chunked_context, chunked_control
@property
def expected_inputs(self):
return self.diffusion_model.expected_inputs
+50
View File
@@ -1,3 +1,5 @@
from unittest import mock
import pytest
import torch
@@ -5,6 +7,25 @@ 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
from coreml_suite.models import CoreMLModelWrapper, get_model_config
@pytest.fixture
def coreml_model():
model = mock.Mock()
model.expected_inputs = {
"sample": {"shape": (2, 4, 64, 64)},
"timestep": {"shape": (2,)},
"encoder_hidden_states": {"shape": (2, 768, 1, 77)},
"additional_residual_0": {"shape": (2, 320, 64, 64)},
"additional_residual_1": {"shape": (2, 640, 32, 32)},
}
return model
@pytest.fixture
def model_config():
return get_model_config()
@pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9])
@@ -36,6 +57,7 @@ def test_merge_chunks(batch_size):
@pytest.mark.parametrize(
"b, target_size, num_chunks",
[
(1, 2, 1),
(1, 1, 1),
(2, 2, 1),
(3, 2, 2),
@@ -71,3 +93,31 @@ def test_chunking_no_control():
chunked = chunk_control(cn, target_size)
assert chunked == [None, None]
def test_chunking_inputs(coreml_model, model_config):
model = CoreMLModelWrapper(model_config, coreml_model)
x = torch.randn(1, 4, 64, 64).to(get_torch_device())
t = torch.randn([1]).to(get_torch_device())
c_crossattn = torch.randn(1, 77, 768).to(get_torch_device())
control = {
"output": [
torch.randn(1, 320, 64, 64).to(get_torch_device()),
torch.randn(1, 640, 32, 32).to(get_torch_device()),
],
}
chunked_x, ts, chunked_context, chunked_control = model.chunk_inputs(
x, t, c_crossattn, control
)
assert len(chunked_x) == 1
assert len(ts) == 1
assert len(chunked_context) == 1
assert len(chunked_control) == 1
assert chunked_x[0].shape == (2, 4, 64, 64)
assert ts[0].shape == (2,)
assert chunked_context[0].shape == (2, 77, 768)
assert chunked_control[0]["output"][0].shape == (2, 320, 64, 64)
assert chunked_control[0]["output"][1].shape == (2, 640, 32, 32)