Rename attention couple -> attention mask
create a node for defining regions rather than using list mechanism
This commit is contained in:
+12
-8
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user