Details
hexagon: matmul and flash-atten scalability updates (#29974)
- 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
Website:
Attestations:
macOS/iOS:
- macOS Apple Silicon (arm64)
- macOS Apple Silicon (arm64, KleidiAI enabled) DISABLED
- macOS Intel (x64)
- iOS XCFramework
Linux:
- Ubuntu x64 (CPU)
- Ubuntu arm64 (CPU)
- Ubuntu s390x (CPU)
- Ubuntu x64 (Vulkan)
- Ubuntu arm64 (Vulkan)
- Ubuntu x64 (CUDA 12) - CUDA 12.8 libraries
- Ubuntu x64 (CUDA 13) - CUDA 13.4 libraries
- Ubuntu arm64 (CUDA 13) - CUDA 13.4 libraries
- Ubuntu x64 (ROCm 10.0)
- Ubuntu x64 (OpenVINO)
- Ubuntu x64 (SYCL FP32)
- Ubuntu x64 (SYCL FP16)
- Linux arm64 (Snapdragon: CPU, Adreno GPU, Hexagon NPU) - setup guide
Android:
Windows:
- Windows x64 (CPU)
- Windows arm64 (CPU)
- Windows arm64 (OpenCL Adreno)
- Windows x64 (CUDA 12) - CUDA 12.4 DLLs
- Windows x64 (CUDA 13) - CUDA 13.4 DLLs
- Windows arm64 (CUDA 13) - CUDA 13.4 DLLs
- Windows x64 (Vulkan)
- Windows arm64 (Vulkan)
- Windows x64 (OpenVINO)
- Windows x64 (SYCL)
- Windows x64 (ROCm 10.0)
openEuler:
- DISABLED
- openEuler x86 (310p)
- openEuler x86 (910b, ACL Graph)
- openEuler aarch64 (310p)
- openEuler aarch64 (910b, ACL Graph)
UI: