Skip to content

metal: pipeline key blocks for short non-causal NAX attention - #18

Merged
bri-prism merged 1 commit into
prismfrom
feat/yukon-attn-shortq
Sep 28, 2026
Merged

bri-prism merged 1 commit into
prismfrom
feat/yukon-attn-shortq

Conversation

@bri-prism

Copy link
Copy Markdown

A speculative-decoding verify step (DFlash-style draft blocks) sends 9 to 16 queries through the full NAX attention kernel. With one 16-row query unit of work per threadgroup, the kernel walks the keys one 32-row block at a time and the four simdgroups spend most of each step waiting on the same serial chain.

This change adds a path for that shape. When the whole query block fits in one 16-row unit (head dim 128, no causal mask, no sinks; additive and boolean masks are supported), each of the four simdgroups scores one key block of a round of four. They share the row maxima through threadgroup memory, then fold the four blocks into the output in key order. Each output value is produced by the same operations on the same operands as the existing loop, and outputs are bitwise identical to the current kernel on every shape tested. Causal attention, sinks, head dim 256 and queries longer than 16 keep the existing path.

Stacked on #17, which enables the NAX path on generation-17 GPUs. Without it the kernel is not reached on M5.

Adapted from the Yukon MLX.fast promoted submission by Meganpark980320, later carried by terrapinelf (Layr-Labs/mlxfast-bonsai2-27b-engine, Apache-2.0).

Testing

  • Outputs vs the current kernel: bitwise identical on 54 cases (kL 128/1000/4096, qL 9/16/17, fp16/bf16, no mask/causal/additive mask).
  • Against a CPU float32 reference: 540 cases (head dim 128 and 256, kL 128 to 4096, qL 1 to 64, fp16/bf16/fp32, no mask, causal, additive, boolean, sinks). Worst relative error 2.4e-2 (bf16), the same as before the change on every case.
  • A temporary probe build that zeroes the new path's output confirmed it runs for qL 9 to 16 without a causal mask, and that causal and qL 17 fall back.

M5 Pro local microbenchmarks (best of 30x10 calls, 3 alternating runs per build, 24 query heads, 4 KV heads, head dim 128, no mask):

kL qL before (us) after (us) speedup
1024 9 20.0 12.7 1.57x
1024 16 19.3 10.9 1.77x
4096 9 43.7 23.0 1.90x
4096 16 44.2 22.3 1.98x
16384 9 167.6 59.6 2.81x
16384 16 165.8 59.7 2.78x

fp16 is shown; bf16 is the same within a few percent. qL 17 and 64, and all causal cases, are unchanged within noise (0.97x to 1.02x). At kL 256 both kernels take about 14 us and the difference is within run-to-run spread.

A DFlash-style draft block sends 9 to 16 queries through the full NAX
attention kernel, which then walks the keys one 32-row block at a time
with only one query tile of work. When the query block fits in one
16-row unit (head dim 128, no causal mask, no sinks), let the four
simdgroups each score one key block of a round of four, share the row
maxima through threadgroup memory, and then fold the four blocks into
the output in key order. Each output value comes from the same
operations on the same operands as the existing loop; causal, sinks
and longer queries keep the existing path.

Adapted from the Yukon MLX.fast promoted submission by Meganpark980320,
later carried by terrapinelf (Layr-Labs/mlxfast-bonsai2-27b-engine,
Apache-2.0).
@bri-prism

Copy link
Copy Markdown
Author

@khosravipasha could you review this when you get a chance? It adds a pipelined NAX attention path for 9 to 16 query verify blocks (1.6 to 2.8x on the kernel, bitwise identical output). It stacks on #17.

@bri-prism
bri-prism marked this pull request as ready for review September 27, 2026 21:43
@bri-prism
bri-prism changed the base branch from feat/mlxfast-fewrow-kernels to prism September 28, 2026 04:34
@bri-prism
bri-prism merged commit e900f87 into prism Sep 28, 2026
9 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant