From 9cdbedb92aeebbee32e9323dadd4a13955679948 Mon Sep 17 00:00:00 2001 From: Acly Date: Mon, 3 Jun 2024 18:32:07 +0200 Subject: [PATCH] Rename attention couple -> attention mask create a node for defining regions rather than using list mechanism --- __init__.py | 20 ++++++++++++-------- nodes.py | 50 -------------------------------------------------- region.py | 53 +++++++++++++++++++++++++++++++++++++++++++++-------- 3 files changed, 57 insertions(+), 66 deletions(-) diff --git a/__init__.py b/__init__.py index dc2b736..2764724 100644 --- a/__init__.py +++ b/__init__.py @@ -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", } diff --git a/nodes.py b/nodes.py index 4c5cb2a..d3c8c46 100644 --- a/nodes.py +++ b/nodes.py @@ -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")) diff --git a/region.py b/region.py index 2d5fc9a..ff415ae 100644 --- a/region.py +++ b/region.py @@ -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):