Skip to content

QwenImage21 FlexAttention mask fails to compile on MPS because of boolean ~ #14889

Description

@korellas

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.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions