diff --git a/src/diffusers/models/attention_dispatch.py b/src/diffusers/models/attention_dispatch.py index 237115bac5a7..87a4494c554a 100644 --- a/src/diffusers/models/attention_dispatch.py +++ b/src/diffusers/models/attention_dispatch.py @@ -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. diff --git a/tests/models/test_attention_dispatch.py b/tests/models/test_attention_dispatch.py index 2143707aac40..fda528f900dd 100644 --- a/tests/models/test_attention_dispatch.py +++ b/tests/models/test_attention_dispatch.py @@ -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, ) @@ -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.