Commit 0bb496dbd for llama.cpp

commit 0bb496dbd3af0add77ff82c406a915b41e839d56
Author: Xuan-Son Nguyen <son@huggingface.co>
Date:   Mon Oct 5 01:35:49 2026 +0200

    llama: support both embd + raw tokens in batch (#29622)

    * llama: support both embd + raw tokens in batch

    * add to test-llama-archs

    * also check case llm_arch_supports_mixed_batch = false

    * constant graph topology

    * have dedicated input for mixed case

    * rm set_tensor_backend

    * is_embd --> type

    * consolidate m-rope pos handling into one place

    * nits

diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp
index 2af5445df..850737eca 100644
--- a/src/llama-arch.cpp
+++ b/src/llama-arch.cpp
@@ -1178,6 +1178,21 @@ bool llm_arch_supports_rs_rollback(const llm_arch & arch) {
     }
 }

+// these models pick weights, routing or input meaning per ubatch based on token vs embd input
+bool llm_arch_supports_mixed_batch(const llm_arch & arch) {
+    switch (arch) {
+        case LLM_ARCH_COGVLM:
+        case LLM_ARCH_DEEPSEEK4:
+        case LLM_ARCH_GRANITE_SWITCH:
+        case LLM_ARCH_EAGLE3:
+        case LLM_ARCH_DFLASH:
+        case LLM_ARCH_GEMMA4_ASSISTANT:
+            return false;
+        default:
+            return true;
+    }
+}
+
 bool llm_arch_supports_sm_tensor(const llm_arch & arch) {
     switch (arch) {
         case LLM_ARCH_GROK:
diff --git a/src/llama-arch.h b/src/llama-arch.h
index 148d293ce..24068fb7a 100644
--- a/src/llama-arch.h
+++ b/src/llama-arch.h
@@ -825,3 +825,4 @@ bool llm_arch_is_hybrid         (const llm_arch & arch);
 bool llm_arch_is_diffusion      (const llm_arch & arch);
 bool llm_arch_supports_sm_tensor(const llm_arch & arch);
 bool llm_arch_supports_rs_rollback(const llm_arch & arch);
+bool llm_arch_supports_mixed_batch(const llm_arch & arch);
diff --git a/src/llama-batch.cpp b/src/llama-batch.cpp
index 1b3d70627..ecd48dd80 100644
--- a/src/llama-batch.cpp
+++ b/src/llama-batch.cpp
@@ -12,7 +12,7 @@
 #include <algorithm>
 #include <sstream>

-llama_batch_allocr::llama_batch_allocr(uint32_t n_pos_per_embd) : n_pos_per_embd(n_pos_per_embd) {
+llama_batch_allocr::llama_batch_allocr(uint32_t n_pos_per_embd, bool allow_mixed) : n_pos_per_embd(n_pos_per_embd), allow_mixed(allow_mixed) {
     const char * LLAMA_BATCH_DEBUG = getenv("LLAMA_BATCH_DEBUG");
     debug = LLAMA_BATCH_DEBUG ? atoi(LLAMA_BATCH_DEBUG) : 0;

@@ -49,25 +49,49 @@ 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)
-    // all entries must carry the same combination
+    // all entries must carry the same combination, or be a mix of token and embd entries
     //

-    const bool has_token = batch_inp.tokens[0].id != LLAMA_TOKEN_NULL;
-    const bool has_embd  = batch_inp.tokens[0].has_embd;
+    int32_t n_tok_only  = 0;
+    int32_t n_embd_only = 0;
+    int32_t n_both      = 0;

-    for (int32_t i = 1; i < n_tok; ++i) {
-        if ((batch_inp.tokens[i].id != LLAMA_TOKEN_NULL) != has_token ||
-             batch_inp.tokens[i].has_embd                != has_embd) {
-            LLAMA_LOG_ERROR("%s: all entries in the batch must have the same content types\n", __func__);
+    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;
+
+        if (!is_tok && !is_emb) {
+            LLAMA_LOG_ERROR("%s: entry %d has neither a token id nor an embedding\n", __func__, i);
             return false;
         }
+
+        n_tok_only  += is_tok && !is_emb;
+        n_embd_only += is_emb && !is_tok;
+        n_both      += is_tok &&  is_emb;
     }

-    if (!has_token && !has_embd) {
-        LLAMA_LOG_ERROR("%s: batch has neither token ids nor embeddings\n", __func__);
+    if (n_both > 0 && n_both != n_tok) {
+        LLAMA_LOG_ERROR("%s: entries with both a token id and an embedding cannot be mixed with other entries\n", __func__);
         return false;
     }

+    const bool mixed = n_tok_only > 0 && n_embd_only > 0;
+
+    if (mixed && !allow_mixed) {
+        LLAMA_LOG_ERROR("%s: this model or context does not support batches mixing token and embedding entries\n", __func__);
+        return false;
+    }
+
+    const bool has_token = n_tok_only  > 0 || n_both > 0;
+    const bool has_embd  = n_embd_only > 0 || n_both > 0;
+
+    if (mixed) {
+        is_embd_vec.resize(n_tok);
+        for (int32_t i = 0; i < n_tok; ++i) {
+            is_embd_vec[i] = batch_inp.tokens[i].has_embd ? 1 : 0;
+        }
+    }
+
     //
     // build flat token/embd array
     //
@@ -75,6 +99,10 @@ bool llama_batch_allocr::init(
     if (has_token) {
         token_vec.resize(n_tok);
         for (int32_t i = 0; i < n_tok; ++i) {
+            if (mixed && is_embd_vec[i]) {
+                token_vec[i] = 0; // placeholder
+                continue;
+            }
             const llama_token id = batch_inp.tokens[i].id;
             if (id < 0 || id >= batch_inp.n_vocab) {
                 LLAMA_LOG_ERROR("%s: invalid token[%d] = %d\n", __func__, i, id);
@@ -84,29 +112,35 @@ bool llama_batch_allocr::init(
         }
     }

-    if (has_embd) {
+    if (mixed) {
+        embd_vec.assign((size_t) n_tok*n_embd, 0.0f);
+        for (int32_t i = 0; i < n_tok; ++i) {
+            if (is_embd_vec[i]) {
+                const float * src = batch_inp.embd.data() + batch_inp.tokens[i].embd_off;
+                std::copy(src, src + n_embd, embd_vec.data() + (size_t) i*n_embd);
+            }
+        }
+    } else if (has_embd) {
         embd_vec = batch_inp.embd;
     }

     //
-    // build flat pos array
-    // token batch:     pos[i]            = tokens[i].pos[0]
-    // embedding batch: pos[j*n_tok + i]  = tokens[i].pos[j]  (section-major)
+    // 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)
+    // embd entry:  tokens[i].pos as-is
     //

-    {
-        const int32_t n_pos_total = has_token ? n_tok : n_tok * (int32_t) n_pos_per_embd;
-        pos.resize(n_pos_total);
-        if (has_token) {
-            for (int32_t i = 0; i < n_tok; ++i) {
-                pos[i] = batch_inp.tokens[i].pos[0];
-            }
-        } else {
-            for (int32_t i = 0; i < n_tok; ++i) {
-                for (uint32_t j = 0; j < n_pos_per_embd; ++j) {
-                    pos[(int32_t) j * n_tok + i] = batch_inp.tokens[i].pos[j];
-                }
+    pos.resize((size_t) n_tok*n_pos_per_embd);
+    for (int32_t i = 0; i < n_tok; ++i) {
+        const auto & tok = batch_inp.tokens[i];
+        const bool expand = tok.id != LLAMA_TOKEN_NULL;
+        for (uint32_t j = 0; j < n_pos_per_embd; ++j) {
+            llama_pos p = tok.pos[j];
+            if (expand) {
+                // expand [p] to [p, p, p, 0] for M-RoPE
+                p = j < 3 ? tok.pos[0] : 0;
             }
+            pos[(size_t) j*n_tok + i] = p;
         }
     }

@@ -264,6 +298,7 @@ bool llama_batch_allocr::init(
             /*.seq_id_unq   =*/ this->seq_id_unq.data(),
             /*.seq_idx      =*/ this->seq_idx.data(),
             /*.output       =*/ batch.logits,
+            /*.type         =*/ is_embd_vec.empty() ? nullptr : is_embd_vec.data(),
             /*.decision_order =*/ decision_order.empty() ? nullptr : decision_order.data(),
             /*.data         =*/ {},
         };
@@ -294,6 +329,21 @@ bool llama_batch_allocr::init(
     //

     if (n_pos_per_embd > 1) {
+        // in a mixed batch, the first entry of each seq picks the rule
+        std::vector<int8_t> seq_first_embd(n_seq_max, batch.token ? 0 : 1);
+        if (mixed) {
+            std::vector<bool> seen(n_seq_max, false);
+            for (int32_t i = 0; i < batch.n_tokens; ++i) {
+                for (int32_t s = 0; s < batch.n_seq_id[i]; ++s) {
+                    const llama_seq_id sid = batch.seq_id[i][s];
+                    if (!seen[sid]) {
+                        seen[sid] = true;
+                        seq_first_embd[sid] = is_embd_vec[i];
+                    }
+                }
+            }
+        }
+
         // M-RoPE case: allow position to "jump" forward only (non-continuous positions are allowed)
         for (uint32_t s = 0; s < n_seq_max; ++s) {
             if (seq_pos[s].empty()) {
@@ -302,7 +352,7 @@ bool llama_batch_allocr::init(

             const llama_pos p0 = mem ? mem->seq_pos_max(s) : -1;

-            if (batch.token) {
+            if (!seq_first_embd[s]) {
                 if (p0 >= 0 && p0 >= seq_pos_min(s)) {
                     LLAMA_LOG_ERROR(
                             "%s: the tokens of sequence %d in the input batch have inconsistent sequence positions:\n"
@@ -471,6 +521,7 @@ llama_ubatch llama_batch_allocr::ubatch_reserve(uint32_t n_seq_tokens, uint32_t
         /*.seq_id_unq   =*/ udata->seq_id_unq.data(),
         /*.seq_idx      =*/ udata->seq_idx.data(),
         /*.output       =*/ udata->output.data(),
+        /*.type         =*/ nullptr,
         /*.decision_order =*/ nullptr,
         /*.data         =*/ std::move(udata),
     };
@@ -769,6 +820,7 @@ void llama_batch_allocr::clear() {

     token_vec   .clear();
     embd_vec    .clear();
+    is_embd_vec .clear();
     seq_id_data .clear();
     pos         .clear();
     n_seq_id    .clear();
@@ -799,7 +851,20 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u

     auto udata = std::make_shared<llama_ubatch::data_t>();

-    const int64_t n_embd_all = batch.embd ? (int64_t) n_tokens*n_embd : 0;
+    const bool mixed_batch = !is_embd_vec.empty();
+
+    // a ubatch with a single kind of rows is emitted as a plain token or embd ubatch
+    uint32_t n_embd_rows = 0;
+    if (mixed_batch) {
+        for (int32_t idx : idxs) {
+            n_embd_rows += is_embd_vec[idx];
+        }
+    }
+    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 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;

     udata->token     .resize(n_tokens);
@@ -810,6 +875,7 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u
     udata->seq_id_unq.resize(0);
     udata->seq_idx   .resize(LLAMA_MAX_SEQ, -1);
     udata->output    .resize(n_tokens);
+    udata->type      .resize(mixed ? n_tokens : 0);
     udata->decision_order.resize(decision_order.empty() ? 0 : n_tokens);

     udata->batch_idxs = idxs;
@@ -818,21 +884,20 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u
     seq_set_t seq_set_unq;

     for (size_t i = 0; i < idxs.size(); ++i) {
-        if (batch.token) {
+        if (use_token) {
             udata->token[i] = batch.token[idxs[i]];
         }

-        if (batch.embd) {
+        if (use_embd) {
             memcpy(udata->embd.data() + i*n_embd, batch.embd + (int64_t) idxs[i]*n_embd, n_embd*sizeof(float));
         }

+        if (mixed) {
+            udata->type[i] = is_embd_vec[idxs[i]];
+        }
+
         for (size_t j = 0; j < (size_t)n_pos_per_embd; ++j) {
-            // if we are using M-RoPE
-            //     if the current batch is text, we need to broadcast the same position across all RoPE sections
-            //     otherwise, the input batch is image embeddings, we copy the positions as-is
-            // if we are not using M-RoPE, there is only one position per token (this loop runs only once)
-            size_t src_off = batch.token ? 0 : j*batch.n_tokens;
-            udata->pos[j*n_tokens + i] = batch.pos[src_off + idxs[i]];
+            udata->pos[j*n_tokens + i] = batch.pos[j*batch.n_tokens + idxs[i]];
         }

         udata->n_seq_id[i] = batch.n_seq_id[idxs[i]];
@@ -875,14 +940,15 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u
         /*.n_seqs_unq   =*/ (uint32_t) udata->seq_id_unq.size(),
         /*.n_pos        =*/ n_pos_per_embd,

-        /*.token        =*/ batch.token ? udata->token.data() : nullptr,
-        /*.embd         =*/ batch.embd ? udata->embd.data() : nullptr,
+        /*.token        =*/ use_token ? udata->token.data() : nullptr,
+        /*.embd         =*/ use_embd  ? udata->embd.data()  : nullptr,
         /*.pos          =*/ udata->pos.data(),
         /*.n_seq_id     =*/ udata->n_seq_id.data(),
         /*.seq_id       =*/ udata->seq_id.data(),
         /*.seq_id_unq   =*/ udata->seq_id_unq.data(),
         /*.seq_idx      =*/ udata->seq_idx.data(),
         /*.output       =*/ udata->output.data(),
+        /*.type         =*/ mixed ? udata->type.data() : nullptr,
         /*.decision_order =*/ udata->decision_order.empty() ? nullptr : udata->decision_order.data(),
         /*.data         =*/ std::move(udata),
     };
@@ -933,6 +999,7 @@ void llama_batch_allocr::ubatch_print(const llama_ubatch & ubatch, int debug) {
         LLAMA_LOG_DEBUG("%s:   seq_id_unq = %s\n", __func__, ss_seq_id_unq.str().c_str());
         LLAMA_LOG_DEBUG("%s:   seq_idx    = %s\n", __func__, ss_seq_idx.str().c_str());
         LLAMA_LOG_DEBUG("%s:   output     = %p\n", __func__, (void *) ubatch.output);
+        LLAMA_LOG_DEBUG("%s:   type       = %p\n", __func__, (void *) ubatch.type);
         LLAMA_LOG_DEBUG("%s:   n_outputs  = %d\n", __func__, n_outputs);

         if (debug > 0) {
@@ -963,7 +1030,7 @@ void llama_batch_allocr::ubatch_print(const llama_ubatch & ubatch, int debug) {
                     }
                 }

-                if (ubatch.token) {
+                if (ubatch.token && !(ubatch.is_mixed() && ubatch.type[i])) {
                     LLAMA_LOG_DEBUG("%s:  %4d: id = %6d (%16s), pos = %4d, n_seq_id = %2d, seq_id = [%s], output = %d\n",
                             __func__, i, ubatch.token[i], vocab->token_to_piece(ubatch.token[i]).c_str(),
                             ubatch.pos[i], ubatch.n_seq_id[i], ss.str().c_str(), ubatch.output[i]);
diff --git a/src/llama-batch.h b/src/llama-batch.h
index 32f103bfc..ff62e26f7 100644
--- a/src/llama-batch.h
+++ b/src/llama-batch.h
@@ -29,6 +29,11 @@ struct llama_ubatch {
         return n_pos >= 3;
     }

+    // mixed: type picks token or embd per row, pos has n_pos sections for all rows
+    bool is_mixed() const {
+        return type != nullptr;
+    }
+
     uint32_t b_equal_seqs; // note: this is a boolean, but we use an int32_t for alignment
                            //       otherwise address sanitizer complains
     // TODO: whole_seqs for embeddings?
@@ -52,6 +57,7 @@ struct llama_ubatch {
     llama_seq_id *  seq_id_unq; // [n_seqs_unq]       | s   | seq_id
     int32_t      *  seq_idx;    // [LLAMA_MAX_SEQ]    | -   | seq_idx
     int8_t       *  output;     // [n_tokens]         | i   | -
+    int8_t       *  type;       // [n_tokens]         | i   | -     (mixed ubatch only, 0 - token, 1 - embd)
     int32_t      *  decision_order; // [n_tokens], NULL if no entry has one, see llama_batch_ext_set_decision_order()

     struct data_t {
@@ -63,6 +69,7 @@ struct llama_ubatch {
         std::vector<llama_seq_id>   seq_id_unq;
         std::vector<int32_t>        seq_idx;
         std::vector<int8_t>         output;
+        std::vector<int8_t>         type;
         std::vector<int32_t>        batch_idxs;  // original batch index for each token
         std::vector<int32_t>        decision_order;

@@ -73,6 +80,9 @@ struct llama_ubatch {
     std::shared_ptr<data_t> data;
 };

+// crash if a mixed ubatch reaches code that expects only tokens or only embd
+#define ASSERT_EMBD_OR_TOKEN(ubatch) GGML_ASSERT(!(ubatch).is_mixed() && "mixed token/embd ubatch is not supported here")
+
 struct llama_hparams;

 // MTP hook batches carry the target model's hidden state (n_embd_out size).
@@ -134,7 +144,7 @@ struct llama_batch_ext {
 // a helper for sanitizing, fulfilling and splitting a batch
 class llama_batch_allocr {
 public:
-    llama_batch_allocr(uint32_t n_pos_per_embd);
+    llama_batch_allocr(uint32_t n_pos_per_embd, bool allow_mixed = false);

     // convert a llama_batch_ext to internal llama_batch and sanitize it
     bool init(
@@ -192,12 +202,15 @@ private:
     //       ref: https://github.com/ggml-org/llama.cpp/issues/13694#issuecomment-2983871762
     const uint32_t n_pos_per_embd;

+    const bool allow_mixed;
+
     uint32_t n_embd;
     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<llama_seq_id>   seq_id_data;  // flat storage for seq_id pointers below

     std::vector<llama_pos>      pos;
diff --git a/src/llama-context.cpp b/src/llama-context.cpp
index 96b5464e6..07b4c6148 100644
--- a/src/llama-context.cpp
+++ b/src/llama-context.cpp
@@ -87,7 +87,9 @@ llama_context::llama_context(
     model(model),
     cvec(std::make_unique<llama_adapter_cvec>()),
     loras(std::make_unique<llama_adapter_loras>()),
-    balloc(std::make_unique<llama_batch_allocr>(model.hparams.n_pos_per_embd())) {
+    // MTP uses the embd input for the hidden state
+    balloc(std::make_unique<llama_batch_allocr>(model.hparams.n_pos_per_embd(),
+                llm_arch_supports_mixed_batch(model.arch) && params.ctx_type == LLAMA_CONTEXT_TYPE_DEFAULT)) {
     // TODO warning when creating llama_context with awkward ctx size that is not a power of 2,
     //     may need to be backend-dependent
     LLAMA_LOG_INFO("%s: constructing llama_context\n", __func__);
diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp
index 16a3ed2ab..1112ad885 100644
--- a/src/llama-graph.cpp
+++ b/src/llama-graph.cpp
@@ -73,13 +73,55 @@ void llm_graph_input_embd::set_input(const llama_ubatch * ubatch) {
         ggml_backend_tensor_set(tokens, ubatch->token, 0, n_tokens*ggml_element_size(tokens));
     }

-    if (ubatch->embd) {
+    if (ubatch->embd && embd && !ubatch->is_mixed()) {
         GGML_ASSERT(n_embd == embd->ne[0]);

         const int64_t n_tokens = ubatch->n_tokens;

         ggml_backend_tensor_set(embd, ubatch->embd, 0, n_tokens*n_embd*ggml_element_size(embd));
     }
+
+    if (ubatch->is_mixed() && embd) {
+        GGML_ASSERT(mixed_tokens && mixed_slots && mixed_embd && "mixed token/embd ubatch is not supported here");
+
+        std::vector<int32_t> ids;
+        std::vector<int64_t> slots;
+        for (uint32_t i = 0; i < ubatch->n_tokens; ++i) {
+            if (!ubatch->type[i]) {
+                ids.push_back(ubatch->token[i]);
+                slots.push_back(i);
+            }
+        }
+        GGML_ASSERT((int64_t) ids.size() == mixed_tokens->ne[0]);
+        GGML_ASSERT(n_embd == mixed_embd->ne[0]);
+
+        ggml_backend_tensor_set(mixed_tokens, ids.data(),    0, ggml_nbytes(mixed_tokens));
+        ggml_backend_tensor_set(mixed_slots,  slots.data(),  0, ggml_nbytes(mixed_slots));
+        ggml_backend_tensor_set(mixed_embd,   ubatch->embd,  0, ggml_nbytes(mixed_embd));
+    }
+
+    if (scale_rows) {
+        const int64_t n_tokens = ubatch->n_tokens;
+
+        std::vector<float> data(n_tokens);
+        for (int64_t i = 0; i < n_tokens; ++i) {
+            const bool is_embd = !ubatch->token || (ubatch->is_mixed() && ubatch->type[i]);
+            data[i] = is_embd ? 1.0f : scale_tok;
+        }
+        ggml_backend_tensor_set(scale_rows, data.data(), 0, ggml_nbytes(scale_rows));
+    }
+}
+
+// number of token rows of the mixed path, a non-mixed ubatch is sized for the worst case
+static int64_t llm_graph_n_tok_rows(const llama_ubatch & ubatch) {
+    if (!ubatch.is_mixed()) {
+        return ubatch.n_tokens;
+    }
+    int64_t n = 0;
+    for (uint32_t i = 0; i < ubatch.n_tokens; ++i) {
+        n += !ubatch.type[i];
+    }
+    return n;
 }

 bool llm_graph_input_embd::can_reuse(const llm_graph_params & params) {
@@ -87,11 +129,16 @@ bool llm_graph_input_embd::can_reuse(const llm_graph_params & params) {

     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 &= (!mixed_tokens) || mixed_tokens->ne[0] == llm_graph_n_tok_rows(params.ubatch);
+    res &= (!mixed_embd)   || mixed_embd->ne[1]   == params.ubatch.n_tokens;
+    res &= (!scale_rows) || scale_rows->ne[1] == params.ubatch.n_tokens;

     return res;
 }

 void llm_graph_input_embd_h::set_input(const llama_ubatch * ubatch) {
+    ASSERT_EMBD_OR_TOKEN(*ubatch);
+
     const int64_t n_tokens = ubatch->n_tokens;

     if (ubatch->token) {
@@ -128,21 +175,7 @@ void llm_graph_input_pos::set_input(const llama_ubatch * ubatch) {
     if (ubatch->pos && pos) {
         const int64_t n_tokens = ubatch->n_tokens;

-        if (ubatch->token && n_pos_per_embd == 4) {
-            // in case we're using M-RoPE with text tokens, convert the 1D positions to 4D
-            // the 3 first dims are the same, and 4th dim is all 0
-            std::vector<llama_pos> pos_data(n_tokens*n_pos_per_embd);
-            // copy the first dimension
-            for (int i = 0; i < n_tokens; ++i) {
-                pos_data[               i] = ubatch->pos[i];
-                pos_data[    n_tokens + i] = ubatch->pos[i];
-                pos_data[2 * n_tokens + i] = ubatch->pos[i];
-                pos_data[3 * n_tokens + i] = 0; // 4th dim is 0
-            }
-            ggml_backend_tensor_set(pos, pos_data.data(), 0, pos_data.size()*ggml_element_size(pos));
-        } else {
-            ggml_backend_tensor_set(pos, ubatch->pos, 0, n_tokens*n_pos_per_embd*ggml_element_size(pos));
-        }
+        ggml_backend_tensor_set(pos, ubatch->pos, 0, n_tokens*n_pos_per_embd*ggml_element_size(pos));
     }
 }

@@ -2377,7 +2410,7 @@ ggml_tensor * llm_graph_context::build_moe_ffn(
 }

 // input embeddings with optional lora
-ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd) const {
+ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd, float tok_scale) const {
     const int64_t n_embd_inp = hparams.n_embd_inp();
     const int64_t n_embd     = hparams.n_embd;

@@ -2394,15 +2427,9 @@ ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd) const {
     cb(inp->embd, "inp_embd", -1);
     ggml_set_input(inp->embd);

-    // select one of the 2 inputs, based on the batch contents
-    // ref: https://github.com/ggml-org/llama.cpp/pull/18550
-    std::array<ggml_tensor *, 2> inps;
-
-    // token embeddings path (ubatch.token != nullptr)
-    {
-        auto & cur = inps[0];
-
-        cur = ggml_get_rows(ctx0, tok_embd, inp->tokens);
+    // token embeddings with lora and padding
+    auto build_tok = [&](ggml_tensor * ids) {
+        ggml_tensor * cur = ggml_get_rows(ctx0, tok_embd, ids);

         // apply lora for embedding tokens if needed
         for (const auto & lora : *loras) {
@@ -2416,7 +2443,7 @@ ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd) const {

             ggml_tensor * inpL_delta = ggml_scale(ctx0, ggml_mul_mat(
                         ctx0, lw->b, // non-transposed lora_b
-                        ggml_get_rows(ctx0, lw->a, inp->tokens)
+                        ggml_get_rows(ctx0, lw->a, ids)
                         ), scale);

             cur = ggml_add(ctx0, cur, inpL_delta);
@@ -2425,19 +2452,48 @@ ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd) const {
         if (n_embd_inp != n_embd) {
             cur = ggml_pad(ctx0, cur, hparams.n_embd_inp() - n_embd, 0, 0, 0);
         }
-    }
+
+        return cur;
+    };
+
+    // select one of the 3 inputs, based on the batch contents
+    // ref: https://github.com/ggml-org/llama.cpp/pull/18550
+    std::array<ggml_tensor *, 3> inps = {};
+
+    // token embeddings path (ubatch.token != nullptr)
+    inps[0] = build_tok(inp->tokens);

     // vector embeddings path (ubatch.embd != nullptr)
-    {
-        auto & cur = inps[1];
+    inps[1] = inp->embd;
+
+    // mixed path (ubatch.is_mixed()): set_rows the token rows into a copy of the embd rows, with its own inputs as select branches must not share tensors
+    // TODO: use inp->tokens and inp->embd once ggml_build_forward_select allows it
+    const bool has_mixed = llm_arch_supports_mixed_batch(arch) && cparams.ctx_type == LLAMA_CONTEXT_TYPE_DEFAULT;
+    if (has_mixed) {
+        const int64_t n_tok_rows = llm_graph_n_tok_rows(ubatch);
+
+        inp->mixed_tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tok_rows);
+        cb(inp->mixed_tokens, "inp_mixed_tokens", -1);
+        ggml_set_input(inp->mixed_tokens);

-        cur = inp->embd;
+        inp->mixed_slots = ggml_new_tensor_1d(ctx0, GGML_TYPE_I64, n_tok_rows);
+        cb(inp->mixed_slots, "inp_mixed_slots", -1);
+        ggml_set_input(inp->mixed_slots);
+
+        inp->mixed_embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_embd_inp, ubatch.n_tokens);
+        cb(inp->mixed_embd, "inp_mixed_embd", -1);
+        ggml_set_input(inp->mixed_embd);
+
+        // note: set_rows writes into its destination, so it gets a copy of the input
+        inps[2] = ggml_set_rows(ctx0, ggml_dup(ctx0, inp->mixed_embd), build_tok(inp->mixed_tokens), inp->mixed_slots);
     }

     assert(ggml_are_same_shape (inps[0], inps[1]));
     assert(ggml_are_same_stride(inps[0], inps[1]));

-    ggml_tensor * cur = ggml_build_forward_select(gf, inps.data(), inps.size(), ubatch.token ? 0 : 1);
+    const int idx = ubatch.is_mixed() ? 2 : ubatch.token ? 0 : 1;
+
+    ggml_tensor * cur = ggml_build_forward_select(gf, inps.data(), has_mixed ? 3 : 2, idx);

     if (n_embd_inp != n_embd) {
         cur = ggml_view_2d(ctx0, cur, n_embd, n_tokens, cur->nb[1], 0);
@@ -2445,16 +2501,31 @@ ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd) const {

     res->t_inp_embd = cur;

-    // For Granite architecture
     // NOTE: For deepstack models, only apply scale to token inputs (ie text-only input).
     //  Raw embeddings are assumed to be multimodal inputs that should not be scaled.
-    if (hparams.f_embedding_scale != 0.0f && (ubatch.token || hparams.n_deepstack_layers == 0)) {
+    const bool scale_tok_only = hparams.f_embedding_scale != 0.0f && hparams.n_deepstack_layers > 0;
+
+    // For Granite architecture
+    if (hparams.f_embedding_scale != 0.0f && !scale_tok_only) {
         if (!ggml_is_contiguous(cur)) {
             cur = ggml_cont(ctx0, cur);
         }
         cur = ggml_scale(ctx0, cur, hparams.f_embedding_scale);
     }

+    // scale the token rows only, applied after the select so that the graph is the same for any batch contents
+    inp->scale_tok = tok_scale*(scale_tok_only ? hparams.f_embedding_scale : 1.0f);
+    if (inp->scale_tok != 1.0f) {
+        inp->scale_rows = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, 1, ubatch.n_tokens);
+        cb(inp->scale_rows, "inp_scale_rows", -1);
+        ggml_set_input(inp->scale_rows);
+
+        if (!ggml_is_contiguous(cur)) {
+            cur = ggml_cont(ctx0, cur);
+        }
+        cur = ggml_mul(ctx0, cur, inp->scale_rows);
+    }
+
     cb(cur, "embd", -1);

     res->add_input(std::move(inp));
diff --git a/src/llama-graph.h b/src/llama-graph.h
index 366eccbcd..5cf74c997 100644
--- a/src/llama-graph.h
+++ b/src/llama-graph.h
@@ -134,8 +134,14 @@ public:

     bool can_reuse(const llm_graph_params & params) override;

-    ggml_tensor * tokens = nullptr; // I32 [n_batch]
-    ggml_tensor * embd   = nullptr; // F32 [n_embd, n_batch]
+    ggml_tensor * tokens       = nullptr; // I32 [n_batch]
+    ggml_tensor * embd         = nullptr; // F32 [n_embd, n_batch]
+    ggml_tensor * mixed_tokens = nullptr; // I32 [n_tok_rows], mixed path: ids of the token rows
+    ggml_tensor * mixed_slots  = nullptr; // I64 [n_tok_rows], mixed path: batch index of the token rows
+    ggml_tensor * mixed_embd   = nullptr; // F32 [n_embd, n_batch], mixed path: embd rows, token rows are overwritten
+    ggml_tensor * scale_rows   = nullptr; // F32 [1, n_batch], per-row scale: scale_tok for token rows, 1 for embd rows
+
+    float scale_tok = 1.0f;

     const int64_t n_embd = 0;
 };
@@ -823,6 +829,7 @@ struct llm_graph_params {
             ubatch.n_seq_tokens == other.ubatch.n_seq_tokens &&
             ubatch.n_seqs       == other.ubatch.n_seqs &&
             ubatch.n_seqs_unq   == other.ubatch.n_seqs_unq &&
+            ubatch.is_mixed()   == other.ubatch.is_mixed() &&
             (
                 (!ubatch.token && !other.ubatch.token) ||
                 (!ubatch.embd  && !other.ubatch.embd)  ||
@@ -1166,7 +1173,8 @@ struct llm_graph_context {
     // inputs
     //

-    ggml_tensor * build_inp_embd(ggml_tensor * tok_embd) const;
+    // tok_scale: applied to token rows only
+    ggml_tensor * build_inp_embd(ggml_tensor * tok_embd, float tok_scale = 1.0f) const;
     ggml_tensor * build_inp_pos() const;
     ggml_tensor * build_inp_attn_scale() const;
     ggml_tensor * build_inp_out_ids() const;
diff --git a/src/llama-kv-cache-dsv4.cpp b/src/llama-kv-cache-dsv4.cpp
index 4dbcfb8db..5dd0b4591 100644
--- a/src/llama-kv-cache-dsv4.cpp
+++ b/src/llama-kv-cache-dsv4.cpp
@@ -83,6 +83,7 @@ static llama_ubatch dsv4_build_raw_write_ubatch(const llama_ubatch & ubatch) {
     if (!dsv4_ubatch_has_coupled(ubatch)) {
         return ubatch;
     }
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     if (ubatch.embd) {
         throw std::runtime_error("DSV4 coupled embedding ubatches are not supported");
     }
@@ -164,6 +165,7 @@ static llama_ubatch dsv4_build_raw_write_ubatch(const llama_ubatch & ubatch) {
         /*.seq_id_unq   =*/ data->seq_id_unq.data(),
         /*.seq_idx      =*/ data->seq_idx.data(),
         /*.output       =*/ data->output.data(),
+        /*.type         =*/ nullptr,
         /*.decision_order =*/ nullptr,
         /*.data         =*/ data,
     };
diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp
index 1d0554fa1..8891526cd 100644
--- a/src/llama-kv-cache.cpp
+++ b/src/llama-kv-cache.cpp
@@ -1138,7 +1138,9 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch &
                     ext.y = ubatch.pos[i + ubatch.n_tokens];
                 }

-                if (ubatch.token) {
+                const bool is_embd = !ubatch.token || (ubatch.is_mixed() && ubatch.type[i]);
+
+                if (!is_embd) {
                     ext.tok = ubatch.token[i];
                 } else if (hparams.ple_n_heads > 0) {
                     // embd batch (multimodal input) has no token ids, need to pad it with the correct ID for PLE layers
@@ -1862,10 +1864,13 @@ void llama_kv_cache::get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, st

     // an embd (multimodal) ubatch can repeat one position for a whole image, so positions
     // do not encode the token order; resolve its predecessors by ubatch order instead
+    // same for a mixed ubatch
+    const bool by_order = !ubatch.token || ubatch.is_mixed();
+
     std::vector<uint32_t> ord; // index among the ubatch tokens of the same seq
     std::unordered_map<llama_seq_id, std::vector<uint32_t>> seq_idx;

-    if (!ubatch.token) {
+    if (by_order) {
         ord.resize(n_tokens);
         for (uint32_t i = 0; i < n_tokens; ++i) {
             auto & v = seq_idx[ubatch.seq_id[i][0]];
@@ -1883,7 +1888,7 @@ void llama_kv_cache::get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, st
             const llama_pos d = (llama_pos) (n - j);

             llama_pos p;
-            if (!ubatch.token) {
+            if (by_order) {
                 const auto & v = seq_idx[seq_id];
                 const int64_t k = (int64_t) ord[i] - d;
                 // k >= 0: an earlier token of this very ubatch; k < 0: before the chunk
diff --git a/src/models/cogvlm.cpp b/src/models/cogvlm.cpp
index 750f57a39..bc394b3c1 100644
--- a/src/models/cogvlm.cpp
+++ b/src/models/cogvlm.cpp
@@ -70,6 +70,7 @@ llama_model_cogvlm::graph::graph(const llama_model & model, const llm_graph_para
     // check ubatch to see if we have input tokens (text)
     // or an input embedding vector (image)
     bool is_text;
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     if (ubatch.token) {
         is_text = true;
     } else {
diff --git a/src/models/cohere2moe.cpp b/src/models/cohere2moe.cpp
index 7704cbb87..cf2af012d 100644
--- a/src/models/cohere2moe.cpp
+++ b/src/models/cohere2moe.cpp
@@ -317,6 +317,7 @@ llama_model_cohere2moe::graph_mtp::graph_mtp(const llama_model & model, const ll
     // TODO: make static using `ggml_build_forward_select()`
     //       see llm_graph_context::build_inp_embd() for reference
     ggml_tensor * tok_embd;
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     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);
diff --git a/src/models/deepseek2.cpp b/src/models/deepseek2.cpp
index deca86527..6d217ea0c 100644
--- a/src/models/deepseek2.cpp
+++ b/src/models/deepseek2.cpp
@@ -221,6 +221,7 @@ llama_model_deepseek2::graph_mtp::graph_mtp(const llama_model & model, const llm
     ggml_set_input(inp->embd);

     ggml_tensor * tok_embd;
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     if (ubatch.token) {
         ggml_tensor * tok_embd_w = layer.nextn.embed_tokens
                 ? layer.nextn.embed_tokens
diff --git a/src/models/deepseek32.cpp b/src/models/deepseek32.cpp
index 60cc17c49..849c7a9a0 100644
--- a/src/models/deepseek32.cpp
+++ b/src/models/deepseek32.cpp
@@ -535,6 +535,7 @@ llama_model_deepseek32::graph_mtp::graph_mtp(const llama_model & model, const ll
     ggml_set_input(inp->embd);

     ggml_tensor * tok_embd;
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     if (ubatch.token) {
         ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;

diff --git a/src/models/deepseek4.cpp b/src/models/deepseek4.cpp
index a14388725..223336d49 100644
--- a/src/models/deepseek4.cpp
+++ b/src/models/deepseek4.cpp
@@ -1281,6 +1281,7 @@ llama_model_deepseek4::graph::graph(const llama_model & model, const llm_graph_p
         ggml_tensor * exp_probs_b = layer.ffn_exp_probs_b;

         // may apply exp_probs_b_vl is input is from mtmd
+        ASSERT_EMBD_OR_TOKEN(ubatch);
         const bool is_media = ubatch.embd != nullptr;
         if (is_media) {
             if (layer.ffn_exp_probs_b_vl) {
@@ -1366,6 +1367,7 @@ llama_model_deepseek4::graph_mtp::graph_mtp(const llama_model & model, const llm
     GGML_ASSERT(cparams.nextn_layer_offset >= 0 &&
             cparams.nextn_layer_offset < (int) hparams.n_layer_nextn &&
             "nextn_layer_offset out of range [0, n_layer_nextn)");
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     GGML_ASSERT(ubatch.token && "DEEPSEEK4 MTP requires token input");

     const int64_t hc = hparams.dsv4_hc_mult;
diff --git a/src/models/dflash.cpp b/src/models/dflash.cpp
index c448e63f3..c8c2895b6 100644
--- a/src/models/dflash.cpp
+++ b/src/models/dflash.cpp
@@ -604,6 +604,7 @@ llama_model_dflash::graph<false>::graph(const llama_model & model, const llm_gra
     };

     // KV cache injection
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     if (ubatch.embd) {
         auto inp = std::make_unique<llm_graph_input_embd>(n_embd_inp);

@@ -870,6 +871,7 @@ llama_model_dflash::graph_dsv4::graph_dsv4(const llama_model & model, const llm_
     llm_graph_input_attn_k_iswa * inp_attn = build_attn_inp_k_iswa();

     // KV cache injection: fused target features from the encoder
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     if (ubatch.embd) {
         auto inp = std::make_unique<llm_graph_input_embd>(n_embd_inp);

diff --git a/src/models/gemma-embedding.cpp b/src/models/gemma-embedding.cpp
index 6c97883d8..692e19ee9 100644
--- a/src/models/gemma-embedding.cpp
+++ b/src/models/gemma-embedding.cpp
@@ -77,10 +77,8 @@ llama_model_gemma_embedding::graph::graph(const llama_model & model, const llm_g
     ggml_tensor * cur;
     ggml_tensor * inpL;

-    inpL = build_inp_embd(model.tok_embd);
-
     // important: do not normalize weights for raw embeddings input (i.e. encoded image embeddings)
-    inpL = ggml_scale(ctx0, inpL, ubatch.token ? sqrtf(n_embd) : 1.0f);
+    inpL = build_inp_embd(model.tok_embd, sqrtf(n_embd));
     cb(inpL, "inp_scaled", -1);

     // inp_pos - contains the positions
diff --git a/src/models/gemma3.cpp b/src/models/gemma3.cpp
index f99bbaacd..83cb57f99 100644
--- a/src/models/gemma3.cpp
+++ b/src/models/gemma3.cpp
@@ -85,10 +85,8 @@ llama_model_gemma3::graph<iswa>::graph(const llama_model & model, const llm_grap
     ggml_tensor * cur;
     ggml_tensor * inpL;

-    inpL = build_inp_embd(model.tok_embd);
-
     // important: do not normalize weights for raw embeddings input (i.e. encoded image embeddings)
-    inpL = ggml_scale(ctx0, inpL, ubatch.token ? sqrtf(n_embd) : 1.0f);
+    inpL = build_inp_embd(model.tok_embd, sqrtf(n_embd));
     cb(inpL, "inp_scaled", -1);

     // inp_pos - contains the positions
diff --git a/src/models/gemma3n.cpp b/src/models/gemma3n.cpp
index 4d47ddc62..50d30c63f 100644
--- a/src/models/gemma3n.cpp
+++ b/src/models/gemma3n.cpp
@@ -96,10 +96,8 @@ llama_model_gemma3n::graph::graph(const llama_model & model, const llm_graph_par
     ggml_tensor * cur;
     ggml_tensor * inpL;

-    inpL = build_inp_embd(model.tok_embd);
-
     // important: do not normalize weights for raw embeddings input (i.e. encoded image embeddings)
-    inpL = ggml_scale(ctx0, inpL, ubatch.token ? sqrtf(n_embd) : 1.0f);
+    inpL = build_inp_embd(model.tok_embd, sqrtf(n_embd));
     cb(inpL, "inp_scaled", -1);

     // inp_pos - contains the positions
@@ -325,6 +323,7 @@ ggml_tensor * llama_model_gemma3n::graph::build_inp_per_layer() {
     auto inp = std::make_unique<llm_graph_input_embd>(n_embd);
     ggml_tensor * inp_per_layer;
     float tok_embd_scale = sqrtf((float) n_embd_altup);
+    // mixed ubatch: embd rows have token id 0, same padding row as below
     if (ubatch.token) {
         inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, ubatch.n_tokens);
         ggml_set_input(inp->tokens);
diff --git a/src/models/gemma4.cpp b/src/models/gemma4.cpp
index 65fc7623d..38239eba0 100644
--- a/src/models/gemma4.cpp
+++ b/src/models/gemma4.cpp
@@ -159,10 +159,8 @@ llama_model_gemma4::graph::graph(const llama_model & model, const llm_graph_para
     ggml_tensor * cur;
     ggml_tensor * inpL;

-    inpL = build_inp_embd(model.tok_embd);
-
     // important: do not normalize weights for raw embeddings input (i.e. encoded image emdeddings)
-    inpL = ggml_scale(ctx0, inpL, ubatch.token ? sqrtf(n_embd) : 1.0f);
+    inpL = build_inp_embd(model.tok_embd, sqrtf(n_embd));
     cb(inpL, "inp_scaled", -1);

     // inp_pos - contains the positions
@@ -473,6 +471,7 @@ ggml_tensor * llama_model_gemma4::graph::build_inp_per_layer() {

     ggml_tensor * inp_per_layer;
     float tok_embd_scale = sqrtf((float) n_embd_per_layer);
+    // mixed ubatch: embd rows have token id 0, same padding row as below
     if (ubatch.token) {
         inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, ubatch.n_tokens);
         ggml_set_input(inp->tokens);
diff --git a/src/models/glm-dsa.cpp b/src/models/glm-dsa.cpp
index 44d883274..6a5132cf6 100644
--- a/src/models/glm-dsa.cpp
+++ b/src/models/glm-dsa.cpp
@@ -579,6 +579,7 @@ llama_model_glm_dsa::graph_mtp::graph_mtp(const llama_model & model, const llm_g
     ggml_set_input(inp->embd);

     ggml_tensor * tok_embd;
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     if (ubatch.token) {
         ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;

diff --git a/src/models/glm4-moe.cpp b/src/models/glm4-moe.cpp
index d6ae5783c..8cdbe10ad 100644
--- a/src/models/glm4-moe.cpp
+++ b/src/models/glm4-moe.cpp
@@ -158,6 +158,7 @@ llama_model_glm4_moe::graph_mtp::graph_mtp(const llama_model & model, const llm_
     ggml_set_input(inp->embd);

     ggml_tensor * tok_embd;
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     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);
diff --git a/src/models/granite-switch.cpp b/src/models/granite-switch.cpp
index 7c9a901c8..044c8e699 100644
--- a/src/models/granite-switch.cpp
+++ b/src/models/granite-switch.cpp
@@ -149,6 +149,7 @@ public:
 // K dim-0 is +gain for an adapter token, -gain otherwise; the causal softmax then
 // lets a single visible adapter token dominate so the readback recovers its slot.
 void llm_graph_input_switch::set_input(const llama_ubatch * ubatch) {
+    ASSERT_EMBD_OR_TOKEN(*ubatch);
     if (!ubatch->token) {
         return;
     }
@@ -226,6 +227,7 @@ llama_model_granite_switch::graph::graph(
     const auto & smodel = static_cast<const llama_model_granite_switch &>(model);

     // TODO: support raw embedding input (multimodal / pre-embedded tokens) when needed
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     GGML_ASSERT(ubatch.token && "granite-switch requires token input");

     const int64_t n_embd_head = hparams.n_embd_head_v();
diff --git a/src/models/nemotron-h-moe.cpp b/src/models/nemotron-h-moe.cpp
index b4fb25430..b9b4fdcf1 100644
--- a/src/models/nemotron-h-moe.cpp
+++ b/src/models/nemotron-h-moe.cpp
@@ -34,6 +34,7 @@ llama_model_nemotron_h_moe::graph_mtp::graph_mtp(const llama_model & model, cons
     ggml_set_input(inp->embd);

     ggml_tensor * tok_embd;
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     if (ubatch.token) {
         tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens);
     } else {
diff --git a/src/models/qwen35.cpp b/src/models/qwen35.cpp
index d50f067a5..ab98744a5 100644
--- a/src/models/qwen35.cpp
+++ b/src/models/qwen35.cpp
@@ -529,6 +529,7 @@ llama_model_qwen35::graph_mtp::graph_mtp(const llama_model & model, const llm_gr
     // TODO: make static using `ggml_build_forward_select()`
     //       see llm_graph_context::build_inp_embd() for reference
     ggml_tensor * tok_embd;
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     if (ubatch.token) {
         ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;

diff --git a/src/models/qwen35moe.cpp b/src/models/qwen35moe.cpp
index bdf772625..f0f917af7 100644
--- a/src/models/qwen35moe.cpp
+++ b/src/models/qwen35moe.cpp
@@ -579,6 +579,7 @@ llama_model_qwen35moe::graph_mtp::graph_mtp(const llama_model & model, const llm
     // TODO: make static using `ggml_build_forward_select()`
     //       see llm_graph_context::build_inp_embd() for reference
     ggml_tensor * tok_embd;
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     if (ubatch.token) {
         ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;

diff --git a/src/models/qwen3next.cpp b/src/models/qwen3next.cpp
index b63fc9c6a..340ef28f7 100644
--- a/src/models/qwen3next.cpp
+++ b/src/models/qwen3next.cpp
@@ -653,6 +653,7 @@ llama_model_qwen3next::graph_mtp::graph_mtp(const llama_model & model, const llm
     // TODO: make static using `ggml_build_forward_select()`
     //       see llm_graph_context::build_inp_embd() for reference
     ggml_tensor * tok_embd;
+    ASSERT_EMBD_OR_TOKEN(ubatch);
     if (ubatch.token) {
         ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;

diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp
index f4df6a5c2..416c9263b 100644
--- a/src/models/qwen4exp.cpp
+++ b/src/models/qwen4exp.cpp
@@ -1240,7 +1240,8 @@ void llm_graph_input_qwen4exp_ple::set_input(const llama_ubatch * ubatch) {
         ? (llama_token) hparams.ple_image_token_id
         : (llama_token) hparams.ple_eos_token_id;
     auto tok_of = [&](int64_t k) -> llama_token {
-        return ubatch->token ? ubatch->token[k] : img_tok;
+        const bool is_embd = !ubatch->token || (ubatch->is_mixed() && ubatch->type[k]);
+        return is_embd ? img_tok : ubatch->token[k];
     };

     const int64_t n_tokens = ubatch->n_tokens;
diff --git a/tests/test-batch-alloc.cpp b/tests/test-batch-alloc.cpp
index ad186c693..b085917cf 100644
--- a/tests/test-batch-alloc.cpp
+++ b/tests/test-batch-alloc.cpp
@@ -99,6 +99,15 @@ struct batch_builder {
         const llama_pos pos[GGML_MROPE_SECTIONS] = { p, 0, 0, 0 };
         return add_embd(pos, seq_ids, output);
     }
+
+    int32_t add_tok(llama_token id, llama_pos p, llama_seq_id seq_id, bool output) {
+        const int32_t idx = b.add_token(seq_id);
+        GGML_ASSERT(idx >= 0);
+        GGML_ASSERT(b.set_token_id(idx, id));
+        GGML_ASSERT(b.set_token_pos(idx, &p));
+        GGML_ASSERT(b.set_output(idx, output));
+        return idx;
+    }
 };

 static void test_init(testing & t) {
@@ -431,6 +440,125 @@ static void test_content_types(testing & t) {
     });
 }

+static void test_mixed(testing & t) {
+    llama_vocab vocab;
+
+    t.test("rejected_unless_allowed", [&](testing & t) {
+        batch_builder bb(2, nullptr, 4, 1, /*n_vocab*/ 10);
+        bb.add_tok(3, 0, 0, false);
+        bb.add(1, {0}, true);
+
+        llama_batch_allocr ba_default(1);
+        t.assert_true("rejected by default", !ba_default.init(bb.b, vocab, false));
+
+        llama_batch_allocr ba_mixed(1, true);
+        t.assert_true("accepted when allowed", ba_mixed.init(bb.b, vocab, false));
+    });
+
+    t.test("layout_and_split", [&](testing & t) {
+        batch_builder bb(2, nullptr, 4, 1, /*n_vocab*/ 10);
+        bb.add_tok(3, 0, 0, false);
+        bb.add(1, {0}, false);
+        bb.add(2, {0}, false);
+        bb.add_tok(5, 3, 0, true);
+
+        llama_batch_allocr ba(1, true);
+        t.assert_true(ba.init(bb.b, vocab, false));
+
+        const llama_batch & batch = ba.get_batch();
+        t.assert_true(batch.token != nullptr && batch.embd != nullptr);
+
+        const llama_token exp_tok[4]  = { 3, 0, 0, 5 };
+        const float       exp_embd[8] = { 0, 0, 100, 101, 200, 201, 0, 0 };
+        for (int i = 0; i < 4; ++i) {
+            t.assert_equal(exp_tok[i], batch.token[i]);
+        }
+        for (int i = 0; i < 8; ++i) {
+            t.assert_equal(exp_embd[i], batch.embd[i]);
+        }
+
+        llama_ubatch ub0 = ba.split_simple(3);
+        t.assert_equal(3u, ub0.n_tokens);
+        t.assert_true(ub0.is_mixed());
+        const int8_t exp_is_embd[3] = { 0, 1, 1 };
+        for (int i = 0; i < 3; ++i) {
+            t.assert_equal(exp_is_embd[i], ub0.type[i]);
+            t.assert_equal((llama_pos) i, ub0.pos[i]);
+        }
+        t.assert_equal(3, ub0.token[0]);
+        t.assert_equal(0.0f, ub0.embd[0]);
+        t.assert_equal(100.0f, ub0.embd[2]);
+
+        // token rows only: a plain token ubatch
+        llama_ubatch ub1 = ba.split_simple(3);
+        t.assert_equal(1u, ub1.n_tokens);
+        t.assert_true(!ub1.is_mixed());
+        t.assert_true(ub1.embd == nullptr);
+        t.assert_equal(5, ub1.token[0]);
+        t.assert_equal((llama_pos) 3, ub1.pos[0]);
+
+        t.assert_equal(0u, ba.split_simple(3).n_tokens);
+    });
+
+    t.test("rejects_entry_with_both", [&](testing & t) {
+        // token + embd on one entry is the MTP layout
+        batch_builder bb(2, nullptr, 4, 1, /*n_vocab*/ 10);
+        bb.add_tok(3, 0, 0, false);
+        bb.add(1, {0}, false);
+        const int32_t idx = bb.add_tok(4, 2, 0, true);
+        const auto r = bb.row(idx, bb.n_embd);
+        t.assert_true(bb.b.set_token_embd(idx, { r.data(), 1, bb.n_embd }));
+
+        llama_batch_allocr ba(1, true);
+        t.assert_true(!ba.init(bb.b, vocab, false));
+    });
+
+    t.test("mrope_pos_expanded", [&](testing & t) {
+        const uint32_t n_pos = 4;
+        batch_builder bb(2, nullptr, 4, n_pos, /*n_vocab*/ 10);
+
+        bb.add_tok(3, 10, 0, false);
+        const llama_pos pos1[n_pos] = { 11, 5, 7, 0 };
+        bb.add_embd(pos1, {0}, true);
+
+        llama_batch_allocr ba(n_pos, true);
+        t.assert_true(ba.init(bb.b, vocab, false));
+
+        llama_ubatch ub = ba.split_simple(2);
+        const llama_pos expected[8] = { 10, 11, 10, 5, 10, 7, 0, 0 };
+        for (int i = 0; i < 8; ++i) {
+            t.assert_equal(expected[i], ub.pos[i]);
+        }
+    });
+
+    t.test("mrope_rule_follows_first_entry", [&](testing & t) {
+        const uint32_t n_pos = 4;
+
+        mock_memory mem;
+        mem.ranges[0] = {0, 9};
+
+        llama_batch_allocr ba(n_pos, true);
+
+        // token first: must start after the memory
+        {
+            batch_builder bb(2, &mem, 4, n_pos, /*n_vocab*/ 10);
+            bb.add_tok(3, 9, 0, false);
+            const llama_pos pos[n_pos] = { 10, 1, 1, 0 };
+            bb.add_embd(pos, {0}, true);
+            t.assert_true("token overlapping the memory is rejected", !ba.init(bb.b, vocab, false));
+        }
+
+        // embd first: can overlap the memory
+        {
+            batch_builder bb(2, &mem, 4, n_pos, /*n_vocab*/ 10);
+            const llama_pos pos[n_pos] = { 9, 1, 1, 0 };
+            bb.add_embd(pos, {0}, false);
+            bb.add_tok(3, 10, 0, true);
+            t.assert_true("embd overlapping the memory is allowed", ba.init(bb.b, vocab, false));
+        }
+    });
+}
+
 static void test_split(testing & t) {
     llama_vocab vocab;

@@ -1059,6 +1187,7 @@ int main(int argc, char ** argv) {

     t.test("init",           test_init);
     t.test("content_types",  test_content_types);
+    t.test("mixed",          test_mixed);
     t.test("compat",         test_compat);
     t.test("split",          test_split);
     t.test("keep_tail",      test_keep_tail);
diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp
index 93bf47baa..7a76efccc 100644
--- a/tests/test-llama-archs.cpp
+++ b/tests/test-llama-archs.cpp
@@ -535,6 +535,50 @@ static std::vector<float> get_logits(
     return ret;
 }

+// entries [n/4, n/2) are embd rows, decoded either as token/embd/token chunks or as one mixed batch
+// returns the llama_process() error code
+static int32_t get_logits_mixed(
+        llama_model * model, llama_context * lctx, const std::vector<llama_token> & tokens, const std::vector<float> & embd, bool mixed,
+        std::vector<float> & ret) {
+    const uint32_t n_vocab  = llama_vocab_n_tokens(llama_model_get_vocab(model));
+    const uint32_t n_embd   = llama_model_n_embd_inp(model);
+    const uint32_t n_tokens = tokens.size();
+    const uint32_t i_embd_0 = n_tokens/4;
+    const uint32_t i_embd_1 = n_tokens/2;
+
+    const std::vector<uint32_t> bounds = mixed
+        ? std::vector<uint32_t>{0, n_tokens}
+        : std::vector<uint32_t>{0, i_embd_0, i_embd_1, n_tokens};
+
+    llama_memory_clear(llama_get_memory(lctx), true);
+    llama_batch_ext_ptr batch(llama_batch_ext_init(lctx));
+
+    ret.clear();
+    ret.reserve(n_tokens*n_vocab);
+    for (size_t c = 0; c + 1 < bounds.size(); c++) {
+        llama_batch_ext_clear(batch.get());
+        for (uint32_t i = bounds[c]; i < bounds[c + 1]; i++) {
+            const bool is_embd = i >= i_embd_0 && i < i_embd_1;
+            const int32_t idx = is_embd
+                ? llama_batch_ext_add_embd(batch.get(), 0, { embd.data() + (size_t) (i - i_embd_0)*n_embd, 1, n_embd })
+                : llama_batch_ext_add_token(batch.get(), 0, tokens[i]);
+            GGML_ASSERT(idx >= 0);
+            const llama_pos pos[4] = { (llama_pos) i, (llama_pos) i, (llama_pos) i, 0 };
+            llama_batch_ext_set_pos(batch.get(), idx, pos);
+            llama_batch_ext_set_output_logits(batch.get(), idx, true);
+        }
+        const int32_t err = llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+        if (err != 0) {
+            return err;
+        }
+        for (uint32_t i = 0; i < bounds[c + 1] - bounds[c]; i++) {
+            const float * logits_ith = llama_get_logits_ith(lctx, i);
+            ret.insert(ret.end(), logits_ith, logits_ith + n_vocab);
+        }
+    }
+    return 0;
+}
+
 static bool check_causal_attn_toggle(
         llama_model * model, llama_context * lctx, const std::vector<llama_token> & tokens) {
     const uint32_t n_vocab  = llama_vocab_n_tokens(llama_model_get_vocab(model));
@@ -833,15 +877,15 @@ static int test_backends(const std::string & arch_filter, const size_t seed, con
         max_arch_name_length = std::max(max_arch_name_length, strlen(llm_arch_name(arch)));
     }

-    const std::string template_header  = std::string("|%" + std::to_string(max_arch_name_length) + "s|%") + std::to_string(max_device_label_length) + "s|%6s|%15s|%9s|\n";
+    const std::string template_header  = std::string("|%" + std::to_string(max_arch_name_length) + "s|%") + std::to_string(max_device_label_length) + "s|%6s|%15s|%9s|%15s|\n";
     const std::string template_row_cfg = std::string("|%" + std::to_string(max_arch_name_length) + "s|%") + std::to_string(max_device_label_length) + "s|%6s|";
-    const std::string template_row_res = "%15s %10s|%20s|\n";
+    const std::string template_row_res = "%15s %10s|%20s|%15s %10s|\n";

     bool all_ok = true;
     size_t n_tests = 0;
     size_t n_failed = 0;
     common_log_flush(common_log_main());
-    LOG(template_header.c_str(), "Model arch.", "Device", "Config", "NMSE vs. CPU", "Roundtrip");
+    LOG(template_header.c_str(), "Model arch.", "Device", "Config", "NMSE vs. CPU", "Roundtrip", "Mixed batch");
     LOG("|");
     for (size_t i = 0; i < max_arch_name_length; i++) {
         LOG("-");
@@ -850,7 +894,7 @@ static int test_backends(const std::string & arch_filter, const size_t seed, con
     for (size_t i = 0; i < max_device_label_length; i++) {
         LOG("-");
     }
-    LOG("|------|---------------|---------|\n");
+    LOG("|------|---------------|---------|---------------|\n");
     for (const llm_arch & arch : llm_arch_all()) {
         if (arch == LLM_ARCH_UNKNOWN) {
             continue;
@@ -889,7 +933,9 @@ static int test_backends(const std::string & arch_filter, const size_t seed, con
                 std::vector<float> logits_dev;
                 std::string status_nmse      = "\033[1;33mSKIP\033[0m";
                 std::string status_roundtrip = "\033[1;33mSKIP\033[0m";
+                std::string status_mixed     = "\033[1;33mSKIP\033[0m";
                 char nmse_str[12] = {0};
+                char mixed_str[12] = {0};

                 bool skip = !arch_supported(arch) || (dc.split_mode == LLAMA_SPLIT_MODE_TENSOR && dc.devs.empty());
                 bool test_executed = false;
@@ -910,6 +956,43 @@ static int test_backends(const std::string & arch_filter, const size_t seed, con
                             test_ok = false;
                             status_nmse = "\033[1;31mFAIL\033[0m";
                         }
+
+                        // chunked decode matches a single batch only with causal attention over a memory
+                        llama_context * lctx_dev = model_and_ctx_dev.second.get();
+                        if (!encode && llama_get_memory(lctx_dev) != nullptr) {
+                            std::vector<float> embd_mixed((size_t) llama_model_n_embd_inp(model_and_ctx_dev.first.get())*tokens.size()/4);
+                            std::mt19937 gen(seed);
+                            std::normal_distribution<float> dis(0.0f, stdev);
+                            for (float & v : embd_mixed) {
+                                v = dis(gen);
+                            }
+                            std::vector<float> logits_mixed;
+                            std::vector<float> logits_chunks;
+                            if (llm_arch_supports_mixed_batch(arch)) {
+                                if (get_logits_mixed(model_and_ctx_dev.first.get(), lctx_dev, tokens, embd_mixed, false, logits_chunks) != 0 ||
+                                    get_logits_mixed(model_and_ctx_dev.first.get(), lctx_dev, tokens, embd_mixed, true,  logits_mixed)  != 0) {
+                                    throw std::runtime_error("failed to decode mixed batch");
+                                }
+                                const double nmse_mixed = nmse(logits_chunks, logits_mixed);
+                                snprintf(mixed_str, sizeof(mixed_str), "(%.2e)", nmse_mixed);
+                                status_mixed = "\033[1;32mOK\033[0m";
+                                if (nmse_mixed > 1e-4) {
+                                    test_ok = false;
+                                    status_mixed = "\033[1;31mFAIL\033[0m";
+                                }
+                            } else {
+                                // must be rejected as an invalid batch, mute the expected error log
+                                ud.verbosity = LOG_LEVEL_OUTPUT;
+                                const int32_t err = get_logits_mixed(model_and_ctx_cpu.first.get(), model_and_ctx_cpu.second.get(), tokens, embd_mixed, true, logits_mixed);
+                                ud.verbosity = verbosity;
+                                if (err != -1) {
+                                    test_ok = false;
+                                    status_mixed = "\033[1;31mFAIL\033[0m";
+                                }
+                            }
+                        }
+
+                        // runs after the mixed batch check, as it leaves the context with non-causal attention
                         if (!encode && !check_causal_attn_toggle(model_and_ctx_dev.first.get(), model_and_ctx_dev.second.get(), tokens)) {
                             if (test_ok) {
                                 status_nmse = "\033[1;31mFAIL\033[0m (toggle)";
@@ -954,7 +1037,7 @@ static int test_backends(const std::string & arch_filter, const size_t seed, con
                 }

                 // log the results for this test case
-                LOG(template_row_res.c_str(), status_nmse.c_str(), nmse_str, status_roundtrip.c_str());
+                LOG(template_row_res.c_str(), status_nmse.c_str(), nmse_str, status_roundtrip.c_str(), status_mixed.c_str(), mixed_str);
             }
         }
     }