Fix T2I Adapter sliding context hints

This commit is contained in:
Jedrzej Kosinski
2026-07-17 23:41:20 -07:00
parent 2f9dd25d93
commit 42cdfe6c88
2 changed files with 55 additions and 2 deletions
+1 -1
View File
@@ -219,7 +219,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)
+54 -1
View File
@@ -10,8 +10,11 @@ if comfyui_path:
import torch
from adv_control.control import ControlNetAdvanced
from comfy.controlnet import T2IAdapter
from adv_control.control import ControlNetAdvanced, T2IAdapterAdvanced
from adv_control.nodes_main import AdvancedControlNetApply, AdvancedControlNetInpaintingApply
from adv_control.utils import ControlWeights
class StopControlModel(Exception):
@@ -103,6 +106,56 @@ class ModernControlPreprocessingTests(unittest.TestCase):
self.assertEqual(tuple(control_model.hint.shape), (1, 5, 2, 2, 2))
class T2IAdapterTests(unittest.TestCase):
def test_effect_masks_are_applied_to_adapter_features(self):
control = T2IAdapterAdvanced(SimpleNamespace(), None, channels_in=3)
control.weights = ControlWeights.t2iadapter()
control.latent_keyframes = None
control.tk_mask_cond_hint = None
control._current_timestep_keyframe = SimpleNamespace(strength=1.0)
masks = {
"zero": torch.zeros((1, 1, 8, 8)),
"one": torch.ones((1, 1, 8, 8)),
"half": torch.cat((torch.zeros((1, 1, 8, 4)), torch.ones((1, 1, 8, 4))), dim=3),
}
for name, mask in masks.items():
with self.subTest(name=name):
features = torch.ones((1, 4, 8, 8))
control.mask_cond_hint = mask
control.apply_advanced_strengths_and_masks(features, batched_number=1)
torch.testing.assert_close(features, mask.expand_as(features))
def test_sliding_context_extends_single_hint_to_full_latent_length(self):
control = T2IAdapterAdvanced(SimpleNamespace(), None, channels_in=3)
original_hint = torch.ones((1, 3, 8, 8))
control.cond_hint_original = original_hint
control.cond_hint = None
control.sub_idxs = [2, 3]
control.full_latent_length = 4
control.prepare_mask_cond_hint = lambda **kwargs: None
selected_hint = None
def get_control(adapter, *args, **kwargs):
nonlocal selected_hint
selected_hint = adapter.cond_hint_original.clone()
return sentinel.output
with patch.object(T2IAdapter, "get_control", get_control):
result = control.get_control_advanced(
torch.ones((2, 4, 8, 8)),
torch.ones(2),
{},
1,
{},
)
self.assertIs(result, sentinel.output)
self.assertEqual(tuple(selected_hint.shape), (2, 3, 8, 8))
torch.testing.assert_close(selected_hint, original_hint.repeat(2, 1, 1, 1))
self.assertIs(control.cond_hint_original, original_hint)
class AdvancedInpaintingApplyTests(unittest.TestCase):
def test_source_mask_and_effect_mask_stay_independent(self):
image = torch.ones((1, 2, 2, 3))