Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions src/diffusers/models/attention_dispatch.py
Original file line number Diff line number Diff line change
Expand Up @@ -956,6 +956,11 @@ def _cudnn_attention_forward_op(
if enable_gqa:
raise ValueError("`enable_gqa` is not yet supported for cuDNN attention.")

# The aten op takes an additive bias, so a boolean mask has to be converted the same way
# `F.scaled_dot_product_attention` does before dispatching to it.
if attn_mask is not None and attn_mask.dtype == torch.bool:
attn_mask = torch.zeros_like(attn_mask, dtype=query.dtype).masked_fill_(attn_mask.logical_not(), float("-inf"))

# The backward pass always needs the log-sum-exp, so compute it whenever a gradient may be
# required — not only when the caller asked for it via `return_lse`. Otherwise training with
# this backend (e.g. under context parallelism) would save `lse=None` and produce wrong grads.
Expand Down
31 changes: 30 additions & 1 deletion tests/models/test_attention_dispatch.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,15 +23,17 @@
import torch.nn.functional as F

from diffusers.models._modeling_parallel import ContextParallelConfig, ParallelConfig
from diffusers.models.attention_dispatch import _cudnn_attention_forward_op, dispatch_attention_fn
from diffusers.models.attention_dispatch import attention_backend as attention_backend_ctx
from diffusers.models.attention_dispatch import dispatch_attention_fn

from ..testing_utils import (
assert_tensors_close,
is_attention,
is_context_parallel,
is_kernels_available,
is_torch_compile,
require_torch_accelerator,
require_torch_gpu,
require_torch_multi_accelerator,
torch_device,
)
Expand All @@ -42,6 +44,33 @@
GRAD_RTOL = 2e-2


@is_attention
@require_torch_gpu
class TestCudnnAttentionForwardOp:
@pytest.mark.parametrize("mask_type", ["partial", "fully_masked_row"])
def test_boolean_attn_mask_matches_sdpa(self, mask_type):
batch_size, num_heads, seq_len, head_dim = 1, 2, 16, 64
torch.manual_seed(0)

# The forward op takes `(batch_size, seq_len, num_heads, head_dim)`.
query, key, value = (
torch.randn(batch_size, seq_len, num_heads, head_dim, device=torch_device, dtype=torch.bfloat16)
for _ in range(3)
)
attn_mask = torch.ones(batch_size, num_heads, seq_len, seq_len, device=torch_device, dtype=torch.bool)
if mask_type == "partial":
attn_mask[..., 3, 5:] = False
else:
attn_mask[..., 7, :] = False

out = _cudnn_attention_forward_op(None, query, key, value, attn_mask=attn_mask, _save_ctx=False)
expected = F.scaled_dot_product_attention(
query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2), attn_mask=attn_mask
).transpose(1, 2)

assert_tensors_close(out, expected, atol=1e-2, rtol=1e-2, msg=f"cuDNN forward op with {mask_type} mask")


def _attention_backward_parity_worker(rank, world_size, master_port, cp_dict, attention_backend, return_dict):
"""Op-level worker: check `dispatch_attention_fn` gradients against a single-process reference.

Expand Down
Loading