Fix chunk_inputs
This commit is contained in:
+12
-12
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user