Attention couple: remove base mask parameter

- first mask is first element in the list
- first conditioning is ignored (ksampler input is used instead)
This commit is contained in:
Acly
2024-05-29 13:29:30 +02:00
parent df7ebe554f
commit 72b779204b
+4 -5
View File
@@ -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):