Rename attention couple -> attention mask

create a node for defining regions rather than using list mechanism
This commit is contained in:
Acly
2024-06-03 18:32:07 +02:00
parent ccaef6966f
commit 9cdbedb92a
3 changed files with 57 additions and 66 deletions
+12 -8
View File
@@ -6,11 +6,12 @@ NODE_CLASS_MAPPINGS = {
"ETN_SendImageWebSocket": nodes.SendImageWebSocket,
"ETN_CropImage": nodes.CropImage,
"ETN_ApplyMaskToImage": nodes.ApplyMaskToImage,
"ETN_ListAppend": nodes.ListAppend,
"ETN_ListElement": nodes.ListElement,
"ETN_SplitImageTiles": tile.SplitImageTiles,
"ETN_MergeImageTiles": tile.MergeImageTiles,
"ETN_AttentionCouple": region.AttentionCouple,
"ETN_TileLayout": tile.TileLayout,
"ETN_ExtractImageTile": tile.ExtractImageTile,
"ETN_GenerateTileMask": tile.GenerateTileMask,
"ETN_MergeImageTile": tile.MergeImageTile,
"ETN_DefineRegion": region.DefineRegion,
"ETN_AttentionCouple": region.AttentionMask,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ETN_LoadImageBase64": "Load Image (Base64)",
@@ -20,7 +21,10 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ETN_ApplyMaskToImage": "Apply Mask to Image",
"ETN_ListAppend": "List 🢒 Append",
"ETN_ListElement": "List 🢒 Get Element",
"ETN_SplitImageTiles": "Tiles 🢒 Split Image",
"ETN_MergeImageTiles": "Tiles 🢒 Merge Image",
"ETN_AttentionCouple": "Region 🢒 Attention Couple",
"ETN_TileLayout": "Create Tile Layout",
"ETN_ExtractImageTile": "Extract Image Tile",
"ETN_MergeImageTile": "Merge Image Tile",
"ETN_GenerateTileMask": "Generate Tile Mask",
"ETN_DefineRegion": "Define Region",
"ETN_AttentionMask": "Regions Attention Mask",
}
-50
View File
@@ -159,53 +159,3 @@ class ApplyMaskToImage:
out = out.movedim(1, -1)
return (out,)
class ListWrapper:
content: list[dict[str, Any]]
def __init__(self, initial: list | None = None):
self.content = initial or []
class ListAppend:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {},
"optional": {
"list": ("LIST",),
"image": ("IMAGE",),
"mask": ("MASK",),
"conditioning": ("CONDITIONING",),
},
}
CATEGORY = "external_tooling"
RETURN_TYPES = ("LIST",)
FUNCTION = "append"
def append(self, list: ListWrapper | None = None, **kwargs):
if list is None:
list = ListWrapper()
list.content.append(kwargs)
return (list,)
class ListElement:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"list": ("LIST",),
"index": ("INT", {"default": 0, "min": 0}),
}
}
CATEGORY = "external_tooling"
RETURN_TYPES = ("IMAGE", "MASK", "CONDITIONING")
FUNCTION = "element"
def element(self, list: ListWrapper, index: int):
elem = list.content[index]
return (elem.get("image"), elem.get("mask"), elem.get("conditioning"))
+45 -8
View File
@@ -1,6 +1,7 @@
# Adapted from https://github.com/laksjdjf/cgem156-ComfyUI/blob/main/scripts/attention_couple/node.py
# by @laksjdjf
from typing import NamedTuple
import torch
import torch.nn.functional as F
import math
@@ -44,35 +45,71 @@ def lcm_for_list(numbers: list[int]):
return current_lcm
class AttentionCouple:
class Region(NamedTuple):
previous: "Region" | None
mask: Tensor
conditioning: dict
def to_list(self):
result: list[Region] = []
current = self
while current is not None:
result.append(current)
current = current.previous
return result
class DefineRegion:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"mask": ("MASK",),
"conditioning": ("CONDITIONING",),
},
"optional": {
"regions": ("REGIONS",),
},
}
CATEGORY = "external_tooling/regions"
RETURN_TYPES = ("REGIONS",)
FUNCTION = "define"
def define(self, mask: Tensor, conditioning: dict, regions: Region | None = None):
return (Region(regions, mask, conditioning),)
class AttentionMask:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"regions": ("LIST",),
"regions": ("REGIONS",),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "attention_couple"
CATEGORY = "external_tooling"
FUNCTION = "attention_mask"
CATEGORY = "external_tooling/regions"
mask: Tensor
conds: list[Tensor]
batch_size: int
def attention_couple(self, model: ModelPatcher, regions: ListWrapper):
def attention_mask(self, model: ModelPatcher, regions: Region):
new_model = model.clone()
num_conds = len(regions.content)
region_list = regions.to_list()
num_conds = len(region_list)
mask = torch.stack([r["mask"] for r in regions.content], dim=0)
mask = torch.stack([r["mask"] for r in region_list], dim=0)
mask_sum = mask.sum(dim=0, keepdim=True)
assert mask_sum.sum() > 0, "There are areas that are zero in all masks."
self.mask = mask / mask_sum
self.conds = [r["conditioning"][0][0] for r in regions.content]
self.conds = [r["conditioning"][0][0] for r in region_list]
num_tokens = [cond.shape[1] for cond in self.conds]
def attn2_patch(q: Tensor, k: Tensor, v: Tensor, extra_options: dict):