Simplify no_control
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user