diff --git a/comfy/ldm/modules/attention.py b/comfy/ldm/modules/attention.py index 44199e1ee..87371fb20 100644 --- a/comfy/ldm/modules/attention.py +++ b/comfy/ldm/modules/attention.py @@ -344,12 +344,9 @@ except: pass def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False): - if skip_reshape: - b, _, _, dim_head = q.shape - else: - b, _, dim_head = q.shape - dim_head //= heads - + b = q.shape[0] + dim_head = q.shape[-1] + # check to make sure xformers isn't broken disabled_xformers = False if BROKEN_XFORMERS: @@ -364,41 +361,45 @@ def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_resh return attention_pytorch(q, k, v, heads, mask, skip_reshape=skip_reshape) if skip_reshape: - q, k, v = map( - lambda t: t.reshape(b * heads, -1, dim_head), + # b h k d -> b k h d + q, k, v = map( + lambda t: t.permute(0, 2, 1, 3), (q, k, v), ) + # actually do the reshaping else: + dim_head //= heads q, k, v = map( lambda t: t.reshape(b, -1, heads, dim_head), (q, k, v), ) if mask is not None: + # add a singleton batch dimension + if mask.ndim == 2: + mask = mask.unsqueeze(0) + if mask.ndim != 3: + raise ValueError(f"Bad mask shape {list(mask.shape)}. Valid shapes are [b, Nq, Nk], [1, Nq, Nk], or [Nq, Nk].") + # add a singleton heads dimension + mask = mask.unsqueeze(1) + # pad to a multiple of 8 pad = 8 - mask.shape[-1] % 8 - # if skip_reshape, then q, k, v have merged heads and batch size - if skip_reshape: - mask_out = torch.empty([q.shape[0], q.shape[1], mask.shape[-1] + pad], dtype=q.dtype, device=q.device) - # otherwise, we have separate heads and batch size - else: - mask_out = torch.empty([q.shape[0], q.shape[2], q.shape[1], mask.shape[-1] + pad], dtype=q.dtype, device=q.device) + # the xformers docs says that it's allowed to have a mask of shape (1, Nq, Nk) + # but when using separated heads, the shape has to be (B, H, Nq, Nk) + # in flux, this matrix ends up being over 1GB + # Therefore we create a mask with a singleton heads dimension, and then expand it (ie, create a view) + mask_out = torch.empty([mask.shape[0], 1, q.shape[1], mask.shape[-1] + pad], dtype=q.dtype, device=q.device) mask_out[..., :mask.shape[-1]] = mask + # doesn't this remove the padding again?? mask = mask_out[..., :mask.shape[-1]] + mask = mask.expand(b, heads, -1, -1) out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask) - if skip_reshape: - out = ( - out.unsqueeze(0) - .reshape(b, heads, -1, dim_head) - .permute(0, 2, 1, 3) - .reshape(b, -1, heads * dim_head) - ) - else: - out = ( - out.reshape(b, -1, heads * dim_head) - ) + out = ( + out.reshape(b, -1, heads * dim_head) + ) return out @@ -420,15 +421,31 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha (q, k, v), ) - if SDP_BATCH_LIMIT >= q.shape[0]: + if mask is not None: + # add a batch dimension if there isn't already one + if mask.ndim == 2: + mask = mask.unsqueeze(0) + # add a heads dimension if there isn't already one + if mask.ndim == 3: + mask = mask.unsqueeze(1) + mask = mask.expand(b, heads, -1, -1) + + + if SDP_BATCH_LIMIT >= b: out = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False) out = ( out.transpose(1, 2).reshape(b, -1, heads * dim_head) ) else: - out = torch.empty((q.shape[0], q.shape[2], heads * dim_head), dtype=q.dtype, layout=q.layout, device=q.device) - for i in range(0, q.shape[0], SDP_BATCH_LIMIT): - out[i : i + SDP_BATCH_LIMIT] = torch.nn.functional.scaled_dot_product_attention(q[i : i + SDP_BATCH_LIMIT], k[i : i + SDP_BATCH_LIMIT], v[i : i + SDP_BATCH_LIMIT], attn_mask=mask, dropout_p=0.0, is_causal=False).transpose(1, 2).reshape(-1, q.shape[2], heads * dim_head) + out = torch.empty((b, q.shape[2], heads * dim_head), dtype=q.dtype, layout=q.layout, device=q.device) + for i in range(0, b, SDP_BATCH_LIMIT): + out[i : i + SDP_BATCH_LIMIT] = torch.nn.functional.scaled_dot_product_attention( + q[i : i + SDP_BATCH_LIMIT], + k[i : i + SDP_BATCH_LIMIT], + v[i : i + SDP_BATCH_LIMIT], + attn_mask=None if mask is None else mask[i : i + SDP_BATCH_LIMIT], + dropout_p=0.0, is_causal=False + ).transpose(1, 2).reshape(-1, q.shape[2], heads * dim_head) return out