Switch Attention Couple implementation to one based on ppm

Also removes compatibility code with older ComfyUI
This commit is contained in:
asagi4
2025-05-26 20:11:36 +03:00
parent 2534e002ad
commit c4ac37333d
5 changed files with 185 additions and 213 deletions
+1 -1
View File
@@ -27,7 +27,7 @@ You can have both installed at the same time; none of the nodes conflict.
See [features](#features) below. Things you can control via the prompt:
- Prompt editing and filtering without noodle soup
- LoRA loading and scheduling via ComfyUI's hook system
- Masking, composition and area control (regional prompting)
- Masking, composition and area control (regional prompting), with experimental attention couple support.
- Per-encoder prompts for models with multiple text encoders, such as SDXL and Flux
- Prompt operations like `BREAK` and `AND`
- Weight interpretation types (comfy, A1111, etc.)
-13
View File
@@ -30,21 +30,8 @@ NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
nodes = ["base", "lazy", "tools"]
optional_nodes = ["attnmask"]
if importlib.util.find_spec("comfy.hooks"):
nodes.extend(["hooks"])
else:
log.error("Your ComfyUI version is too old, can't import comfy.hooks. Update your installation.")
for node in nodes:
mod = importlib.import_module(f".prompt_control.nodes_{node}", package=__name__)
NODE_CLASS_MAPPINGS.update(mod.NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(mod.NODE_DISPLAY_NAME_MAPPINGS)
for node in optional_nodes:
try:
mod = importlib.import_module(f".prompt_control.nodes_{node}", package=__name__)
NODE_CLASS_MAPPINGS.update(mod.NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(mod.NODE_DISPLAY_NAME_MAPPINGS)
except ImportError:
log.info(f"Could not import optional nodes: {node}; continuing anyway")
+179
View File
@@ -0,0 +1,179 @@
# Lifted from https://github.com/pamparamm/ComfyUI-ppm/blob/c3e6b673ee2d424405dcb99aeed89f21943c89ac/nodes_ppm/attention_couple_ppm.py
# Original implementation by laksjdjf, hako-mikan, Haoming02 licensed under GPL-3.0
# https://github.com/laksjdjf/cgem156-ComfyUI/blob/1f5533f7f31345bafe4b833cbee15a3c4ad74167/scripts/attention_couple/node.py
# https://github.com/Haoming02/sd-forge-couple/blob/e8e258e982a8d149ba59a4bc43b945467604311c/scripts/attention_couple.py
import itertools
import math
import torch
import torch.nn.functional as F
from comfy.hooks import TransformerOptionsHook, HookGroup, EnumHookScope, set_hooks_for_conditioning
from comfy.model_patcher import ModelPatcher
import logging
log = logging.getLogger("comfyui-prompt-control")
COND = 0
UNCOND = 1
COND_UNCOND_COUPLE = "cond_or_uncond_couple"
def set_cond_attnmask(base_cond, base_mask, conds, masks):
hook = AttentionCoupleHook(base_mask, conds, masks)
group = HookGroup()
group.add(hook)
return set_hooks_for_conditioning(base_cond, hooks=group)
def get_mask(mask, batch_size, num_tokens, extra_options):
activations_shape = extra_options["activations_shape"]
size = activations_shape[-2:]
num_conds = mask.shape[0]
mask_downsample = F.interpolate(mask, size=size, mode="nearest")
mask_downsample_reshaped = mask_downsample.view(num_conds, num_tokens, 1).repeat_interleave(batch_size, dim=0)
return mask_downsample_reshaped
def lcm_for_list(numbers):
current_lcm = numbers[0]
for number in numbers[1:]:
current_lcm = math.lcm(current_lcm, number)
return current_lcm
class Proxy:
def __init__(self, function):
self.function = function
def to(self, *args, **kwargs):
self.function.__self__.to(*args, **kwargs)
return self
def __call__(self, *args, **kwargs):
return self.function(*args, *kwargs)
class AttentionCoupleHook(TransformerOptionsHook):
def __init__(self, base_mask, conds, masks):
super().__init__(hook_scope=EnumHookScope.HookedOnly)
self.transformers_dict = {
"patches": {
"attn2_output_patch": [Proxy(self.attn2_output_patch)],
"attn2_patch": [Proxy(self.attn2_patch)],
}
}
self.conds_kv = []
self.batch_size = 0
self.num_conds = len(conds) + 1
mask = [base_mask] + masks
mask = torch.stack(mask, dim=0)
if mask.sum(dim=0).min() <= 0:
raise ValueError("Masks contain non-filled areas")
self.mask = mask / mask.sum(dim=0, keepdim=True)
self.conds: list[torch.Tensor] = [cond[0][0] for cond in conds]
def on_apply_hooks(self, model: ModelPatcher, transformer_options: dict[str]):
if not self.conds_kv:
attn_patches = model.model_options["transformer_options"].get("patches", {}).get("attn2_patch", [])
has_negpip = any("negpip_attn" in i.__name__ for i in attn_patches)
log.debug("AttentionCouple has_negpip=%s", has_negpip)
self.conds_kv = (
[(cond[:, 0::2], cond[:, 1::2]) for cond in self.conds]
if has_negpip
else [(cond, cond) for cond in self.conds]
)
self.num_tokens_k = [cond[0].shape[1] for cond in self.conds_kv]
self.num_tokens_v = [cond[1].shape[1] for cond in self.conds_kv]
return super().on_apply_hooks(model, transformer_options)
def to(self, *args, **kwargs):
self.conds = [c.to(*args, **kwargs) for c in self.conds]
self.mask = self.mask.to(*args, **kwargs)
self.conds_kv = [(c1.to(*args, **kwargs), c2.to(*args, **kwargs)) for c1, c2 in self.conds_kv]
return self
def attn2_patch(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, extra_options):
cond_or_uncond = extra_options["cond_or_uncond"]
num_chunks = len(cond_or_uncond)
self.batch_size = q.shape[0] // num_chunks
if len(self.conds_kv) > 0:
q_chunks = q.chunk(num_chunks, dim=0)
k_chunks = k.chunk(num_chunks, dim=0)
v_chunks = v.chunk(num_chunks, dim=0)
lcm_tokens_k = lcm_for_list(self.num_tokens_k + [k.shape[1]])
lcm_tokens_v = lcm_for_list(self.num_tokens_v + [v.shape[1]])
conds_k_tensor = torch.cat(
[
cond[0].repeat(self.batch_size, lcm_tokens_k // self.num_tokens_k[i], 1)
for i, cond in enumerate(self.conds_kv)
],
dim=0,
)
conds_v_tensor = torch.cat(
[
cond[1].repeat(self.batch_size, lcm_tokens_v // self.num_tokens_v[i], 1)
for i, cond in enumerate(self.conds_kv)
],
dim=0,
)
qs, ks, vs = [], [], []
cond_or_uncond_couple = []
for i, cond_type in enumerate(cond_or_uncond):
q_target = q_chunks[i]
k_target = k_chunks[i].repeat(1, lcm_tokens_k // k.shape[1], 1)
v_target = v_chunks[i].repeat(1, lcm_tokens_v // v.shape[1], 1)
if cond_type == UNCOND:
qs.append(q_target)
ks.append(k_target)
vs.append(v_target)
cond_or_uncond_couple.append(UNCOND)
else:
qs.append(q_target.repeat(self.num_conds, 1, 1))
ks.append(torch.cat([k_target, conds_k_tensor], dim=0))
vs.append(torch.cat([v_target, conds_v_tensor], dim=0))
cond_or_uncond_couple.extend(itertools.repeat(COND, self.num_conds))
qs = torch.cat(qs, dim=0)
ks = torch.cat(ks, dim=0)
vs = torch.cat(vs, dim=0)
extra_options[COND_UNCOND_COUPLE] = cond_or_uncond_couple
return qs, ks, vs
return q, k, v
def attn2_output_patch(self, out, extra_options):
cond_or_uncond = extra_options[COND_UNCOND_COUPLE]
bs = self.batch_size
mask_downsample = get_mask(self.mask, self.batch_size, out.shape[1], extra_options)
outputs = []
cond_outputs = []
i_cond = 0
for i, cond_type in enumerate(cond_or_uncond):
pos, next_pos = i * bs, (i + 1) * bs
if cond_type == UNCOND:
outputs.append(out[pos:next_pos])
else:
pos_cond, next_pos_cond = i_cond * bs, (i_cond + 1) * bs
masked_output = out[pos:next_pos] * mask_downsample[pos_cond:next_pos_cond]
cond_outputs.append(masked_output)
i_cond += 1
cond_output = torch.stack(cond_outputs).sum(0)
outputs.append(cond_output)
return torch.cat(outputs, dim=0)
-184
View File
@@ -1,184 +0,0 @@
import logging
log = logging.getLogger("comfyui-prompt-control")
from comfy.hooks import TransformerOptionsHook, HookGroup, EnumHookScope
from comfy.ldm.modules.attention import optimized_attention
import torch.nn.functional as F
import torch
from math import sqrt, gcd
def get_mask(mask, batch_size, num_tokens, original_shape):
num_conds = mask.shape[0]
if original_shape[2] * original_shape[3] == num_tokens:
down_sample_rate = 1
elif (original_shape[2] // 2) * (original_shape[3] // 2) == num_tokens:
down_sample_rate = 2
elif (original_shape[2] // 4) * (original_shape[3] // 4) == num_tokens:
down_sample_rate = 4
else:
down_sample_rate = 8
size = (original_shape[2] // down_sample_rate, original_shape[3] // down_sample_rate)
mask_downsample = F.interpolate(mask, size=size, mode="nearest")
mask_downsample = mask_downsample.view(num_conds, num_tokens, 1).repeat_interleave(batch_size, dim=0)
return mask_downsample
def lcm(a, b):
return a * b // gcd(a, b)
def lcm_for_list(numbers):
current_lcm = numbers[0]
for number in numbers[1:]:
current_lcm = lcm(current_lcm, number)
return current_lcm
def attention_couple_simple(base_mask, conds, masks):
num_conds = len(conds) + 1
mask = [base_mask] + masks
mask = torch.stack(mask, dim=0)
assert mask.sum(dim=0).min() > 0, "There are areas that are zero in all masks."
self_mask = mask / mask.sum(dim=0, keepdim=True)
self_conds = [cond[0][0] for cond in conds]
num_tokens = [cond.shape[1] for cond in self_conds]
self_batch_size = None
def attn2_patch(q, k, v, extra_options):
nonlocal self_conds
nonlocal self_mask
nonlocal self_batch_size
assert k.mean() == v.mean(), "k and v must be the same."
device, dtype = q.device, q.dtype
if self_conds[0].device != device:
self_conds = [cond.to(device, dtype=dtype) for cond in self_conds]
if self_mask.device != device:
self_mask = self_mask.to(device, dtype=dtype)
cond_or_unconds = extra_options["cond_or_uncond"]
num_chunks = len(cond_or_unconds)
self_batch_size = q.shape[0] // num_chunks
q_chunks = q.chunk(num_chunks, dim=0)
k_chunks = k.chunk(num_chunks, dim=0)
lcm_tokens = lcm_for_list(num_tokens + [k.shape[1]])
conds_tensor = torch.cat(
[cond.repeat(self_batch_size, lcm_tokens // num_tokens[i], 1) for i, cond in enumerate(self_conds)], dim=0
)
qs, ks = [], []
for i, cond_or_uncond in enumerate(cond_or_unconds):
k_target = k_chunks[i].repeat(1, lcm_tokens // k.shape[1], 1)
if cond_or_uncond == 1: # uncond
qs.append(q_chunks[i])
ks.append(k_target)
else:
qs.append(q_chunks[i].repeat(num_conds, 1, 1))
ks.append(torch.cat([k_target, conds_tensor], dim=0))
qs = torch.cat(qs, dim=0)
ks = torch.cat(ks, dim=0).to(k)
return qs, ks, ks
def attn2_output_patch(out, extra_options):
nonlocal self_conds
nonlocal self_mask
nonlocal self_batch_size
cond_or_unconds = extra_options["cond_or_uncond"]
mask_downsample = get_mask(self_mask, self_batch_size, out.shape[1], extra_options["original_shape"])
outputs = []
pos = 0
for cond_or_uncond in cond_or_unconds:
if cond_or_uncond == 1: # uncond
outputs.append(out[pos : pos + self_batch_size])
pos += self_batch_size
else:
masked_output = (out[pos : pos + num_conds * self_batch_size] * mask_downsample).view(
num_conds, self_batch_size, out.shape[1], out.shape[2]
)
masked_output = masked_output.sum(dim=0)
outputs.append(masked_output)
pos += num_conds * self_batch_size
return torch.cat(outputs, dim=0)
transformers_dict = {"patches": {"attn2_output_patch": [attn2_output_patch], "attn2_patch": [attn2_patch]}}
hook = TransformerOptionsHook(transformers_dict=transformers_dict, hook_scope=EnumHookScope.HookedOnly)
group = HookGroup()
group.add(hook)
return group
class MaskedAttn2:
def __init__(self, mask):
self.mask = mask
def __call__(self, q, k, v, extra_options):
mask = self.mask
orig_shape = extra_options["original_shape"]
_, _, oh, ow = orig_shape
seq_len = q.shape[1]
mask_h = oh / sqrt(oh * ow / seq_len)
mask_h = int(mask_h) + int((seq_len % int(mask_h)) != 0)
mask_w = seq_len // mask_h
r = optimized_attention(q, k, v, extra_options["n_heads"])
mask = F.interpolate(mask.unsqueeze(1), size=(mask_h, mask_w), mode="nearest").squeeze(1)
mask = mask.view(mask.shape[0], -1, 1).repeat(1, 1, r.shape[2])
return mask * r
def create_attention_hook(mask):
attn_replacements = {}
mask = mask.detach().to(device="cuda", dtype=torch.float16)
masked_attention = MaskedAttn2(mask)
for id in [4, 5, 7, 8]: # id of input_blocks that have cross attention
block_indices = range(2) if id in [4, 5] else range(10) # transformer_depth
for index in block_indices:
k = ("input", id, index)
attn_replacements[k] = masked_attention
for id in range(6): # id of output_blocks that have cross attention
block_indices = range(2) if id in [3, 4, 5] else range(10) # transformer_depth
for index in block_indices:
k = ("output", id, index)
attn_replacements[k] = masked_attention
for index in range(10):
k = ("middle", 1, index)
attn_replacements[k] = masked_attention
hook = TransformerOptionsHook(
transformers_dict={"patches_replace": {"attn2": attn_replacements}}, hook_scope=EnumHookScope.HookedOnly
)
group = HookGroup()
group.add(hook)
return group
class AttentionMaskHookExperimental:
@classmethod
def INPUT_TYPES(s):
return {
"required": {"mask": ("MASK",)},
}
RETURN_TYPES = ("HOOKS",)
CATEGORY = "promptcontrol/_testing"
FUNCTION = "apply"
EXPERIMENTAL = True
DESCRIPTION = "Experimental attention masking hook. For testing only"
def apply(self, mask):
return (create_attention_hook(mask),)
NODE_CLASS_MAPPINGS = {"AttentionMaskHookExperimental": AttentionMaskHookExperimental}
NODE_DISPLAY_NAME_MAPPINGS = {}
+5 -15
View File
@@ -9,20 +9,7 @@ from .adv_encode import advanced_encode_from_tokens
from .cutoff import process_cuts
from .parser import parse_cuts
try:
from .nodes_attnmask import attention_couple_simple
from comfy.hooks import set_hooks_for_conditioning
def set_cond_attnmask(base_cond, base_mask, conds, masks):
hook = attention_couple_simple(base_mask, conds, masks)
return set_hooks_for_conditioning(base_cond, hooks=hook)
except ImportError:
def set_cond_attnmask(cond, mask):
log.info("Attention masking is not available")
return cond
from .attention_couple_ppm import set_cond_attnmask
log = logging.getLogger("comfyui-prompt-control")
@@ -524,7 +511,10 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
attn_cond, base_mask = attnmasked_prompts[0]
if len(attnmasked_prompts) > 1:
attn_cond = set_cond_attnmask(
attn_cond, base_mask, [c[0] for c in attnmasked_prompts[1:]], [c[1] for c in attnmasked_prompts[1:]]
attn_cond,
base_mask,
[c[0] for c in attnmasked_prompts[1:]],
[c[1] for c in attnmasked_prompts[1:]],
)
else:
log.warning("You must specify at least two prompt segments with ATTN() for attention couple to work")