Repository navigation
hexagon: matmul and flash-atten scalability updates - #29974
Merged
Merged
Conversation
In row-split mode each core computes its output row shard of every MUL_MAT, but flash_attn was previously partitioning by Q tokens (flat qrow split) instead of by heads. This forced every core to read the full KV cache (all n_kv_heads), negating the memory bandwidth benefit of multicore on flash_attn. Change both HMX and HVX flash_attn kernels to partition by KV heads when n_kv_heads is divisible by n_cores: core i processes heads [i*n_kv_heads/N, (i+1)*n_kv_heads/N) exclusively, reading only its head shard of the KV cache. Falls back to the original token-block split when n_kv_heads % n_cores != 0 (e.g. Gemma-4 with 2 KV heads on 4 cores). Controlled by GGML_HEXAGON_FA_HEAD_SPLIT (default 1 = on). The flag is packed into bit 1 of the existing is_dst_fp32 kparams byte to stay within the 128-byte kernel_params blob limit. Measured gains at 4c row-split (PP t/s, ubatch=1024): Qwen3-0.6B: 6977 -> 11026 (+58%) llama-3.2-3B: 3717 -> 5522 (+49%) Qwen3.5-4B: 2739 -> 2855 (+4%) Gemma-4 MoE: no change (MoE FFN dominates, fallback path) TG is unchanged (flash_attn is a small fraction of decode time relative to the matmul+barrier cost per layer).
Member
Author
|
@lhez for review and ack @jhen0409 @njsyw1997 I tested the heck out of this but there are lots of combos. Let me know if you see any regressions. |
lhez
approved these changes
Oct 5, 2026
Member
Author
|
@lhez can you please approve again. |
lhez
approved these changes
Oct 5, 2026
Contributor
|
Looks good. Tested on SM8850 (v81). Both single op and end to end have no regression. |
edwardyoon
pushed a commit
to edwardyoon/focus-llama
that referenced
this pull request
Oct 8, 2026
* hexagon: head-parallel flash_attn partitioning for row-split multicore In row-split mode each core computes its output row shard of every MUL_MAT, but flash_attn was previously partitioning by Q tokens (flat qrow split) instead of by heads. This forced every core to read the full KV cache (all n_kv_heads), negating the memory bandwidth benefit of multicore on flash_attn. Change both HMX and HVX flash_attn kernels to partition by KV heads when n_kv_heads is divisible by n_cores: core i processes heads [i*n_kv_heads/N, (i+1)*n_kv_heads/N) exclusively, reading only its head shard of the KV cache. Falls back to the original token-block split when n_kv_heads % n_cores != 0 (e.g. Gemma-4 with 2 KV heads on 4 cores). Controlled by GGML_HEXAGON_FA_HEAD_SPLIT (default 1 = on). The flag is packed into bit 1 of the existing is_dst_fp32 kparams byte to stay within the 128-byte kernel_params blob limit. Measured gains at 4c row-split (PP t/s, ubatch=1024): Qwen3-0.6B: 6977 -> 11026 (+58%) llama-3.2-3B: 3717 -> 5522 (+49%) Qwen3.5-4B: 2739 -> 2855 (+4%) Gemma-4 MoE: no change (MoE FFN dominates, fallback path) TG is unchanged (flash_attn is a small fraction of decode time relative to the matmul+barrier cost per layer). * hex-fa: cleanup kern_params and head-split selection * hex-fa: add -fa-head-split option to run.py * hex-mdev: update matmul solver to account for reduced work in row-split scenarios * hex-mmid: better work splitting by expers in multi-dev scenarios * hex-fa: update HMX gating based on the model/n-hvx/ctx-len sweep * hex-fa: precompute softcap/scale on the host * hexagon: flatten matmul into 2d to use HMX in multi-sequence * hex-mm: cleanup kparams and use collapse to 3/4D -> 2D mapping * hex-mm: fix typo in collapse fallback * hex-mm: another pass at consistent naming for act tensors * hex-mm: add support for colapsing dims in fused matmuls * hex-build: fix WoS build errors * hex-mm: make sure to enforce dst stride in can_collapse * hex-fa: add a onliner commit for head-split check * hex-fa: remove unused local head_split var * hex-fa: tighten up can_split checks * hex-mm: update unfused paths to use act instead src1 * hex-mm: make sure to check all dsts for splitting * hexagon: fix the second weight chunk address in the batched HMX matmul prologue * hexagon: F16 activation and ragged N in the HMX matmul * hex-mm: tighten the ragged/split checks in mdev cases * hex-mm: enable MM fusion for F16 activations * hex-mm: pass tiled sizes to the solver in fused paths * hex-mmid: remove scalar divs from expert mapping loops * hex-mmid: proper cacheline safety enforcement for mdev splits * hex-mm: improve solver for mdev split scanarios and tail handling * hex-mm: remove redundant checks * hex-mm: fix fused HMX MUL_MAT_NX drops the final partial tile for quantized weights * hex-mm: better handling of ragged shapes (removes scalar memset of vtcm) --------- Co-authored-by: ebateni <ebateni@qti.qualcomm.com> Co-authored-by: Jhen-Jie Hong <iainst0409@gmail.com> Co-authored-by: Yiwei Shao <yiwei@aizip.ai> (cherry picked from commit 8345f33)
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.
Overview
This PR is a combination of Flash Attention and MatMul changes from a draft by @ebateni, #29779 by @jhen0409 and ##29626 by @njsyw1997 with a bunch of further rework and optimizations by me. All that stuff was targeting the same areas, and the PRs would've needed quite a bit of rebasing and followup. So I just combined them here and addressed gaps and things, and tested all together on all my setups.
The changes include:
Head-parallel Flash Attention partitioning
HMX Gating for Flash Attention
Matmul Workload Scaling for Row-Split Multi-Device
Expert-Level Work Splitting in Multi-Device MoE
Outer Dimension Collapse for Multi-Sequence & Fused Matmuls
F16 Activations & Ragged N in HMX Matmuls
FA and MM kernel parameters (kernel_params) cleanup
Additional information
I'm seeing really nice perf improvements across the board.
Here are some numbers from an older Galaxy S24U (Hexagon v75).
Please see the other PRs I listed above for additional numbers.
Requirements