Describe the bug
build_qwenimage21_block_causal_mask currently returns
allowed & ~is_padding from its FlexAttention mask_mod:
|
def mask_mod(batch_idx, head_idx, q_idx, kv_idx): |
|
is_padding = (q_idx >= seq_len) | (kv_idx >= seq_len) |
|
q_image_id, kv_image_id = image_ids[q_idx], image_ids[kv_idx] |
|
same_image_block = (q_image_id == kv_image_id) & (q_image_id >= 0) |
|
allowed = ((q_idx >= kv_idx) | same_image_block) & key_valid[batch_idx, kv_idx] |
|
return allowed & ~is_padding |
|
|
|
return create_block_mask( |
is_padding is a boolean tensor. When this block mask is consumed by
torch.compile(flex_attention, fullgraph=True) on MPS, the MPS FlexAttention
lowerer rejects the resulting aten.bitwise_not.default operation:
NotImplementedError: flex_attention on MPS does not support op
aten.bitwise_not.default in score_mod/mask_mod yet
This is related to #14821, which reports a separate earlier fullgraph-compile
failure in QwenImage21Rope. This report is specifically about the
FlexAttention mask lowering path on MPS.
A narrowly scoped candidate change is:
- return allowed & ~is_padding
+ return allowed & torch.logical_not(is_padding)
Because is_padding is boolean, this preserves the mask's boolean meaning.
In my MPS environment, torch.logical_not is accepted by the FlexAttention
lowerer. I can submit this change with a focused MPS regression test if
this approach is acceptable.
Reproduction
Use Diffusers revision 80c7ed262aeffbeb43ef13ae04baeb9b84515a69 and run this on a machine
where torch.backends.mps.is_available() is true. It uses the real QwenImage21
mask helper, synthetic tensors only, and downloads no model weights.
import torch
from torch.nn.attention.flex_attention import flex_attention
from diffusers.models.transformers.transformer_qwenimage21 import (
build_qwenimage21_block_causal_mask,
)
assert torch.backends.mps.is_available()
device = torch.device("mps")
batch_size, heads, sequence_length, head_dim = 2, 2, 97, 128
image_ids = torch.full(
(sequence_length,), -1, dtype=torch.long, device=device
)
key_valid = torch.ones(
batch_size, sequence_length, dtype=torch.bool, device=device
)
block_mask = build_qwenimage21_block_causal_mask(
image_ids, key_valid, batch_size, device
)
padded_length = block_mask.shape[-1]
q = torch.randn(batch_size, heads, padded_length, head_dim, device=device)
k = torch.randn_like(q)
v = torch.randn_like(q)
compiled_flex_attention = torch.compile(
flex_attention, fullgraph=True, dynamic=False
)
compiled_flex_attention(q, k, v, block_mask=block_mask)
torch.mps.synchronize()
With the current expression, this raises the error below. With the proposed
replacement, the no-weights reproduction completes locally.
Logs
torch._inductor.exc.InductorError: LoweringException:
NotImplementedError: flex_attention on MPS does not support op
aten.bitwise_not.default in score_mod/mask_mod yet
System Info
- Diffusers failing checkout:
80c7ed262aeffbeb43ef13ae04baeb9b84515a69
before the candidate patch.
- Current
main inspected at
e0abab83b5df05de9e7abd788643c1a7c1e42e28; it still contains the same
expression. I have not runtime-tested that exact main commit.
- Python: 3.12.13
- PyTorch: 2.13.0
- Transformers: 5.15.1
- Platform: macOS 26.5.1 (build 25F80)
- Accelerator: Apple M3 Ultra, MPS
Scope and limits
This report is a compatibility fix only; it makes no speed claim. I have not
tested CUDA or upstream CI. Local validation included the no-weights mask
reproduction and patched real-generation checks, but the latter are not needed
to reproduce this compiler-lowering failure.
Describe the bug
build_qwenimage21_block_causal_maskcurrently returnsallowed & ~is_paddingfrom its FlexAttentionmask_mod:diffusers/src/diffusers/models/transformers/transformer_qwenimage21.py
Lines 291 to 298 in e0abab8
is_paddingis a boolean tensor. When this block mask is consumed bytorch.compile(flex_attention, fullgraph=True)on MPS, the MPS FlexAttentionlowerer rejects the resulting
aten.bitwise_not.defaultoperation:This is related to #14821, which reports a separate earlier fullgraph-compile
failure in
QwenImage21Rope. This report is specifically about theFlexAttention mask lowering path on MPS.
A narrowly scoped candidate change is:
Because
is_paddingis boolean, this preserves the mask's boolean meaning.In my MPS environment,
torch.logical_notis accepted by the FlexAttentionlowerer. I can submit this change with a focused MPS regression test if
this approach is acceptable.
Reproduction
Use Diffusers revision
80c7ed262aeffbeb43ef13ae04baeb9b84515a69and run this on a machinewhere
torch.backends.mps.is_available()is true. It uses the real QwenImage21mask helper, synthetic tensors only, and downloads no model weights.
With the current expression, this raises the error below. With the proposed
replacement, the no-weights reproduction completes locally.
Logs
System Info
80c7ed262aeffbeb43ef13ae04baeb9b84515a69before the candidate patch.
maininspected ate0abab83b5df05de9e7abd788643c1a7c1e42e28; it still contains the sameexpression. I have not runtime-tested that exact
maincommit.Scope and limits
This report is a compatibility fix only; it makes no speed claim. I have not
tested CUDA or upstream CI. Local validation included the no-weights mask
reproduction and patched real-generation checks, but the latter are not needed
to reproduce this compiler-lowering failure.