Skip to content

hexagon: matmul and flash-atten scalability updates - #29974

Merged
max-krasnyansky merged 30 commits into
ggml-org:masterfrom
qualcomm:hexagon-mm-fa-updates
Oct 5, 2026
Merged

max-krasnyansky merged 30 commits into
ggml-org:masterfrom
qualcomm:hexagon-mm-fa-updates

Conversation

@max-krasnyansky

Copy link
Copy Markdown
Member

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

    • Partitions multi-device Flash Attention across KV heads when n_kv_heads % n_cores == 0, falling back to token-block partitioning otherwise.
    • Improves prompt processing up to +58% (e.g. Qwen3-0.6B: 6,977 -> 11,026 t/s) and cuts DDR traffic by avoiding redundant KV-cache reads across cores.
  • HMX Gating for Flash Attention

    • I did a sweep across models, thread counts, and context lengths to pick cross-over points between HMX and HVX.
    • Lowers attention latency on longer contexts by enabling HMX instead of always using HVX.
  • Matmul Workload Scaling for Row-Split Multi-Device

    • Scales down M in the host HMX solver by sess->mdev.count when row-splitting across devices.
    • Prevents VTCM over-allocation and tunes chunk sizes/pipelining to each device's actual slice rather than global M.
  • Expert-Level Work Splitting in Multi-Device MoE

    • Distributes active experts round-robin across cores when active experts >= n_cores, falling back to row splitting per expert otherwise.
    • Improves MoE decode speed by avoiding sync overhead and ragged tile fragmentation when experts process few tokens.
  • Outer Dimension Collapse for Multi-Sequence & Fused Matmuls

    • Flattens 3D/4D batched activations (ne12 * ne13 > 1) with 2D weights into 2D matmuls across MUL_MAT, fused MUL_MAT_ADD, and fused MUL_MAT_NX.
    • Speeds up batched decode and multi-sequence prefill by using HMX instead of falling back to HVX.
  • F16 Activations & Ragged N in HMX Matmuls

    • Adds F16 activation support with F16-to-F16 transfers, non-32-aligned output N with zero-padded weight tails, and 128-byte stride checks for mdev row splitting.
    • Expands HMX support to F16 activation models and ragged layer shapes instead of falling back to HVX.
  • FA and MM kernel parameters (kernel_params) cleanup

    • Packs int32_t fields to uint8_t in htp_mm_kernel_params and standardizes tensor naming across host and DSP.
    • Host precomputation of Softcap and Log2 Scale

Additional information

I'm seeing really nice perf improvements across the board.
Here are some numbers from an older Galaxy S24U (Hexagon v75).

| model           |     size | params | ngl | threads | n_ubatch |  fa | dev  |          test | master      t/s | this PR     t/s |
| --------------- | -------: | -----: | --: | ------: | -------: | --: | ---- | ------------: | --------------: | --------------: |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 |         pp500 | 1282.47 ± 16.75 |  1270.82 ± 6.54 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 |         pp800 | 1251.63 ± 42.81 | 1339.83 ± 18.36 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 |        pp1024 | 1119.04 ± 18.24 | 1210.25 ± 29.38 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 |        pp2048 | 1009.25 ± 13.42 | 1084.28 ± 22.66 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 |        pp4096 |   930.07 ± 9.65 |   995.45 ± 2.55 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 |         tg128 |    25.50 ± 0.12 |    28.79 ± 0.14 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 |  pp512 @ d500 | 1102.31 ± 52.67 | 1231.71 ± 23.58 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 |  tg128 @ d500 |    24.91 ± 0.07 |    28.63 ± 0.29 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 |  pp512 @ d800 |   964.45 ± 6.65 | 1063.96 ± 15.19 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 |  tg128 @ d800 |    24.36 ± 0.12 |    27.69 ± 0.09 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 | pp512 @ d1024 |   930.87 ± 6.71 |  1004.56 ± 6.67 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 | tg128 @ d1024 |    24.17 ± 0.02 |    25.69 ± 0.13 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 | pp512 @ d2048 |   888.54 ± 9.42 |   957.00 ± 7.66 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 | tg128 @ d2048 |    23.26 ± 0.07 |    25.49 ± 0.17 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 | pp512 @ d4096 |   828.30 ± 8.89 |  908.07 ± 10.07 |
| gemma4 E2B Q4_0 | 3.10 GiB | 4.63 B |  99 |       6 |     1024 |   1 | HTP0 | tg128 @ d4096 |    22.76 ± 0.05 |    25.55 ± 0.11 |


| model           |     size | params | ngl | threads | n_ubatch |  fa | dev  |          test |             t/s |             t/s |
| ----------------| -------: | -----: | --: | ------: | -------: | --: | ---- | ------------: | --------------: | --------------: |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 |         pp500 |  1262.95 ± 3.49 |  1404.47 ± 5.83 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 |         pp800 |  1311.29 ± 2.54 |  1464.83 ± 4.45 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 |        pp1024 |  1294.15 ± 2.32 |  1436.50 ± 3.79 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 |        pp2048 |  1258.65 ± 2.20 |  1406.46 ± 2.35 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 |        pp4096 |  1297.21 ± 4.69 |  1415.21 ± 1.91 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 |         tg128 |    31.62 ± 0.12 |    31.49 ± 0.07 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 |  pp512 @ d500 |  1259.23 ± 1.40 |  1399.44 ± 3.45 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 |  tg128 @ d500 |    30.73 ± 0.04 |    31.46 ± 0.03 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 |  pp512 @ d800 |  1227.21 ± 3.79 |  1360.80 ± 2.93 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 |  tg128 @ d800 |    30.58 ± 0.04 |    31.27 ± 0.04 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 | pp512 @ d1024 |  1221.62 ± 1.15 |  1351.74 ± 2.95 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 | tg128 @ d1024 |    29.90 ± 0.11 |    30.93 ± 0.09 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 | pp512 @ d2048 | 1313.61 ± 11.83 |  1357.18 ± 9.59 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 | tg128 @ d2048 |    28.46 ± 0.02 |    29.67 ± 0.06 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 | pp512 @ d4096 | 1199.40 ± 11.36 |  1244.86 ± 7.59 |
| qwen35 2B Q4_0  | 1.13 GiB | 1.88 B |  99 |       6 |     1024 |   1 | HTP0 | tg128 @ d4096 |    27.77 ± 0.27 |    27.93 ± 0.09 |

Please see the other PRs I listed above for additional numbers.

Requirements

  • I have read and agree with the contributing guidelines
  • AI usage disclosure: Yes, the other PRs included here use AI tools and I used Anti.G to review and refactoring

ebateni and others added 28 commits October 4, 2026 11:21
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).
@max-krasnyansky
max-krasnyansky requested a review from a team as a code owner October 5, 2026 05:50
@github-actions github-actions Bot added the ggml changes relating to the ggml tensor library for machine learning label Oct 5, 2026
@max-krasnyansky

Copy link
Copy Markdown
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.

@jhen0409 jhen0409 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Tested on IQ-9075 (same approach in #29779) and looks no regression.

@max-krasnyansky

Copy link
Copy Markdown
Member Author

@lhez can you please approve again.
I did yet another review pass and realized we had some memsets into vtcm in the new ragged shape mm.
Fixed now.

@max-krasnyansky
max-krasnyansky merged commit 8345f33 into ggml-org:master Oct 5, 2026
19 of 22 checks passed
@njsyw1997

Copy link
Copy Markdown
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)
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

ggml changes relating to the ggml tensor library for machine learning Hexagon

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants