Files
llama.cpp/ggml
ggerganov c3148ebe38 ggml-metal: FA tensor kernel: 32-wide chunks, register-light global-max softmax
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).
2026-08-26 17:22:40 +03:00
..
2026-08-24 10:43:04 +03:00
2024-07-13 18:12:39 +02:00