Add Attention Couple nodes for regional prompting

This commit is contained in:
ramyma
2024-06-26 02:59:53 +03:00
parent dd96073561
commit 18a856446a
4 changed files with 278 additions and 5 deletions
+13 -1
View File
@@ -1,7 +1,19 @@
# A8R8 ComfyUI Nodes
[A8R8](https://github.com/ramyma/a8r8) supporting nodes to integrate with [ComfyUI](https://github.com/comfyanonymous/ComfyUI)
****2024/06/26:** Added Attention Couple nodes for regional prompting [example](examples/attention_couple.png) (drag image to Comfy)
## Installation
1. Checkout this repo under `ComfyUI > custom_nodes`
1. Checkout this repo under `ComfyUI > custom_nodes` or use Comfy Manager
2. Start/restart ComfyUI
## Acknowledgments
Attention couple calculation and patch are heavily derived work - GPL-3.0 license for the relevant parts of the code.
Credits to:
* https://github.com/laksjdjf/cgem156-ComfyUI/tree/main/scripts/attention_couple
* https://github.com/Haoming02/sd-forge-couple
+250
View File
@@ -0,0 +1,250 @@
# Heavily derived work for attention couple calculation and patch - GPL-3.0 license
# Credits to https://github.com/laksjdjf/cgem156-ComfyUI/tree/main/scripts/attention_couple
# and https://github.com/Haoming02/sd-forge-couple
import torch
import torch.nn.functional as F
import math
from base64 import b64decode as decode
from io import BytesIO as bIO
from PIL import Image
import numpy as np
from functools import reduce
def repeat_div(value: int, iterations: int) -> int:
for _ in range(iterations):
value = math.ceil(value / 2)
return value
def get_mask(mask, batch_size, num_tokens, original_shape):
image_width: int = original_shape[3]
image_height: int = original_shape[2]
scale = math.ceil(math.log2(math.sqrt(image_height * image_width / num_tokens)))
size = (repeat_div(image_height, scale), repeat_div(image_width, scale))
num_conds = mask.shape[0]
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 // math.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 b64image2tensor(img: str, width: int, height: int) -> torch.Tensor:
image_bytes = decode(img)
image = Image.open(bIO(image_bytes)).convert("L")
if image.width != width or image.height != height:
image = image.resize((width, height), resample=Image.Resampling.NEAREST)
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image).unsqueeze(0)
return image
class AttentionCoupleRegion:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"cond": ("CONDITIONING",),
"mask": ("MASK",),
"weight": (
"FLOAT",
{"default": 1.0, "min": 0.01, "max": 1.0, "step": 0.01},
),
},
}
RETURN_TYPES = ("ATTENTION_COUPLE_REGION",)
RETURN_NAMES = ("region",)
FUNCTION = "attention_couple_region"
CATEGORY = "A8R8"
def attention_couple_region(self, cond, mask, weight):
return ({"cond": cond, "mask": mask, "weight": weight},)
class AttentionCoupleRegions:
@classmethod
def INPUT_TYPES(s):
return {
"required": {},
"optional": {
**reduce(
lambda acc, i: {**acc, f"region_{i}": ("ATTENTION_COUPLE_REGION",)},
range(1, 12),
{},
),
"regions": ("ATTENTION_COUPLE_REGION",),
},
}
RETURN_TYPES = ("ATTENTION_COUPLE_REGION",)
RETURN_NAMES = ("regions",)
FUNCTION = "attention_couple_regions"
CATEGORY = "A8R8"
def attention_couple_regions(self, **kwargs):
regions = kwargs.get("regions")
if regions:
assert isinstance(
regions, list
), "Regions has to be a list of regions, a single item was passed to regions."
regions = [kwargs.get(f"region_{i}") for i in range(1, 12)] + (
regions if regions else []
)
flattened_regions = reduce(
lambda acc, region: acc + region
if isinstance(region, list)
else acc + [region]
if region
else acc,
regions,
[],
)
return (flattened_regions,)
class AttentionCouple:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"base_prompt": ("CONDITIONING",),
"global_prompt_weight": (
"FLOAT",
{"default": 0.3, "min": 0.01, "max": 1.0, "step": 0.1},
),
"regions": ("ATTENTION_COUPLE_REGION",),
"width": ("INT", {"default": 1024, "min": 8, "step": 8}),
"height": ("INT", {"default": 1024, "min": 8, "step": 8}),
},
}
# INPUT_IS_LIST = True #(False, False, True,)
RETURN_TYPES = ("MODEL",)
RETURN_NAMES = ("model",)
FUNCTION = "attention_couple"
CATEGORY = "A8R8"
def attention_couple(
self, model, global_prompt_weight, base_prompt, height, width, regions, **kwargs
):
base_mask = torch.zeros((height, width)).unsqueeze(0)
global_mask = (torch.ones((height, width)) * global_prompt_weight).unsqueeze(0)
new_model = model.clone()
num_conds = len(regions) + 1
mask = [base_mask] + [
global_mask
if i == 0
else F.interpolate(
regions[i - 1]["mask"].unsqueeze(0),
size=(height, width),
mode="nearest-exact",
).squeeze(0)
* regions[i - 1]["weight"]
for i in range(0, num_conds)
]
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 = [
base_prompt[0][0] if i == 0 else regions[i - 1]["cond"][0][0]
for i in range(0, num_conds)
]
num_tokens = [cond.shape[1] for cond in self.conds]
num_conds += 1
def attn2_patch(q, k, v, extra_options):
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):
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)
new_model.set_model_attn2_patch(attn2_patch)
new_model.set_model_attn2_output_patch(attn2_output_patch)
return (new_model,)
Binary file not shown.

After

Width:  |  Height:  |  Size: 683 KiB

+15 -4
View File
@@ -3,6 +3,11 @@ from PIL import Image
import torch
import numpy as np
import io
from .attention_couple import (
AttentionCouple,
AttentionCoupleRegion,
AttentionCoupleRegions,
)
class Base64ImageInput:
@@ -13,7 +18,7 @@ class Base64ImageInput:
def INPUT_TYPES(s):
return {
"required": {
"bas64_image": ("STRING", {"multiline": False, "default": ""}),
"base64_image": ("STRING", {"multiline": False, "default": ""}),
},
}
@@ -23,9 +28,9 @@ class Base64ImageInput:
CATEGORY = "A8R8"
def process_input(self, bas64_image):
if bas64_image:
image_bytes = base64.b64decode(bas64_image)
def process_input(self, base64_image):
if base64_image:
image_bytes = base64.b64decode(base64_image)
# Open the image from bytes
image = Image.open(io.BytesIO(image_bytes))
@@ -75,10 +80,16 @@ class Base64ImageOutput:
NODE_CLASS_MAPPINGS = {
"Base64ImageInput": Base64ImageInput,
"Base64ImageOutput": Base64ImageOutput,
"AttentionCouple": AttentionCouple,
"AttentionCoupleRegion": AttentionCoupleRegion,
"AttentionCoupleRegions": AttentionCoupleRegions,
}
# A dictionary that contains the friendly/humanly readable titles for the nodes
NODE_DISPLAY_NAME_MAPPINGS = {
"Base64ImageInput": "Base64Image Input Node",
"Base64ImageOutput": "Base64Image Output Node",
"AttentionCouple": "Attention Couple",
"AttentionCoupleRegion": "Attention Couple Region",
"AttentionCoupleRegions": "Attention Couple Regions",
}