Commit 6f767fe96 for llama.cpp

commit 6f767fe960c3b97cf37fac4626c86400561ca1e4
Author: SXX <song_xiaoxi@126.com>
Date:   Mon Sep 28 21:23:31 2026 +0800

    ggml-cpu: enable tiled flash attention for non-vector-multiple head dims on x86 (#29423)

    * ggml-cpu: enable tiled flash attention for non-vector-multiple head dims on x86

    * add AVX2 support for masked loading and storing in simd_gemm_ukernel_tail

    * ggml-cpu: fix FA softcap handling for padded KV tiles

diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp
index ba00a0a73..a07e1f963 100644
--- a/ggml/src/ggml-cpu/ops.cpp
+++ b/ggml/src/ggml-cpu/ops.cpp
@@ -9037,6 +9037,11 @@ static void ggml_compute_forward_flash_attn_ext_tiled(
             simd_gemm(KQ, (const float *)Q_q, K_f32, Q_TILE_SZ, DK, KV_TILE_SZ);
             ggml_vec_scale_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, scale);

+            if (logit_softcap != 0.0f) {
+                ggml_vec_tanh_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, KQ);
+                ggml_vec_scale_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, logit_softcap);
+            }
+
             // Set padded KQ entries to -inf so softmax gives them zero weight
             if (kv_tile < KV_TILE_SZ) {
                 for (int tq = 0; tq < Q_TILE_SZ; tq++) {
@@ -9046,11 +9051,6 @@ static void ggml_compute_forward_flash_attn_ext_tiled(
                 }
             }

-            if (logit_softcap != 0.0f) {
-                ggml_vec_tanh_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, KQ);
-                ggml_vec_scale_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, logit_softcap);
-            }
-
             if (mask) {
                 ggml_vec_add_f32(tile_rows * KV_TILE_SZ, KQ, KQ, mask32);
             }
@@ -9320,7 +9320,7 @@ static void ggml_compute_forward_flash_attn_ext_f16(
                                 kv_is_f32_or_f16 &&
                                 k->type == v->type &&
                                 neq1 >= Q_TILE_SZ);
-#ifdef GGML_SIMD
+#if defined(GGML_SIMD) && !defined(__x86_64__) && !defined(_M_X64)
 #if defined(__ARM_FEATURE_SVE)
         const int64_t f32_epr = svcntw();
 #else
diff --git a/ggml/src/ggml-cpu/simd-gemm.h b/ggml/src/ggml-cpu/simd-gemm.h
index 2ebd10051..4b9396d54 100644
--- a/ggml/src/ggml-cpu/simd-gemm.h
+++ b/ggml/src/ggml-cpu/simd-gemm.h
@@ -56,6 +56,56 @@ static inline void simd_gemm_ukernel(
     }
 }

+template <int RM>
+static inline void simd_gemm_ukernel_tail(
+    float       * GGML_RESTRICT C,
+    const float * GGML_RESTRICT A,
+    const float * GGML_RESTRICT B,
+    int K, int N, int cols)
+{
+#if defined(__AVX512F__)
+    const __mmask16 mask = (1u << cols) - 1;
+    __m512 acc[RM];
+    for (int64_t i = 0; i < RM; i++) {
+        acc[i] = _mm512_maskz_loadu_ps(mask, C + i * N);
+    }
+    for (int64_t kk = 0; kk < K; kk++) {
+        const __m512 b = _mm512_maskz_loadu_ps(mask, B + kk * N);
+        for (int64_t i = 0; i < RM; i++) {
+            acc[i] = _mm512_mask3_fmadd_ps(_mm512_set1_ps(A[i * K + kk]), b, acc[i], mask);
+        }
+    }
+    for (int64_t i = 0; i < RM; i++) {
+        _mm512_mask_storeu_ps(C + i * N, mask, acc[i]);
+    }
+#elif defined(__AVX2__)
+    const __m256i mask = _mm256_cmpgt_epi32(_mm256_set1_epi32(cols), _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7));
+    __m256 acc[RM];
+    for (int64_t i = 0; i < RM; i++) {
+        acc[i] = _mm256_maskload_ps(C + i * N, mask);
+    }
+    for (int64_t kk = 0; kk < K; kk++) {
+        const __m256 b = _mm256_maskload_ps(B + kk * N, mask);
+        for (int64_t i = 0; i < RM; i++) {
+            acc[i] = GGML_F32_VEC_FMA(acc[i], b, _mm256_set1_ps(A[i * K + kk]));
+        }
+    }
+    for (int64_t i = 0; i < RM; i++) {
+        _mm256_maskstore_ps(C + i * N, mask, acc[i]);
+    }
+#else
+    for (int64_t j = 0; j < cols; j++) {
+        for (int64_t i = 0; i < RM; i++) {
+            float a = C[i * N + j];
+            for (int64_t kk = 0; kk < K; kk++) {
+                a += A[i * K + kk] * B[kk * N + j];
+            }
+            C[i * N + j] = a;
+        }
+    }
+#endif
+}
+
 // C[M x N] += A[M x K] * B[K x N]
 static void simd_gemm(
     float       * GGML_RESTRICT C,
@@ -74,14 +124,8 @@ static void simd_gemm(
         for (; jj + KN <= N; jj += KN) {
             simd_gemm_ukernel<GEMM_RM, 1>(C + jj, A, B + jj, K, N);
         }
-        for (; jj < N; jj++) {
-            for (int64_t i = 0; i < GEMM_RM; i++) {
-                float a = C[i * N + jj];
-                for (int64_t kk = 0; kk < K; kk++) {
-                    a += A[i * K + kk] * B[kk * N + jj];
-                }
-                C[i * N + jj] = a;
-            }
+        if (jj < N) {
+            simd_gemm_ukernel_tail<GEMM_RM>(C + jj, A, B + jj, K, N, N - jj);
         }

         A += GEMM_RM * K;
@@ -97,12 +141,8 @@ static void simd_gemm(
         for (; jj + KN <= N; jj += KN) {
             simd_gemm_ukernel<1, 1>(C + jj, A, B + jj, K, N);
         }
-        for (; jj < N; jj++) {
-            float a = C[jj];
-            for (int64_t kk = 0; kk < K; kk++) {
-                a += A[kk] * B[kk * N + jj];
-            }
-            C[jj] = a;
+        if (jj < N) {
+            simd_gemm_ukernel_tail<1>(C + jj, A, B + jj, K, N, N - jj);
         }

         A += K;
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index e11fb751a..17ac9dafe 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -11009,6 +11009,9 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
     // asymmetric head_dim (hsk != hsv) with one or both sides not 64-aligned
     test_cases.emplace_back(new test_flash_attn_ext(72, 64, 4, {1, 1}, 256, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
     test_cases.emplace_back(new test_flash_attn_ext(64, 72, 4, {1, 1}, 256, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
+    test_cases.emplace_back(new test_flash_attn_ext(65, 67, 4, {1, 1}, 113, 75, true, true, 8.0f, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
+    test_cases.emplace_back(new test_flash_attn_ext(65, 67, 4, {1, 1}, 17, 75, false, false, 0, 1.0f, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
+    test_cases.emplace_back(new test_flash_attn_ext(65, 67, 4, {1, 1}, 113, 75, false, false, 0, 1.0f, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));

     // mixed quant and Q1_0 test cases
     test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q4_0));