Add Attention Couple nodes for regional prompting
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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 |
@@ -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",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user