diff --git a/ggml/src/ggml-metal/kernels/fa.metal b/ggml/src/ggml-metal/kernels/fa.metal index 9fcf096d56..e59e1f35d2 100644 --- a/ggml/src/ggml-metal/kernels/fa.metal +++ b/ggml/src/ggml-metal/kernels/fa.metal @@ -2378,19 +2378,13 @@ void kernel_flash_attn_ext_tensor_impl( auto cT_qk0 = mm_qk.template get_destination_cooperative_tensor(); - // ---- 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(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; } } }