Add CoreMLInputs to handle inputs
This commit is contained in:
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user