From 4130ecc0289b74b08e2817b3b1851a01247d5b90 Mon Sep 17 00:00:00 2001 From: Kawrakow Date: Mon, 6 Jul 2026 11:15:54 +0200 Subject: [PATCH] GLM-DSA: improve TG performance even more (#2068) * GLM-DSA: improve TG performance even more * Fix crash when using MTP * Fix the fix --- ggml/src/ggml-backend.cpp | 5 +- src/graphs/build_deepseek2.cpp | 111 +++++++++++++++++++++++++-------- src/llama-build-context.h | 1 + 3 files changed, 91 insertions(+), 26 deletions(-) diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index 922c0f37..1135a2de 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -2039,7 +2039,10 @@ static void ggml_backend_sched_copy_inputs(ggml_backend_sched_t sched, ggml_back ggml_backend_tensor_get_async(ids_backend, ids_tensor, ids.data(), 0, ggml_nbytes(ids_tensor)); ggml_backend_synchronize(ids_backend); - needs_sync[tensor_backend_id(ids_tensor)] = k_set_sync; + if (auto id = tensor_backend_id(ids_tensor); id >= 0 && id < GGML_SCHED_MAX_BACKENDS) { + needs_sync[id] = k_set_sync; + } + //needs_sync[tensor_backend_id(ids_tensor)] = k_set_sync; unique_ids.resize((n_expert + 31)/32); std::memset(unique_ids.data(), 0, unique_ids.size()*sizeof(uint32_t)); diff --git a/src/graphs/build_deepseek2.cpp b/src/graphs/build_deepseek2.cpp index 63ed7575..102304c8 100644 --- a/src/graphs/build_deepseek2.cpp +++ b/src/graphs/build_deepseek2.cpp @@ -323,6 +323,29 @@ ggml_tensor * llm_build_context::build_deepseek2_tp_attention( return combined; } +static inline int dsa_n_top_k(const llama_context & lctx) { + int n_top_k = lctx.model.hparams.indexer_top_k; + if (lctx.cparams.dsa_top_k >= 0) n_top_k = lctx.cparams.dsa_top_k; + return n_top_k; +} +static inline bool dsa_should_use(const llama_context & lctx) { + return lctx.cparams.mtp_op_type == MTP_OP_NONE && lctx.cparams.dsa && lctx.model.arch == LLM_ARCH_GLM_DSA && !lctx.kv_self.kr_l.empty(); +} +static inline bool dsa_should_use(const llama_context & lctx, int il) { + return dsa_should_use(lctx) && lctx.model.layers[il].indexer_attn_q_b && lctx.kv_self.kr_l.size() > (size_t) il && lctx.kv_self.kr_l[il]; +} +static inline bool dsa_can_use_fast_path(const llama_context & lctx, int n_tokens, int n_kv) { + auto n_top_k = dsa_n_top_k(lctx); + if (n_tokens == 1 && n_top_k < n_kv && n_top_k%256 == 0 && lctx.cparams.flash_attn && (lctx.cparams.mla_attn == 1 || lctx.cparams.mla_attn == 3)) { + // We need this check because we need to pretend that the KV cache is F32 so we can use ggml_get_rows to extract + // the rows selected by the indexer. In order this to work, the KV cache row size must be a multiple of sizeof(float). + auto k = lctx.kv_self.k_l[0]; + auto row_size = ggml_row_size(k->type, k->ne[0]); + return row_size % sizeof(float) == 0; + } + return false; +} + // DSA lightning indexer (GLM-5.2 / DeepSeek-V3.2). CACHE-BACKED: the batch's indexer keys are // (Hadamard-rotated and) written to a persistent per-layer indexer-key cache (kv_self.kr_l[il]) at // kv_head, then the FULL [head_size, n_kv] cached key set is read back and scored against the current @@ -722,6 +745,9 @@ ggml_tensor * llm_build_context::build_deepseek2_layer_attention( ggml_tensor * sparse_mask = KQ_mask; ggml_tensor * sparse_mask_fa = KQ_mask; + bool use_dsa = dsa_should_use(lctx, il); + bool dsa_fast_path = use_dsa && dsa_can_use_fast_path(lctx, n_tokens, n_kv); + // self_attention { ggml_tensor * q = nullptr; @@ -775,27 +801,30 @@ ggml_tensor * llm_build_context::build_deepseek2_layer_attention( // for prefill AND decode (single sequence). Gate: --dsa opt-in (off by default) + // GLM_DSA arch + indexer tensors + cache. When off, the model runs the dense MLA path, // byte-identical to a build without this feature. - if (lctx.cparams.dsa && model.arch == LLM_ARCH_GLM_DSA && model.layers[il].indexer_attn_q_b - && kv_self.kr_l.size() > (size_t) il && kv_self.kr_l[il]) { + if (use_dsa) { // GLM-5.2 IndexShare: "full" layers compute their own lightning-indexer top-k; // "shared" layers reuse the previous full layer's top-k (transformers reference: // shared layer indexer=None, topk_indices=prev_topk_indices). The full/shared map is // hparams.indexer_is_full (GGUF metadata or derived config rule). At a given step all // layers share the same n_kv/n_tokens, so a full layer's argsort is valid to reuse. if (hparams.indexer_is_full[il]) { - auto sorted = build_deepseek2_dsa_indexer(gf, il, q, cur, KQ_mask, inp_pos); - if (lctx.cparams.flash_attn) { - last_sparse_mask_fa = ::build_deepseek2_dsa_fa_mask(lctx, ctx0, KQ_mask, sorted); - } else { - last_sparse_mask = build_deepseek2_dsa_sparse_mask(sorted, KQ_mask); + dsa_last_full_sorted = build_deepseek2_dsa_indexer(gf, il, q, cur, KQ_mask, inp_pos); + if (!dsa_fast_path) { + if (lctx.cparams.flash_attn) { + last_sparse_mask_fa = ::build_deepseek2_dsa_fa_mask(lctx, ctx0, KQ_mask, dsa_last_full_sorted); + } else { + last_sparse_mask = build_deepseek2_dsa_sparse_mask(dsa_last_full_sorted, KQ_mask); + } } } - if (lctx.cparams.flash_attn) { - GGML_ASSERT(last_sparse_mask_fa); - sparse_mask_fa = last_sparse_mask_fa; - } else { - GGML_ASSERT(last_sparse_mask); - sparse_mask = last_sparse_mask; + if (!dsa_fast_path) { + if (lctx.cparams.flash_attn) { + GGML_ASSERT(last_sparse_mask_fa); + sparse_mask_fa = last_sparse_mask_fa; + } else { + GGML_ASSERT(last_sparse_mask); + sparse_mask = last_sparse_mask; + } } } @@ -1009,13 +1038,27 @@ ggml_tensor * llm_build_context::build_deepseek2_layer_attention( cb(q, "q", il); if (lctx.cparams.flash_attn && (lctx.cparams.mla_attn == 1 || lctx.cparams.mla_attn == 3)) { - ggml_tensor * kv_cache_lora = ggml_view_2d(ctx0, kv_self.k_l[il], - kv_lora_rank, n_kv, - ggml_row_size(kv_self.k_l[il]->type, kv_lora_rank + n_embd_head_qk_rope), - ggml_row_size(kv_self.k_l[il]->type, n_embd_head_qk_rope)); - cb(kv_cache_lora, "kv_cache_lora", il); - kqv_compressed = ggml_flash_attn_ext(ctx0, q, kv_cache, kv_cache_lora, sparse_mask_fa, kq_scale, hparams.f_max_alibi_bias, 0.f); + if (dsa_fast_path) { + auto row_size = ggml_row_size(kv_self.k_l[il]->type, kv_self.k_l[il]->ne[0]); + GGML_ASSERT(row_size % sizeof(float) == 0); + GGML_ASSERT(dsa_last_full_sorted); + GGML_ASSERT(ggml_nrows(dsa_last_full_sorted) == 1); + auto k32 = ggml_reshape_4d_ext(ctx0, kv_self.k_l[il], GGML_TYPE_F32, row_size/sizeof(float), + kv_self.k_l[il]->ne[1], kv_self.k_l[il]->ne[2], kv_self.k_l[il]->ne[3]); + k32 = ggml_get_rows(ctx0, k32, dsa_last_full_sorted); + auto k = ggml_reshape_4d_ext(ctx0, k32, kv_self.k_l[il]->type, kv_self.k_l[il]->ne[0], k32->ne[1], k32->ne[2], k32->ne[3]); + auto v = ggml_view_2d(ctx0, k, kv_lora_rank, k->ne[1], k->nb[1], ggml_row_size(k->type, n_embd_head_qk_rope)); + kqv_compressed = ggml_flash_attn_ext(ctx0, q, k, v, dsa_tg_fast_mask, kq_scale, hparams.f_max_alibi_bias, 0.f); + } else { + ggml_tensor * kv_cache_lora = ggml_view_2d(ctx0, kv_self.k_l[il], + kv_lora_rank, n_kv, + ggml_row_size(kv_self.k_l[il]->type, kv_lora_rank + n_embd_head_qk_rope), + ggml_row_size(kv_self.k_l[il]->type, n_embd_head_qk_rope)); + cb(kv_cache_lora, "kv_cache_lora", il); + + kqv_compressed = ggml_flash_attn_ext(ctx0, q, kv_cache, kv_cache_lora, sparse_mask_fa, kq_scale, hparams.f_max_alibi_bias, 0.f); + } cb(kqv_compressed, "kqv_compressed", il); if (use_f32_attn_precision) { @@ -1202,7 +1245,7 @@ ggml_cgraph * llm_build_context::build_deepseek2() { // KQ_mask (mask for 1 head, it will be broadcasted to all heads) struct ggml_tensor * KQ_mask = build_inp_KQ_mask(); - if (lctx.cparams.dsa && model.arch == LLM_ARCH_GLM_DSA) { + if (dsa_should_use(lctx)) { static const int n_sink = []{ const char * e = getenv("DSA_SINK"); return e ? atoi(e) : 1; }(); if (n_sink > 0 && n_sink < (int) n_kv) { lctx.inp_dsa_sink = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_kv, n_tokens); @@ -1210,13 +1253,31 @@ ggml_cgraph * llm_build_context::build_deepseek2() { ggml_set_input(lctx.inp_dsa_sink); ggml_build_forward_expand(gf, lctx.inp_dsa_sink); } - auto minus_inf = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, KQ_mask->ne[0], n_tokens); - minus_inf = ggml_fill_inplace(ctx0, minus_inf, -INFINITY); - ggml_build_forward_expand(gf, minus_inf); - lctx.inp_mask_inf = minus_inf; + if (dsa_can_use_fast_path(lctx, n_tokens, n_kv)) { + // For DSA token generation (i.e., n_tokens = 1), when we have more than n_top_k entries in the KV cache, + // we can simply pick the selected n_top_k entries from the cache via ggml_get_rows, so the mask we + // require is all ON (ideally we shouldn't require a mask, but I think we have places in the FA + // implementation which assume that there is a mask). + int n_top_k = dsa_n_top_k(lctx); + int n_padded = GGML_PAD(n_top_k, 256); + auto minus_inf = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, n_padded*KQ_mask->ne[1] - n_top_k); + minus_inf = ggml_fill_inplace(ctx0, minus_inf, -INFINITY); + auto zero = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, n_top_k); + zero = ggml_fill_inplace(ctx0, zero, 0.0f); + auto mask = ggml_concat(ctx0, zero, minus_inf, 0); + mask = ggml_cast(ctx0, mask, GGML_TYPE_F16); + mask = ggml_reshape_2d(ctx0, mask, n_padded, KQ_mask->ne[1]); + ggml_build_forward_expand(gf, mask); + dsa_tg_fast_mask = mask; + } else { + auto minus_inf = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, KQ_mask->ne[0], n_tokens); + minus_inf = ggml_fill_inplace(ctx0, minus_inf, -INFINITY); + ggml_build_forward_expand(gf, minus_inf); + lctx.inp_mask_inf = minus_inf; + } } - last_sparse_mask = last_sparse_mask_fa = nullptr; + last_sparse_mask = last_sparse_mask_fa = dsa_last_full_sorted = nullptr; // whether to use n_tokens as the matrix dimension during multiplication or n_head // n_tokens is higher during prompt processing, this allows to optimize for this case diff --git a/src/llama-build-context.h b/src/llama-build-context.h index b168cc97..32a66739 100644 --- a/src/llama-build-context.h +++ b/src/llama-build-context.h @@ -102,6 +102,7 @@ struct llm_build_context { struct ggml_tensor * dsa_last_full_sorted = nullptr; struct ggml_tensor * last_sparse_mask_fa = nullptr; struct ggml_tensor * last_sparse_mask = nullptr; + struct ggml_tensor * dsa_tg_fast_mask = nullptr; // TODO: consider making the entire interface noexcept llm_build_context(