attention_pytorch: skip the chunk buffer/copy when steps==1
Addresses CodeRabbit feedback on #14838: the chunking loop always allocated a temporary output buffer and copied SDPA's result into it, even in the common case (any non-MPS device, or MPS below the size threshold) where steps==1 and the loop only ever runs once. Adds a fast path that calls scaled_dot_product_attention directly and reshapes its result, matching the original pre-fix code for that case -- no extra allocation or copy. The chunked path (steps>1) is unchanged. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01HAMqHgFBD6e9wjU8k8U6uW
This commit is contained in:
parent
7b7d48faa2
commit
55760cdeb5
|
|
@ -618,28 +618,31 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha
|
||||||
if SDP_BATCH_LIMIT >= b:
|
if SDP_BATCH_LIMIT >= b:
|
||||||
# See _mps_forced_attn_chunk_steps: on MPS a single attention matrix over
|
# See _mps_forced_attn_chunk_steps: on MPS a single attention matrix over
|
||||||
# ~2^31 elements silently corrupts, so chunk over seq_q to stay under that.
|
# ~2^31 elements silently corrupts, so chunk over seq_q to stay under that.
|
||||||
# steps stays 1 (one iteration over the full sequence, identical to the
|
|
||||||
# previous unchunked call) everywhere this doesn't apply.
|
|
||||||
steps = 1
|
steps = 1
|
||||||
if q.device.type == "mps":
|
if q.device.type == "mps":
|
||||||
elements_full = q.shape[0] * q.shape[1] * q.shape[2] * k.shape[2]
|
elements_full = q.shape[0] * q.shape[1] * q.shape[2] * k.shape[2]
|
||||||
steps = _mps_forced_attn_chunk_steps(elements_full, steps)
|
steps = _mps_forced_attn_chunk_steps(elements_full, steps)
|
||||||
slice_size = math.ceil(q.shape[2] / steps)
|
|
||||||
|
|
||||||
raw_out = torch.empty((b, q.shape[1], q.shape[2], dim_head), dtype=q.dtype, layout=q.layout, device=q.device)
|
if steps == 1:
|
||||||
for i in range(0, q.shape[2], slice_size):
|
out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False, **sdpa_extra)
|
||||||
end = i + slice_size
|
if not skip_output_reshape:
|
||||||
mask_chunk = mask
|
out = out.transpose(1, 2).reshape(b, -1, heads * dim_head)
|
||||||
if mask is not None and mask.shape[-2] != 1:
|
|
||||||
mask_chunk = mask[..., i:end, :]
|
|
||||||
raw_out[:, :, i:end] = comfy.ops.scaled_dot_product_attention(
|
|
||||||
q[:, :, i:end], k, v, attn_mask=mask_chunk, dropout_p=0.0, is_causal=False, **sdpa_extra
|
|
||||||
)
|
|
||||||
|
|
||||||
if skip_output_reshape:
|
|
||||||
out = raw_out
|
|
||||||
else:
|
else:
|
||||||
out = raw_out.transpose(1, 2).reshape(b, -1, heads * dim_head)
|
slice_size = math.ceil(q.shape[2] / steps)
|
||||||
|
raw_out = torch.empty((b, q.shape[1], q.shape[2], dim_head), dtype=q.dtype, layout=q.layout, device=q.device)
|
||||||
|
for i in range(0, q.shape[2], slice_size):
|
||||||
|
end = i + slice_size
|
||||||
|
mask_chunk = mask
|
||||||
|
if mask is not None and mask.shape[-2] != 1:
|
||||||
|
mask_chunk = mask[..., i:end, :]
|
||||||
|
raw_out[:, :, i:end] = comfy.ops.scaled_dot_product_attention(
|
||||||
|
q[:, :, i:end], k, v, attn_mask=mask_chunk, dropout_p=0.0, is_causal=False, **sdpa_extra
|
||||||
|
)
|
||||||
|
|
||||||
|
if skip_output_reshape:
|
||||||
|
out = raw_out
|
||||||
|
else:
|
||||||
|
out = raw_out.transpose(1, 2).reshape(b, -1, heads * dim_head)
|
||||||
else:
|
else:
|
||||||
out = torch.empty((b, 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, b, SDP_BATCH_LIMIT):
|
for i in range(0, b, SDP_BATCH_LIMIT):
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue