diff --git a/adv_control/control.py b/adv_control/control.py index 97e729b..7c57a3f 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -813,7 +813,7 @@ def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, tim if controlnet_config is None: unet_dtype = comfy.model_management.unet_dtype() - controlnet_config = comfy.model_detection.model_config_from_unet(controlnet_data, prefix, unet_dtype, True).unet_config + controlnet_config = comfy.model_detection.model_config_from_unet(controlnet_data, prefix, use_base_if_no_match=True).unet_config load_device = comfy.model_management.get_torch_device() manual_cast_dtype = comfy.model_management.unet_manual_cast(unet_dtype, load_device) if manual_cast_dtype is not None: @@ -953,7 +953,7 @@ def load_svdcontrolnet(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, if controlnet_config is None: unet_dtype = comfy.model_management.unet_dtype() - controlnet_config = comfy.model_detection.model_config_from_unet(controlnet_data, prefix, unet_dtype, True).unet_config + controlnet_config = comfy.model_detection.model_config_from_unet(controlnet_data, prefix, use_base_if_no_match=True).unet_config load_device = comfy.model_management.get_torch_device() manual_cast_dtype = comfy.model_management.unet_manual_cast(unet_dtype, load_device) if manual_cast_dtype is not None: diff --git a/adv_control/control_lllite.py b/adv_control/control_lllite.py index efc2779..e885c42 100644 --- a/adv_control/control_lllite.py +++ b/adv_control/control_lllite.py @@ -233,7 +233,7 @@ class LLLiteModule(torch.nn.Module): mask = prepare_mask_batch(control.mask_cond_hint, (1, 1, h, w)).to(cx.dtype) mask = mask.view(mask.shape[0], 1, h * w).permute(0, 2, 1) if control.tk_mask_cond_hint is not None: - mask_tk = prepare_mask_batch(control.mask_cond_hint, (1, 1, h, w)).to(cx.dtype) + mask_tk = prepare_mask_batch(control.tk_mask_cond_hint, (1, 1, h, w)).to(cx.dtype) mask_tk = mask_tk.view(mask_tk.shape[0], 1, h * w).permute(0, 2, 1) # x in uncond/cond doubles batch size @@ -250,7 +250,7 @@ class LLLiteModule(torch.nn.Module): if mask is None: mask = 1.0 - elif mask_tk is not None: + if mask_tk is not None: mask = mask * mask_tk #logger.info(f"cs: {cx.shape}, x: {x.shape}, is_conv2d: {self.is_conv2d}") @@ -260,7 +260,7 @@ class LLLiteModule(torch.nn.Module): if control.latent_keyframes is not None: cx = cx * control.calc_latent_keyframe_mults(x=cx, batched_number=control.batched_number) if control.weights is not None and control.weights.has_uncond_multiplier: - cond_or_uncond = control.batched_number.cond_or_uncond + cond_or_uncond = control.cond_or_uncond actual_length = cx.size(0) // control.batched_number for idx, cond_type in enumerate(cond_or_uncond): # if uncond, set to weight's uncond_multiplier diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index cba8340..2e632dd 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -773,6 +773,9 @@ def refcn_diffusion_model_wrapper_factory(reference_injections: ReferenceInjecti # if nothing related to reference controlnets, do nothing special if len(ref_controlnets) == 0 and len(context_controlnets) == 0: return executor(x, *args, **kwargs) + adain_controlnets = [] + context_adain_controlnets = [] + orig_forward_timestep_embed = None try: # assign cond and uncond idxs batched_number = len(transformer_options["cond_or_uncond"]) @@ -784,14 +787,12 @@ def refcn_diffusion_model_wrapper_factory(reference_injections: ReferenceInjecti transformer_options[REF_COND_IDXS] = [i for i, z in enumerate(indiv_conds) if z == 0] # check which controlnets do which thing attn_controlnets = [] - adain_controlnets = [] for control in ref_controlnets: if ReferenceType.is_attn(control.ref_opts.reference_type): attn_controlnets.append(control) if ReferenceType.is_adain(control.ref_opts.reference_type): adain_controlnets.append(control) context_attn_controlnets = [] - context_adain_controlnets = [] # for ease of access, store current contextref_cond_idx value if len(context_controlnets) == 0: transformer_options[CONTEXTREF_TEMP_COND_IDX] = -1 @@ -877,7 +878,7 @@ def refcn_diffusion_model_wrapper_factory(reference_injections: ReferenceInjecti finally: # make sure ref banks are cleared no matter what happens - otherwise, RIP VRAM reference_injections.clean_ref_module_mem() - if len(adain_controlnets) > 0 or len(context_adain_controlnets) > 0: + if orig_forward_timestep_embed is not None: openaimodel.forward_timestep_embed = orig_forward_timestep_embed return refcn_diffusion_model_wrapper diff --git a/tests/test_regressions.py b/tests/test_regressions.py new file mode 100644 index 0000000..37ce1bf --- /dev/null +++ b/tests/test_regressions.py @@ -0,0 +1,104 @@ +import os +import sys +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +comfyui_path = os.environ.get("COMFYUI_PATH") +if comfyui_path: + sys.path.insert(0, comfyui_path) + +import torch + +from comfy.controlnet import T2IAdapter + +from adv_control.control import T2IAdapterAdvanced +from adv_control.control_lllite import LLLiteModule +from adv_control.control_reference import REF_CONTROL_LIST_ALL, RefConst, refcn_diffusion_model_wrapper_factory + + +class LLLiteRegressionTests(unittest.TestCase): + def create_module(self): + torch.manual_seed(1) + return LLLiteModule("test", False, 2, 1, 2, 2) + + def create_control(self, effect_mask=None, timestep_mask=None, uncond_multiplier=1.0): + return SimpleNamespace( + sub_idxs=None, + cond_hint=torch.ones((1, 3, 8, 8)), + latent_dims_div2=None, + latent_dims_div4=None, + mask_cond_hint=effect_mask, + tk_mask_cond_hint=timestep_mask, + latent_keyframes=None, + weights=SimpleNamespace( + has_uncond_multiplier=uncond_multiplier != 1.0, + uncond_multiplier=uncond_multiplier, + ), + batched_number=2, + cond_or_uncond=[0, 1], + strength=1.0, + _current_timestep_keyframe=SimpleNamespace(strength=1.0), + ) + + def test_unconditional_multiplier_uses_sampling_condition_order(self): + control = self.create_control(uncond_multiplier=0.25) + output = self.create_module()(torch.ones((2, 1, 2)), control) + + torch.testing.assert_close(output[1], output[0] * 0.25) + + def test_timestep_mask_applies_without_effect_mask(self): + control = self.create_control(timestep_mask=torch.zeros((1, 8, 8))) + output = self.create_module()(torch.ones((2, 1, 2)), control) + + torch.testing.assert_close(output, torch.zeros_like(output)) + + def test_effect_and_timestep_masks_are_combined(self): + control = self.create_control( + effect_mask=torch.ones((1, 8, 8)), + timestep_mask=torch.zeros((1, 8, 8)), + ) + output = self.create_module()(torch.ones((2, 1, 2)), control) + + torch.testing.assert_close(output, torch.zeros_like(output)) + + +class T2IAdapterRegressionTests(unittest.TestCase): + def test_sliding_context_extends_hint_to_full_latent_length(self): + adapter = object.__new__(T2IAdapterAdvanced) + adapter.sub_idxs = [2, 3] + adapter.full_latent_length = 4 + adapter.cond_hint_original = torch.tensor([[[[7.0]]]]) + adapter.cond_hint = None + adapter.prepare_mask_cond_hint = lambda **kwargs: None + + with patch.object(T2IAdapter, "get_control", lambda self, *args, **kwargs: self.cond_hint_original.clone()): + output = adapter.get_control_advanced(torch.empty((2, 4, 1, 1)), None, None, 1, {}) + + self.assertEqual(output.flatten().tolist(), [7.0, 7.0]) + self.assertEqual(adapter.cond_hint_original.flatten().tolist(), [7.0]) + + +class ReferenceRegressionTests(unittest.TestCase): + def test_cleanup_does_not_hide_original_exception(self): + class ReferenceInjections: + cleaned = False + + def clean_ref_module_mem(self): + self.cleaned = True + + reference_injections = ReferenceInjections() + wrapper = refcn_diffusion_model_wrapper_factory(reference_injections) + transformer_options = { + REF_CONTROL_LIST_ALL: [SimpleNamespace(should_run=lambda: True)], + RefConst.REFCN_PRESENT_IN_CONDS: True, + } + + with self.assertRaisesRegex(KeyError, "cond_or_uncond"): + wrapper(lambda *args, **kwargs: None, torch.zeros(1), None, None, None, None, transformer_options) + + self.assertTrue(reference_injections.cleaned) + + +if __name__ == "__main__": + unittest.main()