mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-10-08 04:27:05 +08:00
fix pytorch/xformers attention
This corrects a weird inconsistency with skip_reshape. It also allows masks of various shapes to be passed, which will be automtically expanded (in a memory-efficient way) to a size that is compatible with xformers or pytorch sdpa respectively.
This commit is contained in:
parent
66bdf74f0c
commit
5e8206362c
@ -344,12 +344,9 @@ except:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False):
|
def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False):
|
||||||
if skip_reshape:
|
b = q.shape[0]
|
||||||
b, _, _, dim_head = q.shape
|
dim_head = q.shape[-1]
|
||||||
else:
|
# check to make sure xformers isn't broken
|
||||||
b, _, dim_head = q.shape
|
|
||||||
dim_head //= heads
|
|
||||||
|
|
||||||
disabled_xformers = False
|
disabled_xformers = False
|
||||||
|
|
||||||
if BROKEN_XFORMERS:
|
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)
|
return attention_pytorch(q, k, v, heads, mask, skip_reshape=skip_reshape)
|
||||||
|
|
||||||
if skip_reshape:
|
if skip_reshape:
|
||||||
q, k, v = map(
|
# b h k d -> b k h d
|
||||||
lambda t: t.reshape(b * heads, -1, dim_head),
|
q, k, v = map(
|
||||||
|
lambda t: t.permute(0, 2, 1, 3),
|
||||||
(q, k, v),
|
(q, k, v),
|
||||||
)
|
)
|
||||||
|
# actually do the reshaping
|
||||||
else:
|
else:
|
||||||
|
dim_head //= heads
|
||||||
q, k, v = map(
|
q, k, v = map(
|
||||||
lambda t: t.reshape(b, -1, heads, dim_head),
|
lambda t: t.reshape(b, -1, heads, dim_head),
|
||||||
(q, k, v),
|
(q, k, v),
|
||||||
)
|
)
|
||||||
|
|
||||||
if mask is not None:
|
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
|
pad = 8 - mask.shape[-1] % 8
|
||||||
# if skip_reshape, then q, k, v have merged heads and batch size
|
# the xformers docs says that it's allowed to have a mask of shape (1, Nq, Nk)
|
||||||
if skip_reshape:
|
# but when using separated heads, the shape has to be (B, H, Nq, Nk)
|
||||||
mask_out = torch.empty([q.shape[0], q.shape[1], mask.shape[-1] + pad], dtype=q.dtype, device=q.device)
|
# in flux, this matrix ends up being over 1GB
|
||||||
# otherwise, we have separate heads and batch size
|
# Therefore we create a mask with a singleton heads dimension, and then expand it (ie, create a view)
|
||||||
else:
|
mask_out = torch.empty([mask.shape[0], 1, q.shape[1], mask.shape[-1] + pad], dtype=q.dtype, device=q.device)
|
||||||
mask_out = torch.empty([q.shape[0], q.shape[2], q.shape[1], mask.shape[-1] + pad], dtype=q.dtype, device=q.device)
|
|
||||||
|
|
||||||
mask_out[..., :mask.shape[-1]] = mask
|
mask_out[..., :mask.shape[-1]] = mask
|
||||||
|
# doesn't this remove the padding again??
|
||||||
mask = mask_out[..., :mask.shape[-1]]
|
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)
|
out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask)
|
||||||
|
|
||||||
if skip_reshape:
|
out = (
|
||||||
out = (
|
out.reshape(b, -1, heads * dim_head)
|
||||||
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)
|
|
||||||
)
|
|
||||||
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
@ -420,15 +421,31 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha
|
|||||||
(q, k, v),
|
(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 = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False)
|
||||||
out = (
|
out = (
|
||||||
out.transpose(1, 2).reshape(b, -1, heads * dim_head)
|
out.transpose(1, 2).reshape(b, -1, heads * dim_head)
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
out = torch.empty((q.shape[0], q.shape[2], heads * dim_head), dtype=q.dtype, layout=q.layout, device=q.device)
|
out = torch.empty((b, 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):
|
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=mask, dropout_p=0.0, is_causal=False).transpose(1, 2).reshape(-1, q.shape[2], heads * dim_head)
|
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
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user