From 183d0b2707b3c4fecfc316e807e2c6c0ee07a648 Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Tue, 31 Oct 2023 21:57:32 +0100 Subject: [PATCH] Simplify no_control --- coreml_suite/controlnet.py | 27 ++++----------------------- tests/test_controlnet.py | 19 +++++++++++++++++++ 2 files changed, 23 insertions(+), 23 deletions(-) create mode 100644 tests/test_controlnet.py diff --git a/coreml_suite/controlnet.py b/coreml_suite/controlnet.py index 269acc1..48e0cb0 100644 --- a/coreml_suite/controlnet.py +++ b/coreml_suite/controlnet.py @@ -36,30 +36,11 @@ def extract_residual_kwargs(model, control): def no_control(model): - # Dirty hack to get the expected input shape when doing partial ControlNet - # 0.18215 is the latent scale factor (IDK, it kinda works) - # TODO: Find a better way to do this or tweak the values - - logger.warning( - "No ControlNet input, despite the model supports it. " - "Using random noise as ControlNet residuals. " - "For better results, please use a ControlNet or a model " - "that does not support ControlNet." - ) - residuals_names = [ - name - for name in model.expected_inputs.keys() - if name.startswith("additional_residual") - ] + expected = model.expected_inputs residual_kwargs = { - "additional_residual_{}".format(i): 0.18215 - * torch.randn( - *model.expected_inputs["additional_residual_{}".format(i)]["shape"] - ) - .cpu() - .numpy() - .astype(dtype=np.float16) - for i in range(len(residuals_names)) + k: torch.zeros(*expected[k]["shape"]).cpu().numpy().astype(dtype=np.float16) + for k in model.expected_inputs.keys() + if k.startswith("additional_residual") } return residual_kwargs diff --git a/tests/test_controlnet.py b/tests/test_controlnet.py new file mode 100644 index 0000000..5d3d6c7 --- /dev/null +++ b/tests/test_controlnet.py @@ -0,0 +1,19 @@ +from unittest import mock + +from coreml_suite.controlnet import no_control + + +def test_no_control(): + model = mock.Mock() + model.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) + + assert len(residual_kwargs) == 3 + assert residual_kwargs["additional_residual_0"].shape == (2, 2, 2) + assert residual_kwargs["additional_residual_1"].shape == (2, 4, 4) + assert residual_kwargs["additional_residual_2"].shape == (2, 8, 8)