From 2cda8d2dafa5e87f40418f8babfa57ceb4a147b7 Mon Sep 17 00:00:00 2001 From: Joel Farthing Date: Thu, 13 Aug 2026 04:45:39 -0500 Subject: [PATCH] speculative: Allow --swa-compress with DeepSeek4 MTP (#2309) Co-authored-by: Joel Farthing <262452229+joelfarthing@users.noreply.github.com> --- src/graphs/build_deepseek4.cpp | 40 ++++++++++++++++++++++++---------- src/llama.cpp | 4 ++-- 2 files changed, 30 insertions(+), 14 deletions(-) diff --git a/src/graphs/build_deepseek4.cpp b/src/graphs/build_deepseek4.cpp index 11346f7e..d3fefc8b 100644 --- a/src/graphs/build_deepseek4.cpp +++ b/src/graphs/build_deepseek4.cpp @@ -1064,8 +1064,13 @@ static ggml_tensor * ds4_attention(ggml_cgraph * gf, ggml_context * ctx0, llm_bu } + const bool raw_compacted = kv_self.is_compacted((int) il); + // the raw plan indexes compacted rows, so a dense layer in a compacted context (an MTP + // companion's NextN layer) is cell-addressed and must bypass the plan on store and read alike + const bool plan_addressed = !kv_self.any_compacted() || raw_compacted; + ggml_tensor * raw_k_write = nullptr; - if (hparams.n_head_kv(il) == 1 && lctx.dsv4.inputs.raw_k_write_idxs != nullptr) { + if (plan_addressed && hparams.n_head_kv(il) == 1 && lctx.dsv4.inputs.raw_k_write_idxs != nullptr) { raw_k_write = dsv4_raw_cpy_k(&lctx, ctx0, kv_self.k_l[il], kv, lctx.dsv4.inputs.raw_k_write_src_idxs, lctx.dsv4.inputs.raw_k_write_idxs, gf, n_embd_head, cb, il); if (raw_k_write != nullptr) { @@ -1079,7 +1084,7 @@ static ggml_tensor * ds4_attention(ggml_cgraph * gf, ggml_context * ctx0, llm_bu llm.llm_build_kv_store(lctx, ctx0, hparams, cparams, kv_self, gf, nullptr, kv, n_tokens, llm.kv_head, cb, il); } - const bool raw_compacted = kv_self.is_compacted((int) il); + ggml_tensor * const read_idxs = plan_addressed ? lctx.dsv4.inputs.raw_k_read_idxs : nullptr; ggml_tensor * raw_k = nullptr; ggml_tensor * raw_mask = nullptr; @@ -1091,8 +1096,8 @@ static ggml_tensor * ds4_attention(ggml_cgraph * gf, ggml_context * ctx0, llm_bu n_embd_head, 1, lctx.swa_window_view.w_view, row_size, row_size, row_size*(size_t) lctx.swa_window_view.win_off); raw_mask = KQ_mask_swa_win; - } else if (hparams.n_head_kv(il) == 1 && lctx.dsv4.inputs.raw_k_read_idxs != nullptr) { - raw_k = dsv4_raw_get_k(&lctx, ctx0, kv_self.k_l[il], lctx.dsv4.inputs.raw_k_read_idxs, n_embd_head, cb, il); + } else if (hparams.n_head_kv(il) == 1 && read_idxs != nullptr) { + raw_k = dsv4_raw_get_k(&lctx, ctx0, kv_self.k_l[il], read_idxs, n_embd_head, cb, il); } if (raw_k == nullptr) { raw_k = ggml_view_3d(ctx0, kv_self.k_l[il], @@ -1104,7 +1109,7 @@ static ggml_tensor * ds4_attention(ggml_cgraph * gf, ggml_context * ctx0, llm_bu cb(raw_k, "raw_k", il); const int64_t raw_kq_n_kv = raw_compacted ? lctx.swa_window_view.w_view - : raw_k != nullptr && lctx.dsv4.raw.n_kv > 0 + : raw_k != nullptr && read_idxs != nullptr && lctx.dsv4.raw.n_kv > 0 ? lctx.dsv4.raw.n_kv : (raw_k != nullptr ? raw_k->ne[2] * raw_k->ne[3] : n_kv); const int64_t raw_attn_n_kv = raw_compacted ? lctx.swa_window_view.w_view @@ -1114,7 +1119,7 @@ static ggml_tensor * ds4_attention(ggml_cgraph * gf, ggml_context * ctx0, llm_bu } if (raw_mask == nullptr) { raw_mask = dsv4_build_raw_mask_view(ctx0, KQ_mask, - lctx.dsv4.inputs.raw_k_read_idxs, raw_kq_n_kv, n_tokens, raw_k->ne[3], cb, il); + read_idxs, raw_kq_n_kv, n_tokens, raw_k->ne[3], cb, il); cb(raw_mask, "raw_mask_view", il); raw_mask = dsv4_pad_mask_tokens(ctx0, raw_mask, n_tokens); raw_mask = dsv4_pad_raw_mask_to(ctx0, raw_mask, raw_attn_n_kv, n_tokens); @@ -1151,7 +1156,7 @@ static ggml_tensor * ds4_attention(ggml_cgraph * gf, ggml_context * ctx0, llm_bu raw_k = dsv4_repeat_streams(ctx0, raw_k, extra_k->ne[3]); if (!cparams.flash_attn && !raw_compacted) { raw_mask = dsv4_build_raw_mask_view(ctx0, KQ_mask, - lctx.dsv4.inputs.raw_k_read_idxs, raw_kq_n_kv, n_tokens, extra_k->ne[3], cb, il); + read_idxs, raw_kq_n_kv, n_tokens, extra_k->ne[3], cb, il); raw_mask = dsv4_pad_raw_mask_to(ctx0, raw_mask, raw_attn_n_kv, n_tokens); } if (!cparams.flash_attn && raw_compacted && raw_mask->ne[1] != n_tokens) { @@ -1289,10 +1294,23 @@ ggml_cgraph * llm_build_context::build_deepseek4() { // build only the mask the graph consumes; an input tensor without a consumer is never allocated ggml_tensor * KQ_mask = nullptr; ggml_tensor * KQ_mask_swa_win = nullptr; + const int n_layer_begin = is_mtp ? n_layer - hparams.nextn_predict_layers : 0; + const int n_layer_end = is_mtp ? n_layer : n_layer - hparams.nextn_predict_layers; + if (kv_self.any_compacted()) { - bool KQ_mask_swa_windowed = false; - KQ_mask_swa_win = build_swa_mask_for_graph(hparams.n_swa, /* compacted = */ true, &KQ_mask_swa_windowed); - GGML_ASSERT(KQ_mask_swa_windowed && KQ_mask_swa_win != nullptr); + // an MTP companion walks only its dense NextN layer while its other layers compact + bool walked_compacted = false, walked_dense = false; + for (int il = n_layer_begin; il < n_layer_end; ++il) { + (kv_self.is_compacted(il) ? walked_compacted : walked_dense) = true; + } + if (walked_compacted) { + bool KQ_mask_swa_windowed = false; + KQ_mask_swa_win = build_swa_mask_for_graph(hparams.n_swa, /* compacted = */ true, &KQ_mask_swa_windowed); + GGML_ASSERT(KQ_mask_swa_windowed && KQ_mask_swa_win != nullptr); + } + if (walked_dense) { + KQ_mask = hparams.n_swa > 0 ? build_inp_KQ_mask_swa() : build_inp_KQ_mask(); + } } else { KQ_mask = hparams.n_swa > 0 ? build_inp_KQ_mask_swa() : build_inp_KQ_mask(); } @@ -1332,8 +1350,6 @@ ggml_cgraph * llm_build_context::build_deepseek4() { cb(inpL, "hc_init", -1); } - const int n_layer_begin = is_mtp ? n_layer - hparams.nextn_predict_layers : 0; - const int n_layer_end = is_mtp ? n_layer : n_layer - hparams.nextn_predict_layers; for (int il = n_layer_begin; il < n_layer_end; ++il) { auto cur = ds4_attention(gf, ctx0, *this, inpL, diff --git a/src/llama.cpp b/src/llama.cpp index 0fc6a9cb..33904f22 100644 --- a/src/llama.cpp +++ b/src/llama.cpp @@ -8524,8 +8524,8 @@ struct llama_context * llama_init_from_model( } } - if (ctx->kv_self.any_compacted() && cparams.mtp) { - LLAMA_LOG_ERROR("%s: --swa-compress is not supported together with MTP speculative decoding\n", __func__); + if (ctx->kv_self.any_compacted() && cparams.mtp && model->arch != LLM_ARCH_DEEPSEEK4) { + LLAMA_LOG_ERROR("%s: --swa-compress is not supported together with MTP speculative decoding for this model\n", __func__); llama_free(ctx); return nullptr; }