Don't materialise the PRX attention mask (#14677)
`PRXAttnProcessor2_0` builds the joint [text | image] mask and then expands it
to the full `[B, heads, L_img, L_all]` before handing it to attention. Every
backend broadcasts a mask itself, so the expansion only costs bandwidth:
* `native` / `_native_*`: torch SDPA broadcasts `attn_mask` natively
* `flex`: `_native_flex_attention` does `attn_mask.expand(batch_size,
num_heads, seq_len_q, seq_len_kv)` on a 4-D mask before building the block mask
* `xformers`: same, `attn_mask.expand(...)` for a 4-D mask
* `sage` / `aiter` / `_native_npu`: reject `attn_mask` outright
At PRX-1B's training shape (batch 32, 1024x1024, patch 32 -> 1024 image + 256
text tokens, 28 heads) the expanded mask is 1120 MiB per block, read once per
block per forward, for 16 blocks.
Passing the unexpanded `[B, 1, 1, L_all]` is bitwise identical -- verified with
`torch.equal` on both the block output and the full gradient vector, and against
an fp32 unfused-MATH reference the relative error is unchanged to 6 significant
figures.
Measured on an H200, PRX-1B, batch 32 @ 1024px, bf16 autocast, 5 warmup / 20
timed steps (fake tensors: no dataloader, no text encoder, no loss terms), with
`set_attention_backend("_native_cudnn")`:
8 GPU DDP, compiled 446.1 -> 427.5 ms/step peak 110.2 -> 75.2 GiB
1 GPU, compiled 383.8 -> 373.8 ms/step peak 65.6 -> 63.4 GiB
1 GPU, eager 702.3 -> 649.4 ms/step peak 132.7 -> 97.7 GiB
On the default `native` backend the same change is 809.1 -> 762.1 ms/step eager. D
David Bertoin committed
2ff5e5877f743b799cc389d7b60a054ce3b9ae50
Parent: e192749
Committed by GitHub <noreply@github.com>
on 9/2/2026, 11:15:32 PM