mirror of
https://github.com/ikawrakow/ik_llama.cpp.git
synced 2026-07-21 10:15:35 +00:00
GLM-DSA: improve TG performance even more (#2068)
* GLM-DSA: improve TG performance even more * Fix crash when using MTP * Fix the fix
This commit is contained in:
@@ -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));
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user