Switch Attention Couple implementation to one based on ppm
Also removes compatibility code with older ComfyUI
This commit is contained in:
@@ -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
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
@@ -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 = {}
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user