Add CoreMLInputs to handle inputs

This commit is contained in:
aszc-dev
2023-11-08 22:19:41 +01:00
parent 27f1a19131
commit c26099b334
4 changed files with 89 additions and 86 deletions
+9 -8
View File
@@ -21,11 +21,11 @@ def expand_inputs(inputs):
return expanded
def extract_residual_kwargs(model, control):
if "additional_residual_0" not in model.expected_inputs.keys():
def extract_residual_kwargs(expected_inputs, control):
if "additional_residual_0" not in expected_inputs.keys():
return {}
if control is None:
return no_control(model)
return no_control(expected_inputs)
residual_kwargs = {
"additional_residual_{}".format(i): r.cpu().numpy().astype(np.float16)
@@ -34,12 +34,13 @@ def extract_residual_kwargs(model, control):
return residual_kwargs
def no_control(model):
expected = model.expected_inputs
def no_control(expected_inputs):
shapes_dict = {
k: v["shape"] for k, v in expected_inputs.items() if k.startswith("additional")
}
residual_kwargs = {
k: torch.zeros(*expected[k]["shape"]).cpu().numpy().astype(dtype=np.float16)
for k in model.expected_inputs.keys()
if k.startswith("additional_residual")
k: torch.zeros(*shape).cpu().numpy().astype(dtype=np.float16)
for k, shape in shapes_dict.items()
}
return residual_kwargs
+64 -51
View File
@@ -29,65 +29,20 @@ class CoreMLModelWrapper:
self.dtype = torch.float16
def __call__(self, x, t, context, control, transformer_options, **kwargs):
chunked_in = self.chunk_inputs(
x, t, context, control, kwargs.get("timestep_cond")
)
input_list = [
self.get_np_input_kwargs(*chunked) for chunked in zip(*chunked_in)
]
inputs = CoreMLInputs(x, t, context, control, **kwargs)
input_list = inputs.chunks(self.expected_inputs)
chunked_out = [
self.get_torch_outputs(self.coreml_model(**input_kwargs), x.device)
self.get_torch_outputs(
self.coreml_model(**input_kwargs.coreml_kwargs(self.expected_inputs)),
x.device,
)
for input_kwargs in input_list
]
merged_out = merge_chunks(chunked_out, x.shape)
return merged_out
def get_np_input_kwargs(self, x, t, context, control, ts_cond=None):
sample = x.cpu().numpy().astype(np.float16)
context = context.cpu().numpy().astype(np.float16)
context = context.transpose(0, 2, 1)[:, :, None, :]
t = t.cpu().numpy().astype(np.float16)
model_input_kwargs = {
"sample": sample,
"encoder_hidden_states": context,
"timestep": t,
}
residual_kwargs = extract_residual_kwargs(self.coreml_model, control)
model_input_kwargs |= residual_kwargs
if ts_cond is not None:
model_input_kwargs["timestep_cond"] = (
ts_cond.cpu().numpy().astype(np.float16)
)
return model_input_kwargs
def chunk_inputs(self, x, t, context, control, ts_cond=None):
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(context, context_shape)
chunked_control = [None] * len(chunked_x)
if control is not None:
chunked_control = chunk_control(control, sample_shape[0])
chunked_ts_cond = [None] * len(chunked_x)
if ts_cond is not None:
ts_cond_shape = self.expected_inputs["timestep_cond"]["shape"]
chunked_ts_cond = chunk_batch(ts_cond, ts_cond_shape)
return chunked_x, ts, chunked_context, chunked_control, chunked_ts_cond
@staticmethod
def get_torch_outputs(model_output, device):
return torch.from_numpy(model_output["noise_pred"]).to(device)
@@ -101,3 +56,61 @@ class CoreMLModelWrapperLCM(CoreMLModelWrapper):
def __init__(self, coreml_model):
super().__init__(coreml_model)
self.config = None
class CoreMLInputs:
def __init__(self, x, t, context, control, **kwargs):
self.x = x
self.t = t
self.context = context
self.control = control
self.ts_cond = kwargs.get("timestep_cond")
def coreml_kwargs(self, expected_inputs):
sample = self.x.cpu().numpy().astype(np.float16)
context = self.context.cpu().numpy().astype(np.float16)
context = context.transpose(0, 2, 1)[:, :, None, :]
t = self.t.cpu().numpy().astype(np.float16)
model_input_kwargs = {
"sample": sample,
"encoder_hidden_states": context,
"timestep": t,
}
residual_kwargs = extract_residual_kwargs(expected_inputs, self.control)
model_input_kwargs |= residual_kwargs
if self.ts_cond is not None:
model_input_kwargs["timestep_cond"] = (
self.ts_cond.cpu().numpy().astype(np.float16)
)
return model_input_kwargs
def chunks(self, expected_inputs):
sample_shape = expected_inputs["sample"]["shape"]
timestep_shape = expected_inputs["timestep"]["shape"]
hidden_shape = expected_inputs["encoder_hidden_states"]["shape"]
context_shape = (hidden_shape[0], hidden_shape[3], hidden_shape[1])
chunked_x = chunk_batch(self.x, sample_shape)
ts = list(torch.full((len(chunked_x), timestep_shape[0]), self.t[0]))
chunked_context = chunk_batch(self.context, context_shape)
chunked_control = [None] * len(chunked_x)
if self.control is not None:
chunked_control = chunk_control(self.control, sample_shape[0])
chunked_ts_cond = [None] * len(chunked_x)
if self.ts_cond is not None:
ts_cond_shape = expected_inputs["timestep_cond"]["shape"]
chunked_ts_cond = chunk_batch(self.ts_cond, ts_cond_shape)
return [
CoreMLInputs(x, t, context, control, timestep_cond=ts_cond)
for x, t, context, control, ts_cond in zip(
chunked_x, ts, chunked_context, chunked_control, chunked_ts_cond
)
]
+14 -22
View File
@@ -11,13 +11,13 @@ from coreml_suite.models import (
CoreMLModelWrapper,
get_model_config,
CoreMLModelWrapperLCM,
CoreMLInputs,
)
@pytest.fixture
def coreml_model():
model = mock.Mock()
model.expected_inputs = {
def expected_inputs():
expected = {
"sample": {"shape": (2, 4, 64, 64)},
"timestep": {"shape": (2,)},
"timestep_cond": {"shape": (2, 256)},
@@ -25,7 +25,7 @@ def coreml_model():
"additional_residual_0": {"shape": (2, 320, 64, 64)},
"additional_residual_1": {"shape": (2, 640, 32, 32)},
}
return model
return expected
@pytest.fixture
@@ -72,7 +72,7 @@ def inputs():
}
timestep_cond = torch.randn(1, 256).to(get_torch_device())
return x, t, c_crossattn, control, timestep_cond
return CoreMLInputs(x, t, c_crossattn, control, timestep_cond=timestep_cond)
@pytest.mark.parametrize(
@@ -116,22 +116,14 @@ def test_chunking_no_control():
assert chunked == [None, None]
def test_chunking_inputs(coreml_model, model_config, inputs):
model = CoreMLModelWrapper(model_config, coreml_model)
def test_chunking_inputs(expected_inputs, inputs):
chunked = inputs.chunks(expected_inputs)
chunked_x, ts, chunked_context, chunked_cn, chunked_ts_cond = model.chunk_inputs(
*inputs
)
assert len(chunked) == 1
assert len(chunked_x) == 1
assert len(ts) == 1
assert len(chunked_context) == 1
assert len(chunked_cn) == 1
assert len(chunked_ts_cond) == 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_cn[0]["output"][0].shape == (2, 320, 64, 64)
assert chunked_cn[0]["output"][1].shape == (2, 640, 32, 32)
assert chunked_ts_cond[0].shape == (2, 256)
assert chunked[0].x.shape == (2, 4, 64, 64)
assert chunked[0].t.shape == (2,)
assert chunked[0].context.shape == (2, 77, 768)
assert chunked[0].control["output"][0].shape == (2, 320, 64, 64)
assert chunked[0].control["output"][1].shape == (2, 640, 32, 32)
assert chunked[0].ts_cond.shape == (2, 256)
+2 -5
View File
@@ -1,17 +1,14 @@
from unittest import mock
from coreml_suite.controlnet import no_control
def test_no_control():
model = mock.Mock()
model.expected_inputs = {
expected_inputs = {
"additional_residual_0": {"shape": (2, 2, 2)},
"additional_residual_1": {"shape": (2, 4, 4)},
"additional_residual_2": {"shape": (2, 8, 8)},
}
residual_kwargs = no_control(model)
residual_kwargs = no_control(expected_inputs)
assert len(residual_kwargs) == 3
assert residual_kwargs["additional_residual_0"].shape == (2, 2, 2)