From 7c338bbd2a5f8d0869c1732d0ba647e5c83ef54c Mon Sep 17 00:00:00 2001 From: Acly Date: Mon, 3 Jun 2024 22:50:43 +0200 Subject: [PATCH] Fixes --- __init__.py | 2 +- nodes.py | 2 +- region.py | 11 +++++------ tile.py | 4 +++- 4 files changed, 10 insertions(+), 9 deletions(-) diff --git a/__init__.py b/__init__.py index 2764724..0f04195 100644 --- a/__init__.py +++ b/__init__.py @@ -11,7 +11,7 @@ NODE_CLASS_MAPPINGS = { "ETN_GenerateTileMask": tile.GenerateTileMask, "ETN_MergeImageTile": tile.MergeImageTile, "ETN_DefineRegion": region.DefineRegion, - "ETN_AttentionCouple": region.AttentionMask, + "ETN_AttentionMask": region.AttentionMask, } NODE_DISPLAY_NAME_MAPPINGS = { "ETN_LoadImageBase64": "Load Image (Base64)", diff --git a/nodes.py b/nodes.py index d3c8c46..2e91e5b 100644 --- a/nodes.py +++ b/nodes.py @@ -1,5 +1,5 @@ +from __future__ import annotations from PIL import Image -from typing import Any import numpy as np import base64 import torch diff --git a/region.py b/region.py index ff415ae..ba05c17 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 __future__ import annotations from typing import NamedTuple import torch import torch.nn.functional as F @@ -8,8 +9,6 @@ import math from torch import Tensor, Size from comfy.model_patcher import ModelPatcher -from .nodes import ListWrapper - def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: Size) -> Tensor: h, w = original_shape[2], original_shape[3] @@ -48,7 +47,7 @@ def lcm_for_list(numbers: list[int]): class Region(NamedTuple): previous: "Region" | None mask: Tensor - conditioning: dict + conditioning: list def to_list(self): result: list[Region] = [] @@ -76,7 +75,7 @@ class DefineRegion: RETURN_TYPES = ("REGIONS",) FUNCTION = "define" - def define(self, mask: Tensor, conditioning: dict, regions: Region | None = None): + def define(self, mask: Tensor, conditioning: list, regions: Region | None = None): return (Region(regions, mask, conditioning),) @@ -104,12 +103,12 @@ class AttentionMask: region_list = regions.to_list() num_conds = len(region_list) - mask = torch.stack([r["mask"] for r in region_list], 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 region_list] + 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): diff --git a/tile.py b/tile.py index 0d513bc..b7a4646 100644 --- a/tile.py +++ b/tile.py @@ -1,7 +1,7 @@ +from __future__ import annotations import numpy as np import numpy.typing as npt import torch -from kornia.filters import box_blur from torch import Tensor IntArray = npt.NDArray[np.int_] @@ -76,6 +76,8 @@ class TileLayout: return image[self.rect(self.coord(index))] def mask(self, coord: IntArray, blend: bool): + from kornia.filters import box_blur + size = self.size(coord) padding = self.padding if blend else self.padding - self.blending s = self.start(coord, padding) - self.start(coord)