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));
     });
 }