Support for Qwen 3.5 MTP (dense models only) (#1698)

* qwen-mtp: add dense mtp for one draft

* add support for smaller qwen mtp commit

* qwen-mtp: fix graph for qwen dense variants

* Squashed commit of the following:

commit a92a154b38c7fddc84460f8852c900f8d6ce907e
Author: SamuelOliveirads <samueloliveira32df@gmail.com>
Date:   Mon Apr 20 13:30:21 2026 -0300

    recurrent model: refactor api

commit dfac8f19f6
Author: SamuelOliveirads <samueloliveira32df@gmail.com>
Date:   Mon Apr 20 12:22:29 2026 -0300

    recurrent model: implement recurrent kernel checkpoint

commit 9c44b117f9
Author: SamuelOliveirads <samueloliveira32df@gmail.com>
Date:   Sat Apr 18 11:52:39 2026 -0300

    speculative: fix sampler for checkpoints

commit e7006393bc
Author: SamuelOliveirads <samueloliveira32df@gmail.com>
Date:   Fri Apr 17 14:08:25 2026 -0300

    server: refactor checkpoint state logic

commit 57eabf04df
Merge: dc4797b7 64234e3c
Author: SamuelOliveirads <samueloliveira32df@gmail.com>
Date:   Fri Apr 17 13:53:41 2026 -0300

    Merge branch 'main' into fix/hybrid-cache-speculative

commit dc4797b723
Author: SamuelOliveirads <samueloliveira32df@gmail.com>
Date:   Fri Apr 17 13:12:40 2026 -0300

    reset ngram mod state for rejected tokens

commit 8ff2d943a3
Author: SamuelOliveirads <samueloliveira32df@gmail.com>
Date:   Fri Apr 17 13:08:04 2026 -0300

    server: snapshot recurrent state in tensor

commit d93dfb5e6b
Author: SamuelOliveirads <samueloliveira32df@gmail.com>
Date:   Thu Apr 16 22:36:37 2026 -0300

    fix: save/restore sampler state during speculative checkpoint

    When speculative decoding rejects draft tokens and restores the
    recurrent state checkpoint, the sampler (RNG, grammar, prev tokens)
    must also be restored to maintain consistency. Without this, the
    sampler state reflects the rejected draft tokens, leading to
    potential divergence.

    Uses common_sampler_clone() to snapshot the sampler before the
    speculative batch decode, and restores it on rejection.

commit d670cf85cd
Author: SamuelOliveirads <samueloliveira32df@gmail.com>
Date:   Thu Apr 16 21:53:52 2026 -0300

    server: spec checkpoints for recurrent models

* server: fix leak context between requests

* qwen3: allow mtp to run with split graph

* qwen3 mtp: selects rows before the ffn
This commit is contained in:
Samuel Oliveira Alves
2026-04-28 07:47:50 +02:00
committed by GitHub
parent d6f3e4e28f
commit 67e6346225
10 changed files with 401 additions and 131 deletions
+69 -31
View File
@@ -3265,11 +3265,14 @@ void server_context::apply_checkpoint(server_slot & slot) {
if (do_reset) {
if (has_recurrent) {
// Hybrid/recurrent: do NOT zero n_past. The prompt prefix is already in cache_tokens
// and update_slots() reprocesses from slot.n_past_prompt; dropping to 0 forces a full
// recompute on every turn and — combined with cached state — trips llama_decode ret=-3.
SLT_WRN(slot, "no usable hybrid/recurrent checkpoint; preserving slot state (n_past = %d, n_past_prompt = %d)\n",
(int)slot.n_past, (int)slot.n_past_prompt);
// Without a usable recurrent checkpoint, preserving prefix state leaks stale recurrent memory
// from prior requests into the current prompt. Force a full prompt re-processing fallback.
SLT_WRN(slot, "%s", "no usable hybrid/recurrent checkpoint; forcing full prompt re-processing\n");
slot.n_past = 0;
slot.n_past_prompt = 0;
slot.n_past_se = 0;
slot.ga_i = 0;
common_sampler_reset(slot.ctx_sampling);
} else {
SLT_WRN(slot, "forcing full prompt re-processing due to lack of cache data (likely due to SWA, see %s)\n",
"https://github.com/ggml-org/llama.cpp/pull/13194#issuecomment-2868343055");
@@ -3720,7 +3723,8 @@ void server_context::extend_context(const int32_t n_tokens) {
// Restore recurrent state and re-decode accepted tokens after speculative-decode rejection.
static void restore_speculative_checkpoint(
server_slot & slot, llama_context * ctx, llama_model * model,
const std::vector<llama_token> & ids, int n_draft) {
const std::vector<llama_token> & ids, int n_draft,
const std::vector<float> & mtp_hidden_state_pre, int32_t mtp_n_past_base) {
if (slot.spec_ckpt.per_step_enabled) {
const int step = (int)ids.size() - 1;
llama_spec_ckpt_restore(ctx, slot.id, slot.spec_ckpt.n_past, step);
@@ -3732,6 +3736,15 @@ static void restore_speculative_checkpoint(
common_sampler_accept(slot.ctx_sampling, ctx, id, true);
}
// Update MTP KV cache and hidden state using embeddings collected before checkpoint restore.
if (slot.has_mtp && !mtp_hidden_state_pre.empty()) {
slot.mtp_hidden_state = mtp_hidden_state_pre;
llama_context * mtp_ctx = common_speculative_get_mtp_ctx(slot.spec);
llama_context * mtp_target = mtp_ctx ? mtp_ctx : ctx;
llama_set_draft_input_hidden_state(mtp_target, slot.mtp_hidden_state.data());
mtp_accept_tokens(mtp_target, ids, mtp_n_past_base, slot.id);
}
SLT_DBG(slot, "per-step restore: step=%d (rejected %d drafts)\n",
step, (int)(n_draft - (ids.size() - 1)));
} else {
@@ -3752,6 +3765,9 @@ static void restore_speculative_checkpoint(
}
if (slot.has_mtp) {
for (int j = 0; j < re_batch.n_tokens; j++) {
re_batch.logits[j] = true;
}
llama_set_embeddings(ctx, true);
}
@@ -3759,15 +3775,29 @@ static void restore_speculative_checkpoint(
if (ret != 0) {
SLT_ERR(slot, "failed to re-decode accepted tokens after checkpoint restore: %d\n", ret);
}
if (slot.has_mtp) {
llama_set_embeddings(ctx, false);
const int n_embd = llama_model_n_embd(llama_get_model(ctx));
const float * emb = llama_get_embeddings_ith(ctx, -1);
if (emb) {
slot.mtp_hidden_state.resize(n_embd);
memcpy(slot.mtp_hidden_state.data(), emb, n_embd * sizeof(float));
const int n_accepted = (int)ids.size();
slot.mtp_hidden_state.resize(n_accepted * n_embd);
for (int j = 0; j < n_accepted; j++) {
const float * emb_j = llama_get_embeddings_ith(ctx, j);
if (emb_j) {
memcpy(slot.mtp_hidden_state.data() + j * n_embd, emb_j, n_embd * sizeof(float));
}
}
llama_context * mtp_ctx_rej = common_speculative_get_mtp_ctx(slot.spec);
llama_context * mtp_target_rej = mtp_ctx_rej ? mtp_ctx_rej : ctx;
llama_set_draft_input_hidden_state(mtp_target_rej, slot.mtp_hidden_state.data());
mtp_accept_tokens(mtp_target_rej, ids, slot.spec_ckpt.n_past, slot.id);
if (n_accepted > 1) {
memmove(slot.mtp_hidden_state.data(),
slot.mtp_hidden_state.data() + (n_accepted - 1) * n_embd,
n_embd * sizeof(float));
}
slot.mtp_hidden_state.resize(n_embd);
}
for (llama_token id : ids) {
@@ -3795,30 +3825,28 @@ void server_context::speculative_decoding_accept() {
// the accepted tokens from the speculation
const auto ids = common_sampler_sample_and_accept_n(slot.ctx_sampling, ctx, slot.i_batch_dft, slot.drafted);
int32_t mtp_n_past_base = 0;
std::vector<float> mtp_hidden_state_pre;
if (slot.has_mtp) {
llama_context * mtp_ctx = common_speculative_get_mtp_ctx(slot.spec);
llama_context * mtp_target = mtp_ctx ? mtp_ctx : ctx;
mtp_n_past_base = slot.n_past - (slot.drafted.size() + 1);
const int n_embd = llama_model_n_embd(llama_get_model(ctx));
if (!ids.empty()) {
const float* emb = llama_get_embeddings(ctx);
if (emb) {
slot.mtp_hidden_state.resize(ids.size() * n_embd);
memcpy(slot.mtp_hidden_state.data(), emb, ids.size() * n_embd * sizeof(float));
mtp_hidden_state_pre.resize(ids.size() * n_embd);
for (size_t i = 0; i < ids.size(); i++) {
const float* emb_i = llama_get_embeddings_ith(ctx, slot.i_batch_dft[i]);
if (emb_i) {
memcpy(mtp_hidden_state_pre.data() + i * n_embd, emb_i, n_embd * sizeof(float));
}
}
} else {
const float* emb0 = llama_get_embeddings_ith(ctx, 0);
if (emb0) {
slot.mtp_hidden_state.resize(n_embd);
memcpy(slot.mtp_hidden_state.data(), emb0, n_embd * sizeof(float));
mtp_hidden_state_pre.resize(n_embd);
memcpy(mtp_hidden_state_pre.data(), emb0, n_embd * sizeof(float));
}
}
llama_set_draft_input_hidden_state(mtp_target, slot.mtp_hidden_state.data());
int32_t n_past_base = slot.n_past - (slot.drafted.size() + 1);
mtp_accept_tokens(mtp_target, ids, n_past_base, slot.id);
}
slot.i_batch_dft.clear();
@@ -3846,8 +3874,16 @@ void server_context::speculative_decoding_accept() {
// for recurrent/hybrid models: if any drafts were rejected, restore recurrent state
const bool any_rejected = (ids.size() - 1) < n_draft;
if (any_rejected && slot.spec_ckpt.valid) {
restore_speculative_checkpoint(slot, ctx, model, ids, n_draft);
restore_speculative_checkpoint(slot, ctx, model, ids, n_draft, mtp_hidden_state_pre, mtp_n_past_base);
} else {
if (slot.has_mtp && !mtp_hidden_state_pre.empty()) {
llama_context * mtp_ctx = common_speculative_get_mtp_ctx(slot.spec);
llama_context * mtp_target = mtp_ctx ? mtp_ctx : ctx;
slot.mtp_hidden_state = std::move(mtp_hidden_state_pre);
llama_set_draft_input_hidden_state(mtp_target, slot.mtp_hidden_state.data());
mtp_accept_tokens(mtp_target, ids, mtp_n_past_base, slot.id);
}
llama_kv_cache_seq_rm(ctx, slot.id, slot.n_past, -1);
discard_speculative_checkpoint(slot, ctx);
}
@@ -4232,12 +4268,14 @@ void server_context::process_batch_tokens(int32_t & n_batch) {
}
}
if (mtp_warmup_needed) {
const float* emb = llama_get_embeddings(ctx);
const int n_embd = llama_model_n_embd(llama_get_model(ctx));
const int n_toks = batch_view.n_tokens;
if (emb) {
batch_mtp_hidden_state.resize(n_toks * n_embd);
memcpy(batch_mtp_hidden_state.data(), emb, n_toks * n_embd * sizeof(float));
batch_mtp_hidden_state.resize(n_toks * n_embd);
for (int t = 0; t < n_toks; t++) {
const float* emb_t = llama_get_embeddings_ith(ctx, t);
if (emb_t) {
memcpy(batch_mtp_hidden_state.data() + t * n_embd, emb_t, n_embd * sizeof(float));
}
}
}
}