From eacd450d2b57894c8c8c3e0285e755cb93e8a036 Mon Sep 17 00:00:00 2001 From: Kawrakow Date: Mon, 6 Jul 2026 11:57:40 +0200 Subject: [PATCH] GLM-DSA: much better PP performance (CPU-only) (#2074) * GLM-DSA: much better PP performance (CPU-only) * Cleanup --- ggml/src/ggml.c | 3 +- ggml/src/iqk/iqk_flash_attn.cpp | 56 +++++++++++++++++++++++++++++++-- ggml/src/iqk/iqk_mul_mat.h | 2 +- src/graphs/build_deepseek2.cpp | 6 ++++ 4 files changed, 63 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 53aef9869..4e78b2a6b 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -21919,7 +21919,8 @@ static void ggml_compute_forward_flash_attn_ext_f16( Dk, Dv, neq1, nek1, q->nb[1], k->nb[1], v->nb[1], mask ? mask->nb[1] : 0, q->data, k->data, v->data, mask ? mask->data : NULL, sinks ? sinks->data : NULL, scale, softcap, (float *)dst->data, - params->wdata, (barrier_t)ggml_barrier, (void *)params->shared, ith, nth, dst->op_params[4])) return; + params->wdata, (barrier_t)ggml_barrier, (void *)params->shared, ith, nth, dst->op_params[4], + dst->src[5])) return; // if (max_bias <= 0.0f && q->type == GGML_TYPE_F32 && mask && mask->type == GGML_TYPE_F16) { // //if (ith == 0) printf("k: %ld x %ld x %ld, q: %ld x %ld x %ld, v: %ld x %ld x %ld mask: %ld x %ld x %ld\n", diff --git a/ggml/src/iqk/iqk_flash_attn.cpp b/ggml/src/iqk/iqk_flash_attn.cpp index c562b0c06..e1ea5177d 100644 --- a/ggml/src/iqk/iqk_flash_attn.cpp +++ b/ggml/src/iqk/iqk_flash_attn.cpp @@ -51,6 +51,14 @@ size_t iqk_fa_work_buffer_size(const struct ggml_tensor * dst, int nth) { auto Q = dst->src[0]; auto K = dst->src[1]; auto V = dst->src[2]; + auto indexer = dst->src[5]; + if (indexer && indexer->type == GGML_TYPE_I32 && indexer->ne[0] < K->ne[1] && Q->ne[1] >= nth && + Q->ne[3] == 1 && K->ne[3] == 1 && V->ne[3] == 1 && K->ne[2] == 1) { + auto row_size_k = ggml_row_size(K->type, K->ne[0]); + auto row_size_v = ggml_row_size(V->type, V->ne[0]); + auto work_size = (row_size_k + row_size_v + 64) * indexer->ne[0]; + return work_size * nth; + } int rk2 = Q->ne[2]/K->ne[2]; size_t size = 0; if (Q->ne[1] >= 8 && K->type == GGML_TYPE_Q8_0) { @@ -163,10 +171,54 @@ extern "C" IQK_API bool iqk_flash_attn_noalibi(int type_q, int type_mask, float float softcap, // if > 0, a "soft-cap" operation is applied before softmax float * qkv, // v*softmax(scale*(k*q)) [[maybe_unused]] void * work_buffer_in, [[maybe_unused]] barrier_t barrier, [[maybe_unused]] void * barrier_data, - int ith, int nth, int n_swa) { + int ith, int nth, int n_swa, [[maybe_unused]] ggml_tensor * indexer) { if (type_q != 0 || type_mask != 1 || max_bias > 0) return false; + if (indexer && indexer->type == GGML_TYPE_I32) { + if (indexer->ne[0] < nek1 && neq1 >= nth && neq3 == 1 && nek3 == 1 && nev3 == 1 && nek2 == 1) { + int npt = (neq1 + nth - 1)/nth; + int ith_mid = nth; + int neq1_this_thread = npt; + int first = ith*npt; + if (npt*nth > neq1) { + ith_mid = neq1 - nth*(npt - 1); + if (ith >= ith_mid) { + --neq1_this_thread; + first = ith_mid*npt + (ith - ith_mid)*neq1_this_thread; + } + } + // Workbuffer: we need + // * indexer->ne[0] * sizeof(ggml_half) to extract the mask for a row + // * indexer->ne[0] * ggml_row_size(int_type_k_in, Dk) to extract the selected K cache entries + // * indexer->ne[0] * ggml_row_size(int_type_v, Dv) to extract the selected V cache entries + auto row_size_k = ggml_row_size(ggml_type(int_type_k_in), Dk); + auto row_size_v = ggml_row_size(ggml_type(int_type_v ), Dv); + auto work_size = (row_size_k + row_size_v + 64) * indexer->ne[0]; + auto work_k = (char *)work_buffer_in + ith*work_size; + auto work_v = work_k + row_size_k*indexer->ne[0]; + auto work_m = (uint16_t *)(work_v + row_size_v*indexer->ne[0]); + int nkv = indexer->ne[0]; + for (int iq = first; iq < first + neq1_this_thread; ++iq) { + auto idx = (const int *)((const char *)indexer->data + iq*indexer->nb[1]); + auto M = (const uint16_t *)((const char *)mask + iq*stride_m); + for (int j = 0; j < nkv; ++j) { + std::memcpy(work_k + row_size_k*j, ((const char *)k + idx[j]*stride_k), row_size_k); + std::memcpy(work_v + row_size_v*j, ((const char *)v + idx[j]*stride_v), row_size_v); + work_m[j] = M[idx[j]]; + } + auto this_q = (const char *)q + iq*stride_q; + auto this_qkv = qkv + iq*ne1*nb1/sizeof(float); + if (!iqk_flash_attn_impl(int_type_k_in, int_type_v, + Dk, Dv, neq2, nkv, nbq2, row_size_k, row_size_v, 0, Dv, + (const float *)this_q, work_k, work_v, work_m, nullptr, 0, + scale, softcap, + this_qkv, nullptr, nullptr)) return false; + } + return true; + } + } + if (auto type_k = ggml_type(int_type_k_in), type_v = ggml_type(int_type_v); !are_kv_types_supported(type_k, type_v)) { if (ith == 0) { fprintf(stderr, "\n==================== K cache %s coupled with V cache %s is not a supported combination on the CPU backend.\n", @@ -520,7 +572,7 @@ bool iqk_flash_attn_noalibi([[maybe_unused]] int type_q, [[maybe_unused]] int ty [[maybe_unused]] float softcap, // if > 0, a "soft-cap" operation is applied before softmax [[maybe_unused]] float * qkv, // v*softmax(scale*(k*q)) [[maybe_unused]] void * work_buffer, [[maybe_unused]] barrier_t barrier, [[maybe_unused]] void * barrier_data, - [[maybe_unused]] int ith, [[maybe_unused]] int nth, [[maybe_unused]] int n_swa) { + [[maybe_unused]] int ith, [[maybe_unused]] int nth, [[maybe_unused]] int n_swa, [[maybe_unused]] ggml_tensor * indexer) { return false; } diff --git a/ggml/src/iqk/iqk_mul_mat.h b/ggml/src/iqk/iqk_mul_mat.h index 6e372f938..a51597090 100644 --- a/ggml/src/iqk/iqk_mul_mat.h +++ b/ggml/src/iqk/iqk_mul_mat.h @@ -68,7 +68,7 @@ IQK_API bool iqk_flash_attn_noalibi(int type_q, int type_mask, float max_bias, float softcap, // if > 0, a "soft-cap" operation is applied before softmax float * qkv, // v*softmax(scale*(k*q)) void * work_buffer, barrier_t barrier, void * barrier_data, - int ith, int nth, int n_swa); + int ith, int nth, int n_swa, struct ggml_tensor * indexer); IQK_API void iqk_topk_moe(int n_experts, int n_experts_used, int nrows, const float * logits, float * weights, int32_t * ids, int ith, int nth); diff --git a/src/graphs/build_deepseek2.cpp b/src/graphs/build_deepseek2.cpp index 102304c8d..5aad17060 100644 --- a/src/graphs/build_deepseek2.cpp +++ b/src/graphs/build_deepseek2.cpp @@ -1007,6 +1007,9 @@ ggml_tensor * llm_build_context::build_deepseek2_layer_attention( if (use_f32_attn_precision || q->ne[1] <= 8) { ggml_flash_attn_ext_set_prec(kqv, GGML_PREC_F32); } + if (use_dsa && dsa_last_full_sorted) { + kqv->src[5] = dsa_last_full_sorted; + } cb(kqv, "kqv", il); if (iter == 0) { @@ -1064,6 +1067,9 @@ ggml_tensor * llm_build_context::build_deepseek2_layer_attention( if (use_f32_attn_precision) { ggml_flash_attn_ext_set_prec(kqv_compressed, GGML_PREC_F32); } + if (use_dsa && dsa_last_full_sorted) { + kqv_compressed->src[5] = dsa_last_full_sorted; + } kqv_compressed = ggml_permute(ctx0, kqv_compressed, 0, 2, 1, 3); cb(kqv_compressed, "kqv_compressed_perm", il);