Commit 781dbc5ac for llama.cpp
commit 781dbc5ac98921dbdb5e5b2ec5b7a50960e937d4
Author: Xuan-Son Nguyen <son@huggingface.co>
Date: Sat Oct 10 11:22:38 2026 +0200
spec: properly handle mtmd input for mtp (#30257)
* spec: properly handle mtmd input for mtp
* nits
diff --git a/common/common.cpp b/common/common.cpp
index 28ab9ac6c..9a71a33ff 100644
--- a/common/common.cpp
+++ b/common/common.cpp
@@ -1403,6 +1403,14 @@ std::vector<llama_adapter_lora_ptr> & common_init_result::lora() {
return pimpl->lora;
}
+// only for warmup and probe decodes, fill zeros as dummy input
+static void common_batch_set_zero_state(common_batch & batch, const llama_model * model, std::vector<float> & zeros) {
+ zeros.assign(llama_model_n_embd_out(model), 0.0f);
+ for (int32_t i = 0; i < batch.size(); ++i) {
+ batch.set_embd_state(i, { zeros.data(), 1, zeros.size() });
+ }
+}
+
common_init_result_ptr common_init_from_params(common_params & params, bool model_only) {
common_init_result_ptr res(new common_init_result(params, model_only));
@@ -1509,6 +1517,8 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
if (llama_model_has_decoder(model)) {
tmp.resize(std::min(tmp.size(), (size_t) params.n_batch));
common_batch batch = common_batch_get_one(lctx, tmp);
+ std::vector<float> zeros;
+ common_batch_set_zero_state(batch, model, zeros);
llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
}
llama_memory_clear(llama_get_memory(lctx), true);
@@ -1576,6 +1586,8 @@ common_context_seq_rm_type common_context_can_seq_rm(llama_context * ctx) {
int ret;
{
common_batch batch = common_batch_get_one(ctx, tmp);
+ std::vector<float> zeros;
+ common_batch_set_zero_state(batch, llama_get_model(ctx), zeros);
ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
}
if (ret != 0) {
@@ -2161,7 +2173,7 @@ void common_batch::clear() {
}
int32_t common_batch::add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output) {
- tokens.push_back({ id, { pos, 0, 0, 0 }, seq_id, output, { nullptr, 0, 0 }, {} });
+ tokens.push_back({ id, { pos, 0, 0, 0 }, seq_id, output, { nullptr, 0, 0 }, { nullptr, 0, 0 }, {} });
return size() - 1;
}
@@ -2199,8 +2211,16 @@ bool common_batch::set_embd(int32_t idx, llama_embd embd) {
return true;
}
+bool common_batch::set_embd_state(int32_t idx, llama_embd state) {
+ if (idx < 0 || idx >= size() || tokens[idx].state.data != nullptr) {
+ return false;
+ }
+ tokens[idx].state = state;
+ return true;
+}
+
int32_t common_batch::add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output) {
- token t = { LLAMA_TOKEN_NULL, { 0, 0, 0, 0 }, seq_id, output, embd, {} };
+ token t = { LLAMA_TOKEN_NULL, { 0, 0, 0, 0 }, seq_id, output, embd, { nullptr, 0, 0 }, {} };
for (int32_t j = 0; j < n_pos; ++j) {
t.pos[j] = pos[j];
}
@@ -2245,6 +2265,9 @@ llama_batch_ext * common_batch::get_sub_batch(int32_t off, int32_t n) {
if (t.output) {
llama_batch_ext_set_output_logits(res, idx, true);
}
+ if (t.state.data) {
+ llama_batch_ext_set_embd_state(res, idx, t.state); // contexts without a state input ignore it
+ }
if (t.decision_order != 0) {
llama_batch_ext_set_decision_order(res, idx, (llama_decision_order) t.decision_order);
}
diff --git a/common/common.h b/common/common.h
index d2fe10fd7..d89f803af 100644
--- a/common/common.h
+++ b/common/common.h
@@ -1074,6 +1074,7 @@ struct common_batch {
llama_seq_id seq_id; // the first sequence id, see add_seq()
bool output;
llama_embd embd; // non-owning view of the data passed to add_embd()/set_embd(), data == NULL if none
+ llama_embd state; // non-owning view of the data passed to set_embd_state(), data == NULL if none
std::vector<llama_seq_id> seq_ids_extra; // see add_seq()
int32_t decision_order = 0; // see llama_batch_ext_set_decision_order()
};
@@ -1111,6 +1112,9 @@ struct common_batch {
// attach a token embedding to the entry at idx, can only be set once per entry
bool set_embd(int32_t idx, llama_embd embd);
+ // attach a state embedding (e.g. the target hidden state for MTP) to the entry at idx, can only be set once per entry
+ bool set_embd_state(int32_t idx, llama_embd state);
+
// add an embedding-only entry (no token id)
// pos points to n_pos positions
int32_t add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output);
diff --git a/common/speculative.cpp b/common/speculative.cpp
index d9ddf44d2..c5043a4a8 100644
--- a/common/speculative.cpp
+++ b/common/speculative.cpp
@@ -1541,8 +1541,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
return true;
}
- // TODO: how to make it work with vision tokens?
- if (!batch_in.has_token() || batch_in.has_embd()) {
+ if (!batch_in.has_token() && !batch_in.has_embd()) {
return true;
}
@@ -1581,15 +1580,20 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
const float * h_tgt = llama_get_embeddings_nextn(ctx_tgt);
for (int k = 0; k < n_tokens; ++k) {
- const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
+ const auto & t = batch_in.tokens[k];
+
+ const llama_seq_id seq_id = t.seq_id;
- const int32_t idx = batch.add(batch_in.tokens[k].id, batch_in.tokens[k].pos[0], seq_id, false);
+ // vision tokens carry an embedding instead of an id
+ const int32_t idx = t.id != LLAMA_TOKEN_NULL
+ ? batch.add(t.id, t.pos[0], seq_id, false)
+ : batch.add_embd(t.embd, t.pos.data(), seq_id, false);
const float * h_row = k == i_batch_beg[seq_id]
? pending_h[seq_id].data()
: h_tgt + (size_t) (k - 1) * n_embd;
- batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
+ batch.set_embd_state(idx, { h_row, 1, (size_t) n_embd });
}
auto * mem_dft = llama_get_memory(ctx_dft);
@@ -1679,7 +1683,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
}
const int32_t idx = batch.add(dp.id_last, dp.pos0, seq_id, true);
- batch.set_embd(idx, { pending_h[seq_id].data(), 1, (size_t) n_embd });
+ batch.set_embd_state(idx, { pending_h[seq_id].data(), 1, (size_t) n_embd });
i_last[seq_id] = idx;
@@ -1772,18 +1776,18 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
for (int t = 0; t < n_rows; ++t) {
const llama_token tok = (t == 0) ? dp.id_last : result[t - 1];
const int32_t idx = batch.add(tok, dp.pos0 + t, seq_id, t == n_rows - 1);
- batch.set_embd(idx, { chain_h[seq_id].data() + (size_t) t * n_embd, 1, (size_t) n_embd });
+ batch.set_embd_state(idx, { chain_h[seq_id].data() + (size_t) t * n_embd, 1, (size_t) n_embd });
i_last[seq_id] = idx;
}
} else if (is_mem_shared) {
// note: with shared memory (e.g. Gemma4 assistants) we use the same position for all draft tokens
// ref: https://github.com/huggingface/transformers/blob/effde20942e3f82a1b97449f60b3a48c5ff96145/docs/source/en/model_doc/gemma4_assistant.md?plain=1#L36-L37
const int32_t idx = batch.add(id, dp.pos0, seq_id, true);
- batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
+ batch.set_embd_state(idx, { h_row, 1, (size_t) n_embd });
i_last[seq_id] = idx;
} else {
const int32_t idx = batch.add(id, dp.pos0 + i + 1, seq_id, true);
- batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
+ batch.set_embd_state(idx, { h_row, 1, (size_t) n_embd });
i_last[seq_id] = idx;
}
}
diff --git a/include/llama.h b/include/llama.h
index 60329024f..cfc69cee2 100644
--- a/include/llama.h
+++ b/include/llama.h
@@ -1056,6 +1056,7 @@ extern "C" {
// "state" here means extra hidden state carried over from a previous stage, e.g.:
// - MTP: state from N layers of the target model
// - Qwen3 VL (deepstack): state from N layers of the vision encoder
+ // Returns false if the context does not take a state embedding (currently only MTP contexts do)
LLAMA_API bool llama_batch_ext_set_embd_state(
struct llama_batch_ext * batch,
int32_t idx,
diff --git a/src/llama-batch.cpp b/src/llama-batch.cpp
index ecd48dd80..3c4b52008 100644
--- a/src/llama-batch.cpp
+++ b/src/llama-batch.cpp
@@ -31,9 +31,10 @@ bool llama_batch_allocr::init(
bool output_all) {
clear();
- this->vocab = &vocab;
- this->n_embd = batch_inp.n_embd > 0 ? batch_inp.n_embd : batch_inp.n_embd_inp;
- this->n_seq_max = batch_inp.n_seq_max;
+ this->vocab = &vocab;
+ this->n_embd = batch_inp.n_embd > 0 ? batch_inp.n_embd : batch_inp.n_embd_inp;
+ this->n_embd_state = batch_inp.n_embd_state;
+ this->n_seq_max = batch_inp.n_seq_max;
const int32_t n_tok = (int32_t) batch_inp.tokens.size();
@@ -48,14 +49,17 @@ bool llama_batch_allocr::init(
//
// determine the content types of the batch
- // an entry can carry a token id, a token embedding, or both (e.g. MTP hook batches)
+ // an entry can carry a token id, a token embedding, or both
// all entries must carry the same combination, or be a mix of token and embd entries
+ // a state embedding (e.g. MTP hook batches) is set on all entries or on none
//
int32_t n_tok_only = 0;
int32_t n_embd_only = 0;
int32_t n_both = 0;
+ const bool has_state = batch_inp.tokens[0].has_state;
+
for (int32_t i = 0; i < n_tok; ++i) {
const bool is_tok = batch_inp.tokens[i].id != LLAMA_TOKEN_NULL;
const bool is_emb = batch_inp.tokens[i].has_embd;
@@ -65,6 +69,11 @@ bool llama_batch_allocr::init(
return false;
}
+ if (batch_inp.tokens[i].has_state != has_state) {
+ LLAMA_LOG_ERROR("%s: all entries in the batch must have the same state embedding presence\n", __func__);
+ return false;
+ }
+
n_tok_only += is_tok && !is_emb;
n_embd_only += is_emb && !is_tok;
n_both += is_tok && is_emb;
@@ -124,6 +133,10 @@ bool llama_batch_allocr::init(
embd_vec = batch_inp.embd;
}
+ if (has_state) {
+ state_vec = batch_inp.state;
+ }
+
//
// build flat pos array, section-major: pos[j*n_tok + i] = section j of entry i
// token entry: [p, p, p, 0] (M-RoPE text position)
@@ -292,6 +305,7 @@ bool llama_batch_allocr::init(
/*.n_pos =*/ n_pos_per_embd,
/*.token =*/ batch.token,
/*.embd =*/ batch.embd,
+ /*.embd_state =*/ state_vec.empty() ? nullptr : state_vec.data(),
/*.pos =*/ batch.pos,
/*.n_seq_id =*/ batch.n_seq_id,
/*.seq_id =*/ batch.seq_id,
@@ -493,6 +507,7 @@ llama_ubatch llama_batch_allocr::ubatch_reserve(uint32_t n_seq_tokens, uint32_t
udata->token .resize(n_tokens);
udata->embd .clear();
+ udata->embd_state.clear();
udata->pos .resize(n_pos_all);
udata->n_seq_id .resize(n_tokens);
udata->seq_id .resize(n_tokens);
@@ -515,6 +530,7 @@ llama_ubatch llama_batch_allocr::ubatch_reserve(uint32_t n_seq_tokens, uint32_t
/*.token =*/ udata->token.data(),
/*.embd =*/ nullptr,
+ /*.embd_state =*/ nullptr,
/*.pos =*/ udata->pos.data(),
/*.n_seq_id =*/ udata->n_seq_id.data(),
/*.seq_id =*/ udata->seq_id.data(),
@@ -821,6 +837,7 @@ void llama_batch_allocr::clear() {
token_vec .clear();
embd_vec .clear();
is_embd_vec .clear();
+ state_vec .clear();
seq_id_data .clear();
pos .clear();
n_seq_id .clear();
@@ -863,12 +880,15 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u
const bool mixed = mixed_batch && n_embd_rows > 0 && n_embd_rows < n_tokens;
const bool use_token = batch.token && !(mixed_batch && n_embd_rows == n_tokens);
const bool use_embd = batch.embd && !(mixed_batch && n_embd_rows == 0);
+ const bool has_state = !state_vec.empty();
- const int64_t n_embd_all = use_embd ? (int64_t) n_tokens*n_embd : 0;
- const int64_t n_pos_all = (int64_t) n_tokens*n_pos_per_embd;
+ const int64_t n_embd_all = use_embd ? (int64_t) n_tokens*n_embd : 0;
+ const int64_t n_state_all = has_state ? (int64_t) n_tokens*n_embd_state : 0;
+ const int64_t n_pos_all = (int64_t) n_tokens*n_pos_per_embd;
udata->token .resize(n_tokens);
udata->embd .resize(n_embd_all);
+ udata->embd_state.resize(n_state_all);
udata->pos .resize(n_pos_all);
udata->n_seq_id .resize(n_tokens);
udata->seq_id .resize(n_tokens);
@@ -896,6 +916,10 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u
udata->type[i] = is_embd_vec[idxs[i]];
}
+ if (has_state) {
+ memcpy(udata->embd_state.data() + i*n_embd_state, state_vec.data() + (int64_t) idxs[i]*n_embd_state, n_embd_state*sizeof(float));
+ }
+
for (size_t j = 0; j < (size_t)n_pos_per_embd; ++j) {
udata->pos[j*n_tokens + i] = batch.pos[j*batch.n_tokens + idxs[i]];
}
@@ -942,6 +966,7 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u
/*.token =*/ use_token ? udata->token.data() : nullptr,
/*.embd =*/ use_embd ? udata->embd.data() : nullptr,
+ /*.embd_state =*/ has_state ? udata->embd_state.data() : nullptr,
/*.pos =*/ udata->pos.data(),
/*.n_seq_id =*/ udata->n_seq_id.data(),
/*.seq_id =*/ udata->seq_id.data(),
@@ -993,6 +1018,7 @@ void llama_batch_allocr::ubatch_print(const llama_ubatch & ubatch, int debug) {
LLAMA_LOG_DEBUG("%s: token = %p\n", __func__, (void *) ubatch.token);
LLAMA_LOG_DEBUG("%s: embd = %p\n", __func__, (void *) ubatch.embd);
+ LLAMA_LOG_DEBUG("%s: embd_state = %p\n", __func__, (void *) ubatch.embd_state);
LLAMA_LOG_DEBUG("%s: pos = %p\n", __func__, (void *) ubatch.pos);
LLAMA_LOG_DEBUG("%s: n_seq_id = %p\n", __func__, (void *) ubatch.n_seq_id);
LLAMA_LOG_DEBUG("%s: seq_id = %p\n", __func__, (void *) ubatch.seq_id);
@@ -1110,19 +1136,25 @@ void llama_batch_free(struct llama_batch batch) {
// llama_batch_ext
size_t llama_batch_ext_select_n_embd_inp(llama_context_type ctx_type, llm_arch arch, const llama_hparams & hparams) {
- if (ctx_type == LLAMA_CONTEXT_TYPE_MTP) {
- return hparams.n_embd_out();
- }
+ GGML_UNUSED(ctx_type);
if (arch == LLM_ARCH_DFLASH) {
return hparams.n_embd_inp_enc();
}
return hparams.n_embd_inp();
}
+size_t llama_batch_ext_select_n_embd_state(llama_context_type ctx_type, const llama_hparams & hparams) {
+ if (ctx_type == LLAMA_CONTEXT_TYPE_MTP) {
+ return hparams.n_embd_out();
+ }
+ return 0;
+}
+
llama_batch_ext::llama_batch_ext(llama_context * ctx) :
n_tokens_max(llama_n_batch(ctx)),
n_embd_inp(llama_batch_ext_select_n_embd_inp(ctx->get_cparams().ctx_type, llama_get_model(ctx)->arch, llama_get_model(ctx)->hparams)),
n_embd_inp_enc(llama_get_model(ctx)->hparams.n_embd_inp_enc()),
+ n_embd_state(llama_batch_ext_select_n_embd_state(ctx->get_cparams().ctx_type, llama_get_model(ctx)->hparams)),
n_seq_max(llama_n_seq_max(ctx)),
mem(llama_get_memory(ctx)),
n_vocab(llama_vocab_n_tokens(llama_model_get_vocab(llama_get_model(ctx)))),
@@ -1141,6 +1173,7 @@ llama_batch_ext::llama_batch_ext(
n_tokens_max(n_tokens_max),
n_embd_inp(n_embd_inp),
n_embd_inp_enc(n_embd_inp_enc),
+ n_embd_state(0),
n_seq_max(n_seq_max),
mem(mem),
n_vocab(n_vocab),
@@ -1151,6 +1184,7 @@ llama_batch_ext::llama_batch_ext(
void llama_batch_ext::clear() {
tokens.clear();
embd .clear();
+ state .clear();
n_embd = 0;
}
@@ -1233,6 +1267,38 @@ bool llama_batch_ext::set_token_embd(int32_t idx, llama_embd embd_in) {
return true;
}
+bool llama_batch_ext::set_token_state(int32_t idx, llama_embd state_in) {
+ if (idx < 0 || idx >= (int32_t) tokens.size()) {
+ return false;
+ }
+ if (!state_in.data) {
+ return false;
+ }
+ if (n_embd_state == 0) {
+ return false; // this context does not take state embeddings
+ }
+
+ const size_t n_total = state_in.n_rows * state_in.n_embd;
+ if (n_total != n_embd_state) {
+ LLAMA_LOG_ERROR("%s: state size mismatch, got %zu rows x %zu = %zu, expected %zu\n",
+ __func__, state_in.n_rows, state_in.n_embd, n_total, n_embd_state);
+ return false;
+ }
+
+ token & t = tokens[idx];
+
+ if (t.has_state) {
+ LLAMA_LOG_ERROR("%s: state for token %d is already set\n", __func__, idx);
+ return false;
+ }
+
+ t.has_state = true;
+ t.state_off = state.size();
+ state.insert(state.end(), state_in.data, state_in.data + n_total);
+
+ return true;
+}
+
bool llama_batch_ext::set_token_pos(int32_t idx, const llama_pos * pos_in) {
if (idx < 0 || idx >= (int32_t) tokens.size()) {
return false;
@@ -1320,11 +1386,7 @@ bool llama_batch_ext_set_embd_token(llama_batch_ext * batch, int32_t idx, llama_
}
bool llama_batch_ext_set_embd_state(llama_batch_ext * batch, int32_t idx, llama_embd embd) {
- // TODO
- GGML_UNUSED(batch);
- GGML_UNUSED(idx);
- GGML_UNUSED(embd);
- return false;
+ return batch->set_token_state(idx, embd);
}
bool llama_batch_ext_set_output_embd(llama_batch_ext * batch, int32_t idx, bool value) {
@@ -1393,7 +1455,13 @@ void llama_batch_compat::init(llama_batch_ext & dst, const llama_batch & batch_i
t.id = batch_inp.token[i];
}
- if (has_embd) {
+ // legacy MTP hook batches carry the hidden state next to the token ids
+ if (has_embd && has_token && batch_ext->n_embd_state > 0) {
+ t.has_state = true;
+ t.state_off = batch_ext->state.size();
+ const float * src = batch_inp.embd + (size_t) i * batch_ext->n_embd_state;
+ batch_ext->state.insert(batch_ext->state.end(), src, src + batch_ext->n_embd_state);
+ } else if (has_embd) {
t.has_embd = true;
t.embd_off = batch_ext->embd.size();
const float * src = batch_inp.embd + (size_t) i * n_embd_row;
diff --git a/src/llama-batch.h b/src/llama-batch.h
index ff62e26f7..8882c9941 100644
--- a/src/llama-batch.h
+++ b/src/llama-batch.h
@@ -48,10 +48,11 @@ struct llama_ubatch {
// seq_idx: indices of the unique sequence ids in the ubatch in [0, n_seqs_unq)
// used for extracting sequence pooled embeddings
- // // size | idx | val
- llama_token * token; // [n_tokens] | i | id, token
- float * embd; // [n_embd, n_tokens] | i | embd
- llama_pos * pos; // [n_tokens*n_pos] | i | pos
+ // // size | idx | val
+ llama_token * token; // [n_tokens] | i | id, token
+ float * embd; // [n_embd, n_tokens] | i | embd
+ float * embd_state; // [n_embd_state, n_tokens] | i | hidden state carried over from a previous stage (e.g. MTP)
+ llama_pos * pos; // [n_tokens*n_pos] | i | pos
int32_t * n_seq_id; // [n_tokens] | i | -
llama_seq_id ** seq_id; // [n_tokens] | s | s0, s1, seq_id
llama_seq_id * seq_id_unq; // [n_seqs_unq] | s | seq_id
@@ -63,6 +64,7 @@ struct llama_ubatch {
struct data_t {
std::vector<llama_token> token;
std::vector<float> embd;
+ std::vector<float> embd_state;
std::vector<llama_pos> pos;
std::vector<int32_t> n_seq_id;
std::vector<llama_seq_id *> seq_id; // these point into the seq_id_data below
@@ -85,15 +87,18 @@ struct llama_ubatch {
struct llama_hparams;
-// MTP hook batches carry the target model's hidden state (n_embd_out size).
// DFlash batches carry the fused target features at the encoder input width (n_embd_inp_enc size).
-// Normal batches carry token embeddings (n_embd_inp size).
+// Other batches carry token embeddings (n_embd_inp size).
size_t llama_batch_ext_select_n_embd_inp(llama_context_type ctx_type, llm_arch arch, const llama_hparams & hparams);
+// MTP contexts also take the target model's hidden state (n_embd_out size), 0 = no state input
+size_t llama_batch_ext_select_n_embd_state(llama_context_type ctx_type, const llama_hparams & hparams);
+
struct llama_batch_ext {
const size_t n_tokens_max; // max number of tokens that can be stored in the batch
const size_t n_embd_inp; // decoder embd row width
const size_t n_embd_inp_enc; // encoder embd row width (e.g. eagle3/dflash extracted features)
+ const size_t n_embd_state; // state embd row width, 0 if the context takes no state
const llama_seq_id n_seq_max; // max number of sequences
llama_memory_i * mem; // memory for position inference
const llama_token n_vocab; // max token ID that we accept
@@ -107,6 +112,8 @@ struct llama_batch_ext {
llama_token id = LLAMA_TOKEN_NULL;
bool has_embd = false; // whether embd_off is set
size_t embd_off = 0; // index offset in the embd array
+ bool has_state = false; // whether state_off is set
+ size_t state_off = 0; // index offset in the state array
bool output = false; // TODO: have dedicated output flags
int32_t decision_order = 0; // see llama_batch_ext_set_decision_order()
std::unordered_set<llama_seq_id> seq_ids;
@@ -114,6 +121,7 @@ struct llama_batch_ext {
};
std::vector<token> tokens;
std::vector<float> embd;
+ std::vector<float> state;
llama_batch_ext(llama_context * ctx);
@@ -136,6 +144,7 @@ struct llama_batch_ext {
bool add_seq(int32_t idx, llama_seq_id seq_id);
bool set_token_id(int32_t idx, llama_token id);
bool set_token_embd(int32_t idx, llama_embd embd_in);
+ bool set_token_state(int32_t idx, llama_embd state_in);
bool set_token_pos(int32_t idx, const llama_pos * pos_in);
bool set_output(int32_t idx, bool output_last);
bool set_decision_order(int32_t idx, int32_t order);
@@ -205,12 +214,14 @@ private:
const bool allow_mixed;
uint32_t n_embd;
+ uint32_t n_embd_state;
uint32_t n_seq_max;
uint32_t n_outputs;
std::vector<llama_token> token_vec; // owned token IDs built from llama_batch_ext
std::vector<float> embd_vec; // owned embeddings built from llama_batch_ext
std::vector<int8_t> is_embd_vec; // mixed batch only (= 1 if embd, 0 if text token)
+ std::vector<float> state_vec; // owned state embeddings built from llama_batch_ext, llama_batch has no slot for them
std::vector<llama_seq_id> seq_id_data; // flat storage for seq_id pointers below
std::vector<llama_pos> pos;
diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp
index 1514b8aeb..b22b95e6d 100644
--- a/src/llama-graph.cpp
+++ b/src/llama-graph.cpp
@@ -149,25 +149,21 @@ void llm_graph_input_embd_h::set_input(const llama_ubatch * ubatch) {
GGML_ASSERT(ubatch->embd);
GGML_ASSERT(n_embd == embd->ne[0]);
- ggml_backend_tensor_set(embd, ubatch->embd, 0, n_tokens*n_embd*ggml_element_size(h));
+ ggml_backend_tensor_set(embd, ubatch->embd, 0, n_tokens*n_embd*ggml_element_size(embd));
}
- // TODO: extend llama_ubatch to differentiate between token embeddings and hidden states
- // for now, we assume that the hidden state is always provided as an embedding
- // ref: https://github.com/ggml-org/llama.cpp/pull/23643
- if (ubatch->embd) {
- GGML_ASSERT(n_embd == h->ne[0]);
+ GGML_ASSERT(ubatch->embd_state && "this graph requires a state embedding, see llama_batch_ext_set_embd_state()");
+ GGML_ASSERT(n_embd_state == h->ne[0]);
- ggml_backend_tensor_set(h, ubatch->embd, 0, n_tokens*n_embd*ggml_element_size(h));
- }
+ ggml_backend_tensor_set(h, ubatch->embd_state, 0, n_tokens*n_embd_state*ggml_element_size(h));
}
bool llm_graph_input_embd_h::can_reuse(const llm_graph_params & params) {
bool res = true;
- res &= (!params.ubatch.token) || (tokens && tokens->ne[0] == params.ubatch.n_tokens);
- res &= (!params.ubatch.embd) || (embd && embd->ne[1] == params.ubatch.n_tokens);
- res &= (!params.ubatch.embd) || (h && h->ne[1] == params.ubatch.n_tokens);
+ res &= (!params.ubatch.token) || (tokens && tokens->ne[0] == params.ubatch.n_tokens);
+ res &= (!params.ubatch.embd) || (embd && embd->ne[1] == params.ubatch.n_tokens);
+ res &= (!params.ubatch.embd_state) || (h && h->ne[1] == params.ubatch.n_tokens);
return res;
}
diff --git a/src/llama-graph.h b/src/llama-graph.h
index c469847a8..06bb8c472 100644
--- a/src/llama-graph.h
+++ b/src/llama-graph.h
@@ -149,10 +149,10 @@ public:
const int64_t n_embd = 0;
};
-// similar to llm_graph_input_embd but with an additional hidden state input
+// similar to llm_graph_input_embd but with an additional hidden state input, fed from ubatch.embd_state
class llm_graph_input_embd_h : public llm_graph_input_i {
public:
- llm_graph_input_embd_h(int64_t n_embd) : n_embd(n_embd) {}
+ llm_graph_input_embd_h(int64_t n_embd, int64_t n_embd_state) : n_embd(n_embd), n_embd_state(n_embd_state) {}
virtual ~llm_graph_input_embd_h() = default;
void set_input(const llama_ubatch * ubatch) override;
@@ -161,9 +161,10 @@ public:
ggml_tensor * tokens = nullptr; // I32 [n_batch]
ggml_tensor * embd = nullptr; // F32 [n_embd, n_batch]
- ggml_tensor * h = nullptr; // F32 [n_embd, n_batch]
+ ggml_tensor * h = nullptr; // F32 [n_embd_state, n_batch]
- const int64_t n_embd = 0;
+ const int64_t n_embd = 0;
+ const int64_t n_embd_state = 0;
};
class llm_graph_input_pos : public llm_graph_input_i {
@@ -838,7 +839,8 @@ struct llm_graph_params {
(!ubatch.token && !other.ubatch.token) ||
(!ubatch.embd && !other.ubatch.embd) ||
(ubatch.token && other.ubatch.token && ubatch.embd && other.ubatch.embd)
- );
+ ) &&
+ (!ubatch.embd_state == !other.ubatch.embd_state);
// when we split the batch using "equal_seqs" we have to verify that the participating sequences are the same
// the reason is because the set of attention streams would be different for different sequences
diff --git a/src/llama-kv-cache-dsv4.cpp b/src/llama-kv-cache-dsv4.cpp
index 5dd0b4591..f8caae9d6 100644
--- a/src/llama-kv-cache-dsv4.cpp
+++ b/src/llama-kv-cache-dsv4.cpp
@@ -159,6 +159,7 @@ static llama_ubatch dsv4_build_raw_write_ubatch(const llama_ubatch & ubatch) {
/*.n_pos =*/ ubatch.n_pos,
/*.token =*/ data->token.empty() ? nullptr : data->token.data(),
/*.embd =*/ nullptr,
+ /*.embd_state =*/ nullptr,
/*.pos =*/ data->pos.data(),
/*.n_seq_id =*/ data->n_seq_id.data(),
/*.seq_id =*/ data->seq_id.data(),
diff --git a/src/models/bailingmoe3.cpp b/src/models/bailingmoe3.cpp
index 2458b1c1a..abbc30fdf 100644
--- a/src/models/bailingmoe3.cpp
+++ b/src/models/bailingmoe3.cpp
@@ -438,15 +438,17 @@ llama_model_bailingmoe3::graph_mtp::graph_mtp(const llama_model & model, const l
const int64_t kv_lora_rank = hparams.n_lora_kv;
const float kq_scale = 1.0f / sqrtf((float) qk_head_dim);
- auto inp = std::make_unique<llm_graph_input_embd>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
- inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd, n_tokens);
+ inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_inp(), n_tokens);
ggml_set_input(inp->embd);
- ggml_set_name(inp->embd, "mtp_h_input");
+ inp->h = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd, n_tokens);
+ ggml_set_input(inp->h);
+ ggml_set_name(inp->h, "mtp_h_input");
- ggml_tensor * tok_embd = ggml_get_rows(ctx0, model.tok_embd, inp->tokens);
- ggml_tensor * h_norm = build_norm(inp->embd, layer.nextn.hnorm, nullptr, LLM_NORM_RMS, il);
+ ggml_tensor * tok_embd = ubatch.token ? ggml_get_rows(ctx0, model.tok_embd, inp->tokens) : inp->embd;
+ ggml_tensor * h_norm = build_norm(inp->h, layer.nextn.hnorm, nullptr, LLM_NORM_RMS, il);
ggml_tensor * e_norm = build_norm(tok_embd, layer.nextn.enorm, nullptr, LLM_NORM_RMS, il);
ggml_tensor * cur = ggml_mul_mat(ctx0, layer.nextn.eh_proj, ggml_concat(ctx0, e_norm, h_norm, 0));
cb(cur, "mtp_eh_proj", il);
diff --git a/src/models/cohere2moe.cpp b/src/models/cohere2moe.cpp
index a379e2f60..19362e5c1 100644
--- a/src/models/cohere2moe.cpp
+++ b/src/models/cohere2moe.cpp
@@ -297,7 +297,7 @@ llama_model_cohere2moe::graph_mtp::graph_mtp(const llama_model & model, const ll
const llm_norm_type cohere2moe_norm_type = hparams.f_norm_rms_eps == 0.0f ? LLM_NORM : LLM_NORM_RMS;
// TODO: extract in a common llm_graph_context::build_inp_embd_h()
- auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
diff --git a/src/models/deepseek2.cpp b/src/models/deepseek2.cpp
index f1067a516..2d56dfb22 100644
--- a/src/models/deepseek2.cpp
+++ b/src/models/deepseek2.cpp
@@ -206,7 +206,7 @@ llama_model_deepseek2::graph_mtp::graph_mtp(const llama_model & model, const llm
GGML_ASSERT(layer.ffn_down_shexp);
GGML_ASSERT(layer.ffn_up_shexp);
- auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
diff --git a/src/models/deepseek32.cpp b/src/models/deepseek32.cpp
index 76f763c76..98f34a200 100644
--- a/src/models/deepseek32.cpp
+++ b/src/models/deepseek32.cpp
@@ -520,7 +520,7 @@ llama_model_deepseek32::graph_mtp::graph_mtp(const llama_model & model, const ll
const float kq_scale = 1.0f * mscale * mscale / sqrtf(float(n_embd_head_k));
// TODO: extract in a common llm_graph_context::build_inp_embd_h()
- auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
diff --git a/src/models/deepseek4.cpp b/src/models/deepseek4.cpp
index 0974b63b3..5df4c04c5 100644
--- a/src/models/deepseek4.cpp
+++ b/src/models/deepseek4.cpp
@@ -1378,20 +1378,26 @@ llama_model_deepseek4::graph_mtp::graph_mtp(const llama_model & model, const llm
GGML_ASSERT(layer.nextn.enorm && "MTP block missing nextn.enorm");
GGML_ASSERT(layer.nextn.hnorm && "MTP block missing nextn.hnorm");
- auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_out());
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd_out());
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
- inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_out(), n_tokens);
+ inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_inp(), n_tokens);
ggml_set_input(inp->embd);
inp->h = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_out(), n_tokens);
ggml_set_input(inp->h);
ggml_set_name(inp->h, "mtp_h_input");
- ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
- ggml_tensor * tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens);
+ ggml_tensor * tok_embd;
+ if (ubatch.token) {
+ ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
+
+ tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens);
+ } else {
+ tok_embd = inp->embd;
+ }
cb(tok_embd, "mtp_tok_embd", il);
ggml_tensor * h_state = ggml_reshape_3d(ctx0, inp->h, n_embd, hc, n_tokens);
diff --git a/src/models/gemma4-assistant.cpp b/src/models/gemma4-assistant.cpp
index 74d06151e..b6c29183e 100644
--- a/src/models/gemma4-assistant.cpp
+++ b/src/models/gemma4-assistant.cpp
@@ -86,9 +86,10 @@ llama_model_gemma4_assistant::graph::graph(const llama_model & model, const llm_
const int64_t n_embd_backbone = hparams.n_embd_inp();
ggml_tensor * inp_tokens;
+ ggml_tensor * inp_embd;
ggml_tensor * inp_h;
{
- auto inp = std::make_unique<llm_graph_input_embd>(n_embd_backbone);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(n_embd_backbone, n_embd_backbone);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, ubatch.n_tokens);
cb(inp->tokens, "inp_tokens", -1);
@@ -97,18 +98,23 @@ llama_model_gemma4_assistant::graph::graph(const llama_model & model, const llm_
res->t_inp_tokens = inp->tokens;
inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_embd_backbone, ubatch.n_tokens);
- cb(inp->embd, "inp_h", -1);
+ cb(inp->embd, "inp_embd", -1);
ggml_set_input(inp->embd);
- inp_h = inp->embd;
+ inp_embd = inp->embd;
res->t_inp_embd = inp->embd;
+ inp->h = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_embd_backbone, ubatch.n_tokens);
+ cb(inp->h, "inp_h", -1);
+ ggml_set_input(inp->h);
+ inp_h = inp->h;
+
res->add_input(std::move(inp));
}
GGML_ASSERT(cparams.ctx_other != nullptr);
const auto * model_other = llama_get_model(cparams.ctx_other);
- ggml_tensor * x = ggml_get_rows(ctx0, model_other->tok_embd, inp_tokens);
+ ggml_tensor * x = ubatch.token ? ggml_get_rows(ctx0, model_other->tok_embd, inp_tokens) : inp_embd;
x = ggml_scale(ctx0, x, sqrtf((float) n_embd_backbone));
cb(x, "inp_embd_target", -1);
diff --git a/src/models/glm-dsa.cpp b/src/models/glm-dsa.cpp
index 3e28ef271..596d86a61 100644
--- a/src/models/glm-dsa.cpp
+++ b/src/models/glm-dsa.cpp
@@ -560,7 +560,7 @@ llama_model_glm_dsa::graph_mtp::graph_mtp(const llama_model & model, const llm_g
const float kq_scale = 1.0f * mscale * mscale / sqrtf(float(n_embd_head_k));
// TODO: extract in a common llm_graph_context::build_inp_embd_h()
- auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
diff --git a/src/models/glm4-moe.cpp b/src/models/glm4-moe.cpp
index 4b41b5958..341d18a94 100644
--- a/src/models/glm4-moe.cpp
+++ b/src/models/glm4-moe.cpp
@@ -143,7 +143,7 @@ llama_model_glm4_moe::graph_mtp::graph_mtp(const llama_model & model, const llm_
GGML_ASSERT(layer.nextn.hnorm && "MTP block missing nextn.hnorm");
GGML_ASSERT(layer.ffn_gate_inp && "MTP block missing ffn_gate_inp");
- auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
diff --git a/src/models/glm5-next.cpp b/src/models/glm5-next.cpp
index 48b07f1e2..14224e64f 100644
--- a/src/models/glm5-next.cpp
+++ b/src/models/glm5-next.cpp
@@ -568,20 +568,25 @@ llama_model_glm5_next::graph_mtp::graph_mtp(const llama_model & model, const llm
ggml_tensor * inp_out_ids = build_inp_out_ids();
- auto inp = std::make_unique<llm_graph_input_embd_h>(n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
- inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_embd, n_tokens);
+ inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_inp(), n_tokens);
ggml_set_input(inp->embd);
inp->h = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_embd, n_tokens);
ggml_set_input(inp->h);
ggml_set_name(inp->h, "mtp_h_input");
- ggml_tensor * tok_embd = ggml_get_rows(ctx0,
- layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd, inp->tokens);
+ ggml_tensor * tok_embd;
+ if (ubatch.token) {
+ tok_embd = ggml_get_rows(ctx0,
+ layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd, inp->tokens);
+ } else {
+ tok_embd = inp->embd;
+ }
cb(tok_embd, "mtp_tok_embd", il);
ggml_tensor * h = inp->h;
diff --git a/src/models/hy-v3.cpp b/src/models/hy-v3.cpp
index c4c55c3f9..037191da8 100644
--- a/src/models/hy-v3.cpp
+++ b/src/models/hy-v3.cpp
@@ -245,19 +245,22 @@ llama_model_hy_v3::graph_mtp::graph_mtp(const llama_model & model, const llm_gra
GGML_ASSERT(layer.nextn.enorm && "MTP block missing nextn.enorm");
GGML_ASSERT(layer.nextn.hnorm && "MTP block missing nextn.hnorm");
- auto inp = std::make_unique<llm_graph_input_embd>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
- inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd, n_tokens);
+ inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_inp(), n_tokens);
ggml_set_input(inp->embd);
- ggml_set_name(inp->embd, "mtp_h_input");
+
+ inp->h = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd, n_tokens);
+ ggml_set_input(inp->h);
+ ggml_set_name(inp->h, "mtp_h_input");
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
- ggml_tensor * h_input = inp->embd;
- ggml_tensor * tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens);
+ ggml_tensor * h_input = inp->h;
+ ggml_tensor * tok_embd = ubatch.token ? ggml_get_rows(ctx0, tok_embd_w, inp->tokens) : inp->embd;
cb(tok_embd, "mtp_tok_embd", il);
res->add_input(std::move(inp));
diff --git a/src/models/mimo2.cpp b/src/models/mimo2.cpp
index a70350191..238dcfd91 100644
--- a/src/models/mimo2.cpp
+++ b/src/models/mimo2.cpp
@@ -282,18 +282,21 @@ llama_model_mimo2::graph_mtp::graph_mtp(const llama_model & model, const llm_gra
const float freq_scale_l = model.get_rope_freq_scale(cparams, il);
const float v_scale = hparams.f_attn_value_scale;
- auto inp = std::make_unique<llm_graph_input_embd>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
- inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd, n_tokens);
+ inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_inp(), n_tokens);
ggml_set_input(inp->embd);
- ggml_set_name(inp->embd, "mtp_h_input");
+
+ inp->h = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd, n_tokens);
+ ggml_set_input(inp->h);
+ ggml_set_name(inp->h, "mtp_h_input");
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
- ggml_tensor * h_input = inp->embd;
- ggml_tensor * tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens);
+ ggml_tensor * h_input = inp->h;
+ ggml_tensor * tok_embd = ubatch.token ? ggml_get_rows(ctx0, tok_embd_w, inp->tokens) : inp->embd;
cb(tok_embd, "mtp_tok_embd", il);
res->add_input(std::move(inp));
diff --git a/src/models/nemotron-h-moe.cpp b/src/models/nemotron-h-moe.cpp
index f1e3ce3b4..d3b8bb837 100644
--- a/src/models/nemotron-h-moe.cpp
+++ b/src/models/nemotron-h-moe.cpp
@@ -25,7 +25,7 @@ llama_model_nemotron_h_moe::graph_mtp::graph_mtp(const llama_model & model, cons
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
GGML_ASSERT(tok_embd_w != nullptr && "NEMOTRON_H_MOE MTP requires token embeddings");
- auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
diff --git a/src/models/qwen35.cpp b/src/models/qwen35.cpp
index 350cea808..ab3b2aed1 100644
--- a/src/models/qwen35.cpp
+++ b/src/models/qwen35.cpp
@@ -518,7 +518,7 @@ llama_model_qwen35::graph_mtp::graph_mtp(const llama_model & model, const llm_gr
std::copy(std::begin(hparams.rope_sections), std::begin(hparams.rope_sections) + 4, sections);
// TODO: extract in a common llm_graph_context::build_inp_embd_h()
- auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
diff --git a/src/models/qwen35moe.cpp b/src/models/qwen35moe.cpp
index f9cc8a2c6..118deeb2a 100644
--- a/src/models/qwen35moe.cpp
+++ b/src/models/qwen35moe.cpp
@@ -568,7 +568,7 @@ llama_model_qwen35moe::graph_mtp::graph_mtp(const llama_model & model, const llm
std::copy(std::begin(hparams.rope_sections), std::begin(hparams.rope_sections) + 4, sections);
// TODO: extract in a common llm_graph_context::build_inp_embd_h()
- auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
diff --git a/src/models/qwen3next.cpp b/src/models/qwen3next.cpp
index 620161ac7..49c28b134 100644
--- a/src/models/qwen3next.cpp
+++ b/src/models/qwen3next.cpp
@@ -642,7 +642,7 @@ llama_model_qwen3next::graph_mtp::graph_mtp(const llama_model & model, const llm
GGML_ASSERT(layer.ffn_gate_inp && "MTP block missing ffn_gate_inp");
// TODO: extract in a common llm_graph_context::build_inp_embd_h()
- auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp
index ad01b5388..c1771fcad 100644
--- a/src/models/qwen4exp.cpp
+++ b/src/models/qwen4exp.cpp
@@ -540,19 +540,24 @@ llama_model_qwen4exp::graph_mtp::graph_mtp(const llama_model & model, const llm_
int sections[4];
std::copy(std::begin(hparams.rope_sections), std::begin(hparams.rope_sections) + 4, sections);
- auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_out());
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd_out());
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
- inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_out(), n_tokens);
+ inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_inp(), n_tokens);
ggml_set_input(inp->embd);
inp->h = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_out(), n_tokens);
ggml_set_input(inp->h);
ggml_set_name(inp->h, "mtp_h_input");
- ggml_tensor * tok_embd = ggml_get_rows(ctx0, model.tok_embd, inp->tokens);
+ ggml_tensor * tok_embd;
+ if (ubatch.token) {
+ tok_embd = ggml_get_rows(ctx0, model.tok_embd, inp->tokens);
+ } else {
+ tok_embd = inp->embd;
+ }
cb(tok_embd, "mtp_tok_embd", il);
ggml_tensor * h = inp->h;
diff --git a/src/models/step35.cpp b/src/models/step35.cpp
index bfa80fab3..bd43d9ae8 100644
--- a/src/models/step35.cpp
+++ b/src/models/step35.cpp
@@ -380,19 +380,22 @@ llama_model_step35::graph_mtp::graph_mtp(const llama_model & model, const llm_gr
const float freq_base_l = model.get_rope_freq_base(cparams, il);
const float freq_scale_l = model.get_rope_freq_scale(cparams, il);
- auto inp = std::make_unique<llm_graph_input_embd>(hparams.n_embd);
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_inp(), hparams.n_embd);
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
ggml_set_input(inp->tokens);
- inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd, n_tokens);
+ inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_inp(), n_tokens);
ggml_set_input(inp->embd);
- ggml_set_name(inp->embd, "mtp_h_input");
+
+ inp->h = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd, n_tokens);
+ ggml_set_input(inp->h);
+ ggml_set_name(inp->h, "mtp_h_input");
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
- ggml_tensor * h_input = inp->embd;
- ggml_tensor * tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens);
+ ggml_tensor * h_input = inp->h;
+ ggml_tensor * tok_embd = ubatch.token ? ggml_get_rows(ctx0, tok_embd_w, inp->tokens) : inp->embd;
cb(tok_embd, "mtp_tok_embd", il);
res->add_input(std::move(inp));
diff --git a/tests/test-batch-alloc.cpp b/tests/test-batch-alloc.cpp
index b085917cf..c41e91d61 100644
--- a/tests/test-batch-alloc.cpp
+++ b/tests/test-batch-alloc.cpp
@@ -1132,7 +1132,7 @@ static void test_compat(testing & t) {
}
static void test_mtp_embd_width(testing & t) {
- t.test("mtp_uses_n_embd_out", [&](testing & t) {
+ t.test("mtp_keeps_n_embd_inp_and_takes_state_at_n_embd_out", [&](testing & t) {
llama_hparams hparams = {};
hparams.n_embd = 64;
hparams.n_deepstack_layers = 2; // makes n_embd_inp() = 64 + 64*2 = 192
@@ -1141,16 +1141,22 @@ static void test_mtp_embd_width(testing & t) {
t.assert_equal("default context uses n_embd_inp (deepstack-aware)",
(size_t) 192, llama_batch_ext_select_n_embd_inp(LLAMA_CONTEXT_TYPE_DEFAULT, LLM_ARCH_LLAMA, hparams));
- t.assert_equal("MTP context uses n_embd_out instead (target-model hidden state width)",
- (size_t) 96, llama_batch_ext_select_n_embd_inp(LLAMA_CONTEXT_TYPE_MTP, LLM_ARCH_LLAMA, hparams));
+ t.assert_equal("MTP context keeps n_embd_inp for the token embeddings",
+ (size_t) 192, llama_batch_ext_select_n_embd_inp(LLAMA_CONTEXT_TYPE_MTP, LLM_ARCH_LLAMA, hparams));
+
+ t.assert_equal("MTP context takes the target hidden state at n_embd_out",
+ (size_t) 96, llama_batch_ext_select_n_embd_state(LLAMA_CONTEXT_TYPE_MTP, hparams));
+
+ t.assert_equal("default context takes no state",
+ (size_t) 0, llama_batch_ext_select_n_embd_state(LLAMA_CONTEXT_TYPE_DEFAULT, hparams));
});
- t.test("mtp_falls_back_to_n_embd_when_no_override", [&](testing & t) {
+ t.test("mtp_state_falls_back_to_n_embd_when_no_override", [&](testing & t) {
llama_hparams hparams = {};
hparams.n_embd = 64; // no deepstack, no n_embd_out_impl override
- t.assert_equal((size_t) 64, llama_batch_ext_select_n_embd_inp(LLAMA_CONTEXT_TYPE_DEFAULT, LLM_ARCH_LLAMA, hparams));
t.assert_equal((size_t) 64, llama_batch_ext_select_n_embd_inp(LLAMA_CONTEXT_TYPE_MTP, LLM_ARCH_LLAMA, hparams));
+ t.assert_equal((size_t) 64, llama_batch_ext_select_n_embd_state(LLAMA_CONTEXT_TYPE_MTP, hparams));
});
t.test("dflash_uses_n_embd_inp_enc", [&](testing & t) {
@@ -1165,8 +1171,8 @@ static void test_mtp_embd_width(testing & t) {
t.assert_equal("other archs ignore n_embd_inp_enc",
(size_t) 64, llama_batch_ext_select_n_embd_inp(LLAMA_CONTEXT_TYPE_DEFAULT, LLM_ARCH_LLAMA, hparams));
- t.assert_equal("MTP takes precedence over DFlash",
- (size_t) 96, llama_batch_ext_select_n_embd_inp(LLAMA_CONTEXT_TYPE_MTP, LLM_ARCH_DFLASH, hparams));
+ t.assert_equal("MTP context does not change the DFlash input width",
+ (size_t) 128, llama_batch_ext_select_n_embd_inp(LLAMA_CONTEXT_TYPE_MTP, LLM_ARCH_DFLASH, hparams));
});
}