metal: pipeline key blocks for short non-causal NAX attention - #18
Merged
Merged
Conversation
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).
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
marked this pull request as ready for review
September 27, 2026 21:43
bri-prism
changed the base branch from
feat/mlxfast-fewrow-kernels
to
prism
September 28, 2026 04:34
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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
M5 Pro local microbenchmarks (best of 30x10 calls, 3 alternating runs per build, 24 query heads, 4 KV heads, head dim 128, no mask):
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.