From 72b779204b15b9ae8295cfceb8a724dd9b92cbf8 Mon Sep 17 00:00:00 2001 From: Acly Date: Wed, 29 May 2024 13:29:30 +0200 Subject: [PATCH] Attention couple: remove base mask parameter - first mask is first element in the list - first conditioning is ignored (ksampler input is used instead) --- region.py | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/region.py b/region.py index b8dc3b7..d0083b4 100644 --- a/region.py +++ b/region.py @@ -48,7 +48,6 @@ class AttentionCouple: return { "required": { "model": ("MODEL",), - "base_mask": ("MASK",), "regions": ("LIST",), } } @@ -61,16 +60,16 @@ class AttentionCouple: conds: list[Tensor] batch_size: int - def attention_couple(self, model: ModelPatcher, base_mask: Tensor, regions: ListWrapper): + def attention_couple(self, model: ModelPatcher, regions: ListWrapper): new_model = model.clone() - num_conds = len(regions.content) + 1 + num_conds = len(regions.content) - mask = torch.stack([base_mask] + [r["mask"] for r in regions.content], dim=0) + mask = torch.stack([r["mask"] for r in regions.content], 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 regions.content[1:]] num_tokens = [cond.shape[1] for cond in self.conds] def attn2_patch(q: Tensor, k: Tensor, v: Tensor, extra_options: dict):