Fix Attention Couple when combined with hooks on the CLIP (eg. LoRAs)

All clones of the AC hook must maintain the same state. This feels
a bit hacky though; there should be a better way

Fixes #119
This commit is contained in:
asagi4
2025-06-09 17:25:26 +03:00
parent f1de65f257
commit 7b001ace7b
2 changed files with 50 additions and 29 deletions
+25 -16
View File
@@ -72,7 +72,8 @@ class AttentionCoupleHook(TransformerOptionsHook):
}
self.has_negpip = False
# calculate later
# calculate later. All clones must refer to the same kv dict
self.kv = {"k": None, "v": None}
self.conds_k: list[torch.Tensor] = None
self.conds_v: list[torch.Tensor] = None
@@ -99,7 +100,6 @@ class AttentionCoupleHook(TransformerOptionsHook):
largest_shape = max(m.shape for m in masks)
if base_mask is not None:
largest_shape = max(largest_shape, base_mask.shape)
print("largest shape x", largest_shape, [m.shape for m in masks], base_mask.shape)
log.warning("Attention Couple: Masks are irregularly shaped, resizing them all to match the largest")
for i in range(len(masks)):
masks[i] = F.interpolate(masks[i].unsqueeze(1), size=largest_shape[1:], mode="nearest-exact").squeeze(1)
@@ -124,27 +124,36 @@ class AttentionCoupleHook(TransformerOptionsHook):
self.mask = mask / mask.sum(dim=0, keepdim=True)
def on_apply_hooks(self, model: ModelPatcher, transformer_options: dict[str, Any]):
if self.conds_k is None:
if self.kv["k"] is None:
self.has_negpip = model.model_options.get("ppm_negpip", False)
log.debug("AttentionCouple has_negpip=%s", self.has_negpip)
# Skip the base cond here, which is always first
if self.has_negpip:
self.conds_k = [cond[:, 0::2] for cond in self.conds[1:]]
self.conds_v = [cond[:, 1::2] for cond in self.conds[1:]]
self.kv["k"] = [cond[:, 0::2] for cond in self.conds[1:]]
self.kv["v"] = [cond[:, 1::2] for cond in self.conds[1:]]
else:
self.conds_k = self.conds_v = self.conds[1:]
self.kv["k"] = self.kv["v"] = self.conds[1:]
return super().on_apply_hooks(model, transformer_options)
def clone(self):
c: AttentionCoupleHook = super().clone()
c.initialize_regions(self._base_cond, self._conds, self._fill)
c.mask = self.mask
c.conds = self.conds
c.kv = self.kv
c.has_negpip = self.has_negpip
c.base_strength = self.base_strength
c.strengths = self.strengths
c.num_conds = self.num_conds
return c
def to(self, *args, **kwargs):
self.conds = [c.to(*args, **kwargs) for c in self.conds]
self.mask = self.mask.to(*args, **kwargs)
if self.kv["k"] is not None:
self.kv["k"] = [c.to(*args, **kwargs) for c in self.kv["k"]]
self.kv["v"] = [c.to(*args, **kwargs) for c in self.kv["v"]]
return self
def attn2_patch(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, extra_options):
@@ -152,8 +161,11 @@ class AttentionCoupleHook(TransformerOptionsHook):
cond_or_uncond_couple = extra_options[self.COND_UNCOND_COUPLE_OPTION] = list(cond_or_uncond)
num_chunks = len(cond_or_uncond)
lcm_tokens_k = math.lcm(k.shape[1], *(cond.shape[1] for cond in self.conds_k))
lcm_tokens_v = math.lcm(v.shape[1], *(cond.shape[1] for cond in self.conds_v))
conds_k = self.kv["k"]
conds_v = self.kv["v"]
lcm_tokens_k = math.lcm(k.shape[1], *(cond.shape[1] for cond in conds_k))
lcm_tokens_v = math.lcm(v.shape[1], *(cond.shape[1] for cond in conds_v))
q_chunks = q.chunk(num_chunks, dim=0)
k_chunks = k.chunk(num_chunks, dim=0)
v_chunks = v.chunk(num_chunks, dim=0)
@@ -161,20 +173,17 @@ class AttentionCoupleHook(TransformerOptionsHook):
bs = q.shape[0] // num_chunks
conds_k_tensor = conds_v_tensor = torch.cat(
[
cond.repeat(bs, lcm_tokens_k // cond.shape[1], 1) * self.strengths[i]
for i, cond in enumerate(self.conds_k)
],
[cond.repeat(bs, lcm_tokens_k // cond.shape[1], 1) * self.strengths[i] for i, cond in enumerate(conds_k)],
dim=0,
)
).to(k)
if self.has_negpip:
conds_v_tensor = torch.cat(
[
cond.repeat(bs, lcm_tokens_v // cond.shape[1], 1) * self.strengths[i]
for i, cond in enumerate(self.conds_v)
for i, cond in enumerate(conds_v)
],
dim=0,
)
).to(v)
qs, ks, vs = [], [], []
cond_or_uncond_couple.clear()
+25 -13
View File
@@ -585,20 +585,32 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
settings["start_percent"] = start_pct
settings["end_percent"] = end_pct
x = encode_prompt_segment(clip, p, settings, style, normalization)
encoded.extend(x)
encoded.append(x)
base_cond, *attention_couple = encoded
base_cond = [base_cond]
if attention_couple:
fill = base_cond[0][1].get("x-promptcontrol.fill")
if not fill:
ensure_mask(base_cond[0])
# else, set_cond_attnmask will have the base mask fill any unspecified areas
base_cond = set_cond_attnmask(
base_cond,
[ensure_mask(c) for c in attention_couple],
fill=fill,
)
assert all(
len(c) == len(encoded[0]) for c in encoded
), "All encoded prompts didn't produce the same number of conds, I don't know what to do in this situation."
# each call to encode_prompt_segment can produce a number of conds based on any
# scheduled LoRA hooks on the clip model. Zip them together with coupled prompts
for base_cond, *attention_couple in zip(*encoded):
s = base_cond[1]
# If there are LoRAs on the CLIP, we need to fix start_percent and end_percent on the new conds for things to work properly.
s["start_percent"] = s.get("clip_start_percent", s["start_percent"])
s["end_percent"] = s.get("clip_end_percent", s["end_percent"])
s.pop("clip_start_percent", None)
s.pop("clip_end_percent", None)
base_cond = [base_cond]
if attention_couple:
fill = base_cond[0][1].get("x-promptcontrol.fill")
if not fill:
ensure_mask(base_cond[0])
# else, set_cond_attnmask will have the base mask fill any unspecified areas
base_cond = set_cond_attnmask(
base_cond,
[ensure_mask(c) for c in attention_couple],
fill=fill,
)
conds.extend(base_cond)
return conds