Fix control scheduling regressions

This commit is contained in:
Jedrzej Kosinski
2026-07-17 21:19:49 -07:00
parent d25dbc8fc3
commit a4f01ba9cc
4 changed files with 114 additions and 9 deletions
+3 -3
View File
@@ -216,7 +216,7 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase):
del self.cond_hint
self.cond_hint = None
if full_cond_hint_original.size(0) < self.full_latent_length:
actual_cond_hint_orig = extend_to_batch_size(tensor=full_cond_hint_original, batch_size=full_cond_hint_original.size(0))
actual_cond_hint_orig = extend_to_batch_size(tensor=full_cond_hint_original, batch_size=self.full_latent_length)
self.cond_hint_original = actual_cond_hint_orig[self.sub_idxs]
# mask hints
self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number)
@@ -810,7 +810,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:
@@ -950,7 +950,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:
+3 -3
View File
@@ -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
+4 -3
View File
@@ -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
+104
View File
@@ -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()