GLM-DSA: much better PP performance (CPU-only)

This commit is contained in:
Kawrakow
2026-07-06 09:19:24 +00:00
parent 4130ecc028
commit 711fccab5d
4 changed files with 73 additions and 4 deletions
+2 -1
View File
@@ -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",
+64 -2
View File
@@ -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,64 @@ 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) {
//if (ith == 0) {
// printf("%s: could use indexer as indexer->ne[0] = %ld, neq1 = %d, nek1 = %d, nek2 = %d nth = %d\n", __func__, indexer->ne[0], neq1, nek1, nek2, nth);
//}
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);
//iqk_flash_attn_noalibi(type_q, type_mask, max_bias,
// neq3, neq2, nbq3, nbq2, nek3, nek2, nbk3, nbk2, nev3, nev2, nbv3, nbv2, ne2, ne1, nb1,
// int_type_k_in, int_type_v, Dk, Dv, 1, nkv,
// stride_q, row_size_k, row_size_v, 0,
// this_q, work_k, work_v, work_m,
// sinks, scale, softcap, this_qkv,
// nullptr, barrier, barrier_data, 0, 1, 0, nullptr);
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 +582,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;
}
+1 -1
View File
@@ -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);
+6
View File
@@ -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);