Improve mtp acceptance rate (#1499)

* wip: port MTP architecture

Ports the Multi-Token Prediction (MTP) architecture to the older `llama.cpp` codebase used by `ikllama`.

Changes include:
- Updating `llama_batch` to support `mtp_params`.
- Modifying `llama_decode_internal` (and `encode`) to handle MTP operations (Warmup, Update, Draft).
- Adding public APIs for MTP state management (`llama_set_draft_input_hidden_state`).
- Adapting the embedding extraction logic to skip MTP update passes.

* Refactors `server_slot` to support generic speculative decoding (MTP or Draft Model).

* core: enable hybrid outputs (logits + embeddings) for MTP support

* fix(mtp): correct KV-cache slot finding for updates

* fix(mtp): persist hidden states to prevent context corruption during drafting

* refactor(mtp): clean unused code

* fix(mtp): update server to new functions name

* fix(mtp): fix graph and save hidden state

* mtp: refactor integration, context params and kv cache search

* mtp: fix hidden state extraction and speculative acceptance flow

* server: fix MTP warmup for long prompts and reset token buffer

* llama: refactor MTP operation state to context parameters

* server: fix n_past calculation in MTP acceptance

* llama: fix mtp enable flags

* speculative: refactor MTP to use common_speculative interface

* context: remove unused signatures

* clip: fix deprecated enum-enum conversion warning

* common: fix format string crash in help message

* context: fix mtp activation logic

* llamat: always use the extracted embedding

* llama: get all embeddings to kv cache

* llama: revert logit to not run mtp for not supported arch

* llama: allocate all the n_outputs for MTP

* wip

* server-context: get only the last embedding for hidden state

* ggml-backend: fix array of bounds in debug build

* server-context: run mt kv update to each prompt batch

* revert segmentation fault fixes

* glm-mtp(feat): optimize graph embedding and recursive drafting
This commit is contained in:
Samuel Oliveira Alves
2026-03-25 10:20:22 +01:00
committed by GitHub
parent 3c4a7d18bd
commit 1f3e832cb3
10 changed files with 173 additions and 66 deletions
+1 -1
View File
@@ -139,7 +139,7 @@ thinking_tokens thinking_tokens_from_string(const std::string& format);
enum common_speculative_type {
COMMON_SPECULATIVE_TYPE_NONE, // no speculative decoding
COMMON_SPECULATIVE_TYPE_DRAFT, // draft model
COMMON_SPECULATIVE_TYPE_MTP, // MTP model
COMMON_SPECULATIVE_TYPE_MTP, // MTP model
COMMON_SPECULATIVE_TYPE_EAGLE3, // eagle draft model
COMMON_SPECULATIVE_TYPE_NGRAM_SIMPLE, // simple self-speculative decoding
COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K, // self-speculative decoding with n-gram keys only
+26
View File
@@ -678,3 +678,29 @@ common_grammar_trigger common_grammar_trigger::from_json(const json& in) {
}
return out;
}
llama_token common_sampler_sample_speculative(struct common_sampler * gsmpl, struct llama_context * ctx, int idx, float * out_prob) {
GGML_UNUSED(gsmpl);
float * logits = llama_get_logits_ith(ctx, idx);
const int n_vocab = llama_n_vocab(llama_get_model(ctx));
int best_id = 0;
float max_val = logits[0];
for (int i = 1; i < n_vocab; ++i) {
if (logits[i] > max_val) {
max_val = logits[i];
best_id = i;
}
}
if (out_prob) {
double sum_exp = 0.0;
for (int i = 0; i < n_vocab; ++i) {
sum_exp += exp((double)(logits[i] - max_val));
}
*out_prob = (float)(1.0 / sum_exp);
}
return best_id;
}
+3
View File
@@ -227,5 +227,8 @@ std::vector<llama_token> llama_sampling_sample_and_accept_n(struct common_sample
std::vector<llama_token> common_sampler_sample_and_accept_n(struct common_sampler * gsmpl, struct llama_context * ctx, const std::vector<int> & idxs, const std::vector<llama_token> & draft, bool grammar_first = false);
// Greedy argmax sampling for speculative drafting
llama_token common_sampler_sample_speculative(struct common_sampler * gsmpl, struct llama_context * ctx, int idx, float * out_prob = nullptr);
llama_grammar* llama_sampler_init_llg(const llama_vocab* vocab,
const char* grammar_kind, const char* grammar_data);
+9 -14
View File
@@ -1146,27 +1146,22 @@ std::vector<llama_token> mtp_speculative_gen_draft(
break;
}
common_sampler_sample(smpl, ctx, 0, true);
float prob;
llama_token id_next = common_sampler_sample_speculative(smpl, ctx, 0, &prob);
const auto * cur_p = common_sampler_get_candidates(smpl, true);
if (!cur_p || cur_p->size == 0) {
break;
drafts.push_back(id_next);
const float * emb = llama_get_embeddings_ith(ctx, 0);
if (emb) {
llama_set_draft_input_hidden_state(ctx, emb);
}
const llama_token id_next = cur_p->data[0].id;
const float prob = cur_p->data[0].p;
common_sampler_accept(smpl, nullptr, id_next, true);
current_input_id = id_next;
current_n_past++;
if (prob < p_min) {
break;
}
drafts.push_back(id_next);
current_input_id = id_next;
current_n_past++;
}
llama_batch_free(mtp_batch);
llama_set_mtp_op_type(ctx, MTP_OP_NONE);
+73 -31
View File
@@ -209,9 +209,8 @@ void server_context::init() {
if (params_base.has_mtp) {
if (llama_model_n_nextn_layer(model) > 0) {
SRV_INF("%s\n", "MTP detected, configuring for speculative decoding...");
params_base.speculative.type = COMMON_SPECULATIVE_TYPE_MTP;
params_base.pooling_type = LLAMA_POOLING_TYPE_NONE;
slot.has_mtp = true;
slot.params.speculative.type = COMMON_SPECULATIVE_TYPE_MTP;
@@ -404,6 +403,10 @@ void server_slot::reset() {
task.reset();
}
bool server_slot::need_embd() const {
return embedding || has_mtp;
}
bool server_slot::has_budget(gpt_params& global_params) {
if (params.n_predict == -1 && global_params.n_predict == -1) {
return true; // limitless
@@ -1737,6 +1740,9 @@ bool server_context::check_no_mtmd(const int id_task) {
}
void server_context::send_partial_response(server_slot& slot, completion_token_output tkn) {
if (slot.task == nullptr) {
return;
}
auto res = std::make_unique<server_task_result_cmpl_partial>();
res->final_result = false;
res->id = slot.id_task;
@@ -2719,10 +2725,18 @@ void server_context::add_sampled_tokens() {
if (slot.has_mtp) {
if (!slot.mtp_hidden_state.empty()) {
llama_set_draft_input_hidden_state(ctx, slot.mtp_hidden_state.data());
const int n_embd = llama_model_n_embd(llama_get_model(ctx));
const int n_hidden = slot.mtp_hidden_state.size() / n_embd;
llama_set_draft_input_hidden_state(ctx, slot.mtp_hidden_state.data() + (n_hidden - 1) * n_embd);
} else {
LOG_ERROR("MTP hidden state is empty during speculation", {});
llama_set_draft_input_hidden_state(ctx, llama_get_embeddings_ith(ctx, -1));
const float* emb_neg1 = llama_get_embeddings_ith(ctx, -1);
if (emb_neg1) {
const int n_embd = llama_model_n_embd(llama_get_model(ctx));
slot.mtp_hidden_state.resize(n_embd);
memcpy(slot.mtp_hidden_state.data(), emb_neg1, n_embd * sizeof(float));
llama_set_draft_input_hidden_state(ctx, slot.mtp_hidden_state.data());
}
}
}
@@ -3177,7 +3191,7 @@ void server_context::batch_pending_prompt(const int32_t n_ubatch, const int32_t
}
int p0 = system_tokens.size() + slot.cache_tokens.pos_next();
common_batch_add(batch, cur_tok, p0, { slot.id }, slot.embedding);
common_batch_add(batch, cur_tok, p0, { slot.id }, slot.need_embd());
slot.cache_tokens.push_back(cur_tok);
@@ -3280,17 +3294,21 @@ void server_context::speculative_decoding_accept() {
const auto ids = common_sampler_sample_and_accept_n(slot.ctx_sampling, ctx, slot.i_batch_dft, slot.drafted);
if (slot.has_mtp) {
const int n_embd = llama_model_n_embd(llama_get_model(ctx));
const int n_embd = llama_model_n_embd(llama_get_model(ctx));
if (!ids.empty()) {
const float* emb = llama_get_embeddings_ith(ctx, ids.size() - 1);
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));
}
} 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(), emb, n_embd * sizeof(float));
memcpy(slot.mtp_hidden_state.data(), emb0, n_embd * sizeof(float));
}
}
else {
llama_set_draft_input_hidden_state(ctx, llama_get_embeddings_ith(ctx, 0));
}
llama_set_draft_input_hidden_state(ctx, slot.mtp_hidden_state.data());
int32_t n_past_base = slot.n_past - (slot.drafted.size() + 1);
@@ -3343,6 +3361,9 @@ void server_context::speculative_decoding_accept() {
}
} else {
buffer_and_check_string_ban(slot, result);
if (slot.task == nullptr) {
break;
}
}
common_sampler_review(slot.ctx_sampling, slot.token_buffer.size(), slot.rewind_status);
@@ -3669,8 +3690,34 @@ void server_context::process_batch_tokens(int32_t & n_batch) {
continue; // continue loop of n_batch
}
bool mtp_warmup_needed = false;
std::vector<float> batch_mtp_hidden_state;
if (params_base.has_mtp) {
for (auto& slot : slots) {
if ((slot.state == SLOT_STATE_PROCESSING && slot.n_decoded == 0) ||
(slot.state == SLOT_STATE_IDLE && slot.command == SLOT_COMMAND_LOAD_PROMPT)) {
bool has_tokens_for_slot = (batch_view.n_tokens > 0 && batch_view.n_seq_id[0] > 0 && batch_view.seq_id[0][0] == slot.id);
if (has_tokens_for_slot) {
mtp_warmup_needed = true;
break;
}
}
}
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));
}
}
}
for (auto& slot : slots) {
if (slot.state != SLOT_STATE_PROCESSING || slot.i_batch < (int)i || slot.i_batch >= (int)(i + n_tokens)) {
bool is_active_slot = (slot.state == SLOT_STATE_PROCESSING);
if (!is_active_slot || slot.i_batch < (int)i || slot.i_batch >= (int)(i + n_tokens)) {
// save checkpoint during prompt processing
if (slot.command == SLOT_COMMAND_LOAD_PROMPT) {
if (slot.do_checkpoint) {
@@ -3709,21 +3756,20 @@ void server_context::process_batch_tokens(int32_t & n_batch) {
completion_token_output result;
const int tok_idx = slot.i_batch - i;
if (params_base.has_mtp && slot.n_decoded == 0) {
const float* emb_i = llama_get_embeddings_ith(ctx, tok_idx);
if (emb_i) {
const int n_embd = llama_model_n_embd(llama_get_model(ctx));
slot.mtp_hidden_state.resize(n_embd);
memcpy(slot.mtp_hidden_state.data(), emb_i, n_embd * sizeof(float));
}
}
const llama_token id = common_sampler_sample(slot.ctx_sampling, ctx, tok_idx);
common_sampler_accept(slot.ctx_sampling, ctx, id, true);
if (params_base.has_mtp && slot.n_decoded == 0) {
if (batch_view.n_seq_id[0] > 0 && batch_view.seq_id[0][0] == slot.id) {
mtp_update_kv_cache(ctx, batch_view, true);
const float* emb = llama_get_embeddings_ith(ctx, -1);
if (emb) {
const int n_embd = llama_model_n_embd(llama_get_model(ctx));
slot.mtp_hidden_state.resize(n_embd);
memcpy(slot.mtp_hidden_state.data(), emb, n_embd * sizeof(float));
}
}
}
slot.n_decoded += 1;
const int64_t t_current = ggml_time_us();
@@ -3765,15 +3811,11 @@ void server_context::process_batch_tokens(int32_t & n_batch) {
slot.i_batch = -1;
}
if (params_base.has_mtp) {
for (auto& slot : slots) {
if (slot.n_past < slot.n_prompt_tokens) {
if (batch_view.n_seq_id[0] > 0 && batch_view.seq_id[0][0] == slot.id) {
mtp_update_kv_cache(ctx, batch_view, true);
}
}
}
if (mtp_warmup_needed && !batch_mtp_hidden_state.empty()) {
llama_set_draft_input_hidden_state(ctx, batch_mtp_hidden_state.data());
mtp_update_kv_cache(ctx, batch_view, true);
}
// speculative decoding - main model sample and accept
speculative_decoding_accept();
}
+2
View File
@@ -165,6 +165,8 @@ struct server_slot {
void reset();
bool need_embd() const;
bool has_budget(gpt_params& global_params);
bool available() const;
+2 -1
View File
@@ -7369,7 +7369,7 @@ ggml_cgraph * llm_build_context::build_glm4_moe() {
struct ggml_tensor * KQ_mask = build_inp_KQ_mask();
// output token IDs (for last layer cropping)
struct ggml_tensor * inp_out_ids = n_tokens > 1 ? build_inp_out_ids() : nullptr;
struct ggml_tensor * inp_out_ids = (n_tokens > 1 && !lctx.cparams.mtp) ? build_inp_out_ids() : nullptr;
float kq_scale = 1.0f/sqrtf(float(n_embd_head));
@@ -7592,6 +7592,7 @@ struct ggml_tensor * llm_build_context::build_mtp_tail(
cb(cur, "mtp_ffn_out_resid", il);
}
cur = llm_build_norm(ctx0, cur, hparams, mtp_layer.nextn.shared_head_norm, NULL, LLM_NORM_RMS, cb, il);
cb(cur, "result_norm", -1);
if (inp_out_ids) {
cur = ggml_get_rows(ctx0, cur, inp_out_ids);
+1
View File
@@ -437,6 +437,7 @@ llm_expert_gating_func_type gating_op,
bool is_multi = false);
static uint32_t llama_kv_qnext_state_slots(const llama_kv_cache & kv_self);
struct ggml_tensor * build_mtp_tail(
const struct llama_layer & mtp_layer,
struct ggml_tensor * prev_embeddings,
+5
View File
@@ -47,6 +47,9 @@ struct llama_kv_cache {
uint32_t size = 0;
uint32_t used = 0; // used cells (i.e. at least one seq_id)
// Track's main model's head position for MTP KV cache operations
uint32_t mtp_kv_head_hint = 0;
// computed before each graph build
uint32_t n = 0;
@@ -229,7 +232,9 @@ struct llama_context {
std::vector<CacheCopy> cache_copies;
bool update_cache_copies();
bool prepare_mtp_graph_inputs(
struct llama_context & lctx);
void set_mtp_op_type(llama_mtp_op_type value);
};
+51 -19
View File
@@ -1093,7 +1093,13 @@ static bool llama_kv_cache_find_slot(
bool found = false;
if (cache.head < cache.size &&
if (cache.mtp_kv_head_hint < cache.size &&
cache.cells[cache.mtp_kv_head_hint].pos == target_pos &&
cache.cells[cache.mtp_kv_head_hint].has_seq_id(target_seq)) {
cache.head = cache.mtp_kv_head_hint;
found = true;
}
else if (cache.head < cache.size &&
cache.cells[cache.head].pos == target_pos &&
cache.cells[cache.head].has_seq_id(target_seq)) {
found = true;
@@ -3039,7 +3045,7 @@ static void llama_set_inputs(llama_context & lctx, const llama_batch & batch) {
auto tim1 = ggml_time_us();
#endif
const int64_t n_tokens = batch.n_tokens;
if (n_tokens > 1) {
if (n_tokens > 1 && !cparams.mtp) {
GGML_ASSERT(lctx.inp_out_ids && "every model that can must skip unused outputs");
}
@@ -3643,12 +3649,7 @@ static void llama_graph_compute(
static bool prepare_mtp_graph_inputs(struct llama_context & lctx) {
ggml_tensor * dst = lctx.inp_mtp_states;
const float * src = nullptr;
if (lctx.cparams.mtp_op_type == MTP_OP_WARMUP || lctx.cparams.mtp_op_type == MTP_OP_UPDATE_ACCEPTED) {
src = lctx.embd;
} else {
src = lctx.draft_input_hidden_state;
}
const float * src = lctx.draft_input_hidden_state;
if (!src) {
LLAMA_LOG_ERROR("%s: Source hidden state is null\n", __func__);
@@ -3705,6 +3706,8 @@ static int llama_decode_internal(
uint32_t n_outputs = 0;
uint32_t n_outputs_prev = 0;
uint32_t n_outputs_embd = 0;
uint32_t n_outputs_prev_embd = 0;
const auto n_ubatch = cparams.n_ubatch;
@@ -3716,6 +3719,7 @@ static int llama_decode_internal(
// this indicates we are doing pooled embedding, so we ignore batch.logits and output all tokens
const bool embd_pooled = cparams.embeddings && cparams.pooling_type != LLAMA_POOLING_TYPE_NONE;
const bool has_mtp = cparams.mtp && hparams.nextn_predict_layers > 0;
// count outputs
if (batch_all.logits && !embd_pooled) {
@@ -3730,8 +3734,9 @@ static int llama_decode_internal(
}
// reserve output buffer
if (llama_output_reserve(lctx, n_outputs) < n_outputs) {
LLAMA_LOG_ERROR("%s: could not reserve space for batch with %u outputs\n", __func__, n_outputs);
n_outputs_embd = has_mtp ? n_tokens_all : n_outputs;
if (llama_output_reserve(lctx, std::max<size_t>(n_outputs, n_outputs_embd)) < std::max<size_t>(n_outputs, n_outputs_embd)) {
LLAMA_LOG_ERROR("%s: could not reserve space for batch with %zu outputs\n", __func__, std::max<size_t>(n_outputs, n_outputs_embd));
return -2;
};
@@ -3741,10 +3746,17 @@ static int llama_decode_internal(
for (uint32_t i = 0; i < n_tokens_all; ++i) {
if (batch_all.logits[i]) {
lctx.output_ids[i] = i_logits++;
} else {
lctx.output_ids[i] = -1;
}
}
} else if (n_outputs == 1 && n_tokens_all > 0) {
for (uint32_t i = 0; i < n_tokens_all; ++i) {
lctx.output_ids[i] = -1;
}
lctx.output_ids[n_tokens_all - 1] = 0;
} else {
for (uint32_t i = 0; i < n_outputs; ++i) {
for (uint32_t i = 0; i < std::max<uint32_t>(n_outputs, n_outputs_embd); ++i) {
lctx.output_ids[i] = i;
}
}
@@ -3874,6 +3886,10 @@ static int llama_decode_internal(
return 1;
}
if (cparams.mtp_op_type == MTP_OP_NONE) {
kv_self.mtp_kv_head_hint = kv_self.head;
}
if (!kv_self.recurrent) {
// a heuristic, to avoid attending the full cache if it is not yet utilized
// after enough generations, the benefit from this heuristic disappears
@@ -4032,7 +4048,22 @@ static int llama_decode_internal(
if (n_outputs_new) {
GGML_ASSERT( n_outputs_prev + n_outputs_new <= n_outputs);
GGML_ASSERT((n_outputs_prev + n_outputs_new)*n_vocab <= (int64_t) lctx.logits_size);
ggml_backend_tensor_get_async(backend_res, res, logits_out, 0, n_outputs_new*n_vocab*sizeof(float));
if (res->ne[1] == n_tokens && n_outputs_new < n_tokens) {
int32_t i_out = 0;
if (u_batch.logits && !embd_pooled) {
for (uint32_t i = 0; i < n_tokens; i++) {
if (u_batch.logits[i]) {
ggml_backend_tensor_get_async(backend_res, res, logits_out + i_out*n_vocab, i*n_vocab*sizeof(float), n_vocab*sizeof(float));
i_out++;
}
}
} else if (cur_token + n_tokens >= n_tokens_all) {
ggml_backend_tensor_get_async(backend_res, res, logits_out, (n_tokens - 1)*n_vocab*sizeof(float), n_vocab*sizeof(float));
}
} else {
ggml_backend_tensor_get_async(backend_res, res, logits_out, 0, n_outputs_new*n_vocab*sizeof(float));
}
}
}
#if IK_PRINT_TIMING
@@ -4042,7 +4073,7 @@ static int llama_decode_internal(
}
// extract embeddings
if (embd && cparams.mtp_op_type == MTP_OP_NONE) {
if (embd && (cparams.mtp_op_type == MTP_OP_NONE || cparams.mtp_op_type == MTP_OP_DRAFT_GEN)) {
#if IK_PRINT_TIMING
tim1 = ggml_time_us();
#endif
@@ -4054,13 +4085,13 @@ static int llama_decode_internal(
{
// extract token embeddings
GGML_ASSERT(lctx.embd != nullptr);
float * embd_out = lctx.embd + n_outputs_prev*n_embd;
const int32_t n_outputs_new = lctx.n_outputs;
float * embd_out = lctx.embd + n_outputs_prev_embd*n_embd;
const int32_t n_outputs_new_embd = has_mtp ? n_tokens : lctx.n_outputs;
if (n_outputs_new) {
GGML_ASSERT( n_outputs_prev + n_outputs_new <= n_outputs);
GGML_ASSERT((n_outputs_prev + n_outputs_new)*n_embd <= (int64_t) lctx.embd_size);
ggml_backend_tensor_get_async(backend_embd, embd, embd_out, 0, n_outputs_new*n_embd*sizeof(float));
if (n_outputs_new_embd) {
GGML_ASSERT( n_outputs_prev_embd + n_outputs_new_embd <= n_outputs_embd);
GGML_ASSERT((n_outputs_prev_embd + n_outputs_new_embd)*n_embd <= (int64_t) lctx.embd_size);
ggml_backend_tensor_get_async(backend_embd, embd, embd_out, 0, n_outputs_new_embd*n_embd*sizeof(float));
}
} break;
case LLAMA_POOLING_TYPE_MEAN:
@@ -4091,6 +4122,7 @@ static int llama_decode_internal(
#endif
}
n_outputs_prev += lctx.n_outputs;
n_outputs_prev_embd += has_mtp ? n_tokens : lctx.n_outputs;
cur_token += n_tokens;
if (reset_previous) {
// We need to discard this graph. Otherwise, iwith CUDA graphs enabled, the graph will get resused and this will reset the