Simplify no_control

This commit is contained in:
aszc-dev
2023-10-31 21:57:32 +01:00
parent dd438f66cc
commit 901ea6da16
2 changed files with 23 additions and 23 deletions
+4 -23
View File
@@ -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
+19
View File
@@ -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)