ggml-metal: FA tensor kernel: register-light softmax in the per-d-block path

The NBLK != 2 branch (NBLK == 1 for dv <= 128, NBLK == 4 for dv == 512) still
used the shared-memory + 4 simdgroup barriers softmax per chunk, with the O
rescale decoded the query from each accumulator element.  Apply the same
treatment as the NBLK == 2 branch:

- online-softmax state (M, S[QPSG], alpha) in per-thread registers, reduced
  with simd_max/simd_sum: no shared memory and no barriers in the chunk loop;
- one global running max (scalar) so the O rescale is a uniform scalar
  multiply;
- skip the O rescale when the max did not change (it does not overlap the
  tensor core: it serializes between QK^T and PV);
- sinks correction with the final global max;
- drop the now-unused shared memory layout (the threadgroup buffer stays in
  the ABI; the host still allocates it).

Measured on Apple M5 Max (llama-bench, Llama-2-7B Q8_0, pp2048, -fa 1):
  d0:    tensor 1750, vec 1701 tok/s (1.03x)
  d4096: tensor 1014, vec 1025 tok/s (0.99x)
  d8192: tensor  715, vec  717 tok/s (1.00x)
(128/128 FA in isolation: 31.3 vs 30.9 ms, was ~0.92x before)

test-backend-ops -o FLASH_ATTN_EXT: 4809/4809 pass with the tensor path
enabled and disabled.
This commit is contained in:
ggerganov
2026-08-26 17:49:19 +03:00
parent c3148ebe38
commit 3f56636c8b
+46 -89
View File
@@ -2378,19 +2378,13 @@ void kernel_flash_attn_ext_tensor_impl(
auto cT_qk0 = mm_qk.template get_destination_cooperative_tensor<decltype(tK0), decltype(tQ), float>();
// ---- shared memory per lane: partial max/sum per (query, thread) + M/S/alpha ----
constexpr short SHM_FLOATS = NLANES*(QPSG*32 + QPSG*32 + 3*QPSG);
threadgroup float * sh_qmax = (threadgroup float *) shmem_f16;
threadgroup float * sh_qsum = sh_qmax + NLANES*QPSG*32;
threadgroup float * sh_M = sh_qsum + NLANES*QPSG*32;
threadgroup float * sh_S = sh_M + NLANES*QPSG;
threadgroup float * sh_alpha = sh_S + NLANES*QPSG;
(void) SHM_FLOATS;
const int jb = sgitg*QPSG; // query offset within lane (sgitg == 0)
const int jqb = sgitg*QPSG*32; // (query, thread) offset (sgitg == 0)
const int t = tiitg % 32; // thread within the (only) simdgroup
// the softmax state lives in per-thread registers (global-max online
// softmax, see below): no shared memory or simdgroup barriers are needed.
// The threadgroup buffer is still in the ABI (the host allocates it).
(void) shmem_f16;
(void) tiitg;
(void) sgitg;
(void) NLANES;
const int kv = args.ne11;
const int nchunks = (kv + C - 1)/C;
@@ -2765,26 +2759,19 @@ void kernel_flash_attn_ext_tensor_impl(
if (cT_pv.is_valid_element(i)) { cT_pv[i] = 0.0f; }
}
}
for (int j = t; j < QPSG; j += 32) {
sh_M[jb + j] = -FLT_MAX / 2;
sh_S[jb + j] = 0.0f;
sh_alpha[jb + j] = 1.0f;
}
simdgroup_barrier(mem_flags::mem_threadgroup);
// online-softmax state in per-thread registers (uniform across the
// simdgroup after simd_max/simd_sum): one GLOBAL running max (scalar)
// so the O rescale is a trivial scalar multiply; the rescale does not
// overlap the tensor core, so it is skipped when the max is unchanged
float M = -FLT_MAX / 2;
float S[QPSG];
for (int j = 0; j < QPSG; ++j) { S[j] = 0.0f; }
float alpha = 1.0f;
for (int ic = 0; ic < nchunks; ++ic) {
const int k0 = ic*C;
const int kc = min((int) C, kv - k0);
// reset this thread's (q, t) column
{
#pragma clang loop unroll(full)
for (int j = 0; j < QPSG; ++j) {
sh_qmax[jqb + j*32 + t] = -FLT_MAX / 2;
sh_qsum[jqb + j*32 + t] = 0.0f;
}
}
// the last partial chunk is read from the pad buffer (padded with 0);
// the pad buffer holds only the last C kv items per (KV head, batch),
// so its rows are indexed from 0 (not k0)
@@ -2845,27 +2832,21 @@ void kernel_flash_attn_ext_tensor_impl(
}
}
// partial max -> shared (non-owned queries hold -FLT_MAX/2 -> harmless)
// warp-level reduction: one global max (scalar) + per-row sums
float lmax_g = -FLT_MAX / 2;
#pragma clang loop unroll(full)
for (int j = 0; j < QPSG; ++j) {
lmax_g = max(lmax_g, simd_max(lmax[j]));
}
{
const float m_old = M;
M = max(M, lmax_g);
alpha = exp(m_old - M);
#pragma clang loop unroll(full)
for (int j = 0; j < QPSG; ++j) {
sh_qmax[jqb + j*32 + t] = lmax[j];
S[j] = S[j]*alpha;
}
}
simdgroup_barrier(mem_flags::mem_threadgroup);
// M / alpha update: thread t (< QPSG) finalizes query t
if (t < QPSG) {
const int j = t;
float m_new = sh_M[jb + j];
#pragma clang loop unroll(full)
for (int tt = 0; tt < 32; ++tt) {
m_new = max(m_new, sh_qmax[jqb + j*32 + tt]);
}
sh_alpha[jb + j] = exp(sh_M[jb + j] - m_new);
sh_M[jb + j] = m_new;
}
simdgroup_barrier(mem_flags::mem_threadgroup);
// ---- exp in registers, partial sums ----
{
@@ -2874,39 +2855,22 @@ void kernel_flash_attn_ext_tensor_impl(
if (!cT_qk.is_valid_element(i)) { continue; }
auto idx = cT_qk.get_multidimensional_index(i);
const int q = (int) idx[0];
const float p = exp(cT_qk[i] - sh_M[jb + q]);
const float p = exp(cT_qk[i] - M);
cT_qk[i] = p; // P in registers (f32), later used as PV right input
lsum[q] += p;
}
}
{
#pragma clang loop unroll(full)
for (int j = 0; j < QPSG; ++j) {
sh_qsum[jqb + j*32 + t] = lsum[j];
}
#pragma clang loop unroll(full)
for (int j = 0; j < QPSG; ++j) {
S[j] += simd_sum(lsum[j]);
}
simdgroup_barrier(mem_flags::mem_threadgroup);
// S update: thread t (< QPSG) finalizes query t
if (t < QPSG) {
const int j = t;
float s_new = 0.0f;
#pragma clang loop unroll(full)
for (int tt = 0; tt < 32; ++tt) {
s_new += sh_qsum[jqb + j*32 + tt];
}
sh_S[jb + j] = sh_S[jb + j]*sh_alpha[jb + j] + s_new;
}
simdgroup_barrier(mem_flags::mem_threadgroup);
// ---- rescale O accumulator (per element: query = idx0) ----
// NOTE: skip the first chunk: alpha_1 = exp(-FLT_MAX/2 - m_new) = 0
if (ic > 0) {
// ---- rescale O accumulator: scalar alpha, skipped when M is
// unchanged (the pass does not overlap the tensor core) ----
if (ic > 0 && alpha != 1.0f) {
#pragma clang loop unroll(full)
for (uint i = 0; i < cT_pv.get_capacity(); ++i) {
if (!cT_pv.is_valid_element(i)) { continue; }
auto idx = cT_pv.get_multidimensional_index(i);
cT_pv[i] *= sh_alpha[jb + (int) idx[0]];
if (cT_pv.is_valid_element(i)) { cT_pv[i] *= alpha; }
}
}
@@ -2915,31 +2879,24 @@ void kernel_flash_attn_ext_tensor_impl(
auto cT_pr = mm_pvb.template get_right_input_cooperative_tensor<half, float, float>(cT_qk);
mm_pvb.run(tVb, cT_pr, cT_pv);
}
simdgroup_barrier(mem_flags::mem_threadgroup);
}
// ---- final: sinks, O /= S, element-wise write for this d block ----
{
simdgroup_barrier(mem_flags::mem_threadgroup);
// sinks: a virtual sink position with score sinks[iq2] and no O
// contribution: M' = max(M, sink), S' = S*exp(M - M') + exp(sink - M'),
// O' = O*exp(M - M'). The per-query O factor is stashed in sh_alpha
// (reused; the last-chunk alpha is not needed at this point).
// sinks: O and S are both relative to the final global max M, so
// the correction uses M (uniform across queries)
float alpha_final;
float s_sink = 0.0f;
if (has_sinks) {
const float s_sink = ((device const float *) sinks)[iq2];
for (int j = t; j < QPSG; j += 32) {
const float m = sh_M[jb + j];
const float m2 = max(m, s_sink);
sh_alpha[jb + j] = exp(m - m2);
sh_S[jb + j] = sh_S[jb + j]*sh_alpha[jb + j] + exp(s_sink - m2);
}
} else {
for (int j = t; j < QPSG; j += 32) {
sh_alpha[jb + j] = 1.0f;
s_sink = ((device const float *) sinks)[iq2];
}
const float m2 = has_sinks ? max(M, s_sink) : M;
alpha_final = has_sinks ? exp(M - m2) : 1.0f;
if (has_sinks) {
for (int j = 0; j < QPSG; ++j) {
S[j] = S[j]*alpha_final + exp(s_sink - m2);
}
}
simdgroup_barrier(mem_flags::mem_threadgroup);
// output layout is (DV, heads, batch, batch3) with DV innermost:
// offset = (iq3*ne2*ne1 + iq2 + (iq1 + sgitg*QPSG + j)*ne1)*DV + d
@@ -2958,8 +2915,8 @@ void kernel_flash_attn_ext_tensor_impl(
const int j = (int) idx[0]; // query
const int d = b*PVM + (int) idx[1]; // head dim
if (j >= 0 && j < QPSG && d >= 0 && d < (int) DV && iq1 + j < args.ne01) {
const float s = sh_S[jb + j];
op[(uint64_t) j*args.ne1*DV + (uint) d] = (s == 0.0f) ? 0.0f : cT_pv[i]*sh_alpha[jb + j]/s;
const float s = S[j];
op[(uint64_t) j*args.ne1*DV + (uint) d] = (s == 0.0f) ? 0.0f : cT_pv[i]*alpha_final/s;
}
}
}