This commit is contained in:
Acly
2024-06-03 22:50:43 +02:00
parent 9cdbedb92a
commit 7c338bbd2a
4 changed files with 10 additions and 9 deletions
+1 -1
View File
@@ -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)",
+1 -1
View File
@@ -1,5 +1,5 @@
from __future__ import annotations
from PIL import Image
from typing import Any
import numpy as np
import base64
import torch
+5 -6
View File
@@ -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):
+3 -1
View File
@@ -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)