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));