mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-27 21:46:57 +02:00
Performance work on the MPP tensor API flash attention kernel (NBLK == 2 branch, i.e. dv > 128, e.g. 256/256 prefill). The QK^T + softmax was already hoisted out of the d-block loop (previous commit); this closes the remaining gap to the vec kernel: - C = 32 kv items per chunk (was 64): the smaller K/V operand and P tiles fit registers better. The pad buffer is written with matching 32-wide chunks when the tensor path is active (the reserved pad space is 64-wide, a superset of what the vec kernel needs). - online softmax with ONE GLOBAL running max (scalar) instead of per-row maxes: the max cancels in O/S, so any constant >= max score seen so far is valid; a global max makes the per-chunk O rescale a uniform scalar multiply with no per-element index decode. The rescale also does not overlap the tensor core (it serializes between QK^T and PV), so it is skipped entirely when the max did not change (alpha == 1) - with random or causal scores the max stabilizes after a few chunks. - softmax state in per-thread registers (M, S[QPSG], alpha, lmax/lsum), no shared memory or simdgroup barriers in this branch. - sinks correction uses the final global max (O and S are both relative to it); this also removes the per-row max tracking. Measured on Apple M5 Max (256/256, nq=512, f16 K/V, mask, paired A/B): kv=20000: vec 25.05 ms, tensor 23.60 ms (1.06x) kv=10000: vec 8.48 ms, tensor 8.10 ms (1.05x) (previously ~0.71-0.73x before the hoisting, ~0.82-0.84x after) test-backend-ops -o FLASH_ATTN_EXT: 4809/4809 pass with the tensor path enabled and disabled. Note: 16 queries per threadgroup is 2x faster in isolated microbenchmarks (K/V-bandwidth bound per query) but spills in the full kernel and is ~30% slower; 8 queries stays the sweet spot (documented in the kernel).