From 7b001ace7b736fda4a413b51608374aa92d8539b Mon Sep 17 00:00:00 2001 From: asagi4 <130366179+asagi4@users.noreply.github.com> Date: Mon, 9 Jun 2025 16:34:19 +0300 Subject: [PATCH] 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 --- prompt_control/attention_couple_ppm.py | 41 ++++++++++++++++---------- prompt_control/prompts.py | 38 ++++++++++++++++-------- 2 files changed, 50 insertions(+), 29 deletions(-) diff --git a/prompt_control/attention_couple_ppm.py b/prompt_control/attention_couple_ppm.py index b95832d..bad9cc4 100644 --- a/prompt_control/attention_couple_ppm.py +++ b/prompt_control/attention_couple_ppm.py @@ -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() diff --git a/prompt_control/prompts.py b/prompt_control/prompts.py index bc78030..6cdf5d9 100644 --- a/prompt_control/prompts.py +++ b/prompt_control/prompts.py @@ -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