From 617c33da0f00e6c10b5adfe735f7fece7497feee Mon Sep 17 00:00:00 2001 From: Kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 21 Aug 2024 16:08:46 +0300 Subject: [PATCH] attn mask bugfix --- library/flux_models.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/library/flux_models.py b/library/flux_models.py index ecff935..6320420 100644 --- a/library/flux_models.py +++ b/library/flux_models.py @@ -707,9 +707,10 @@ class DoubleStreamBlock(nn.Module): # make attention mask if not None attn_mask = None if txt_attention_mask is not None: - attn_mask = txt_attention_mask # b, seq_len + # F.scaled_dot_product_attention expects attn_mask to be bool for binary mask + attn_mask = txt_attention_mask.to(torch.bool) # b, seq_len attn_mask = torch.cat( - (attn_mask, torch.ones(attn_mask.shape[0], img.shape[1]).to(attn_mask.device)), dim=1 + (attn_mask, torch.ones(attn_mask.shape[0], img.shape[1], device=attn_mask.device, dtype=torch.bool)), dim=1 ) # b, seq_len + img_len # broadcast attn_mask to all heads