Commit f1ea20621 for llama.cpp

commit f1ea206218210afb913ae2f5d2c51faed35915da
Author: Xuan-Son Nguyen <son@huggingface.co>
Date:   Mon Sep 28 19:52:45 2026 +0200

    batch: migrate speculative, mtmd and server to batch_ext (#29385)

    * adapt common

    * add common_batch

    * wip

    * wip: spec

    * cont

    * common_speculative_process

    * server_batch to use common_batch

    * rm some stale calls

    Assisted-by: Claude Fable 5.1

    * migrate mtmd

    * handle imrope, handle return val of add()/add_embd()

    * add spec zeros vector

    * add warning on zero fill path

diff --git a/common/common.cpp b/common/common.cpp
index 6099f2ecc..6f443f1bf 100644
--- a/common/common.cpp
+++ b/common/common.cpp
@@ -613,34 +613,6 @@ std::string string_from(const struct llama_context * ctx, const std::vector<llam
     return buf.str();
 }

-std::string string_from(const struct llama_context * ctx, const struct llama_batch & batch) {
-    std::stringstream buf;
-
-    buf << "[ ";
-
-    bool first = true;
-    for (int i = 0; i < batch.n_tokens; ++i) {
-        if (!first) {
-            buf << ", ";
-        } else {
-            first = false;
-        }
-
-        auto detokenized = common_token_to_piece(ctx, batch.token[i]);
-
-        buf << "\n"          << std::to_string(i)
-            << ", token '"   << detokenized << "'"
-            << ", pos "      << std::to_string(batch.pos[i])
-            << ", n_seq_id " << std::to_string(batch.n_seq_id[i])
-            << ", seq_id "   << std::to_string(batch.seq_id[i][0])
-            << ", logits "   << std::to_string(batch.logits[i]);
-    }
-
-    buf << " ]";
-
-    return buf.str();
-}
-
 void string_process_escapes(std::string & input) {
     std::size_t input_len = input.length();
     std::size_t output_idx = 0;
@@ -1491,7 +1463,8 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
         }

         if (llama_model_has_encoder(model)) {
-            llama_encode(lctx, llama_batch_get_one(tmp.data(), tmp.size()));
+            common_batch batch = common_batch_get_one(lctx, tmp);
+            llama_process(lctx, LLAMA_PROCESS_TYPE_ENCODE, batch.get());
             llama_token decoder_start_token_id = llama_model_decoder_start_token(model);
             if (decoder_start_token_id == LLAMA_TOKEN_NULL) {
                 decoder_start_token_id = bos;
@@ -1500,7 +1473,9 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
             tmp.push_back(decoder_start_token_id);
         }
         if (llama_model_has_decoder(model)) {
-            llama_decode(lctx, llama_batch_get_one(tmp.data(), std::min(tmp.size(), (size_t) params.n_batch)));
+            tmp.resize(std::min(tmp.size(), (size_t) params.n_batch));
+            common_batch batch = common_batch_get_one(lctx, tmp);
+            llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
         }
         llama_memory_clear(llama_get_memory(lctx), true);
         llama_synchronize(lctx);
@@ -1564,9 +1539,13 @@ common_context_seq_rm_type common_context_can_seq_rm(llama_context * ctx) {
     tmp.push_back(0);
     tmp.push_back(0);

-    int ret = llama_decode(ctx, llama_batch_get_one(tmp.data(), tmp.size()));
+    int ret;
+    {
+        common_batch batch = common_batch_get_one(ctx, tmp);
+        ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+    }
     if (ret != 0) {
-        COM_ERR("llama_decode() failed: %d\n", ret);
+        COM_ERR("llama_process() failed: %d\n", ret);
         res = COMMON_CONTEXT_SEQ_RM_TYPE_NO;
         goto done;
     }
@@ -2153,31 +2132,140 @@ float lr_opt::get_lr(float epoch) const {
 }

 bool common_replay_last_token(struct llama_context * ctx, llama_token last_token, int32_t pos) {
-    llama_batch batch = llama_batch_get_one(&last_token, 1);
-    batch.pos = &pos;
-    if (llama_decode(ctx, batch)) {
+    common_batch batch(ctx);
+    batch.add(last_token, pos, 0, true);
+
+    if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
         LOG_ERR("%s: failed to replay last token\n", __func__);
         return false;
     }
     return true;
 }

-llama_batch_ext_ptr common_batch_ext_get_one(llama_context * ctx, const llama_tokens & tokens) {
-    llama_batch_ext_ptr batch(llama_batch_ext_init(ctx));
+common_batch::common_batch(llama_context * ctx) : batch(llama_batch_ext_init(ctx)) {
+    const auto rope_type = llama_model_rope_type(llama_get_model(ctx));
+    n_pos = rope_type == LLAMA_ROPE_TYPE_MROPE || rope_type == LLAMA_ROPE_TYPE_IMROPE ? GGML_MROPE_SECTIONS : 1;
+}
+
+void common_batch::clear() {
+    tokens.clear();
+    llama_batch_ext_clear(batch.get());
+}
+
+int32_t common_batch::add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output) {
+    const int32_t idx = llama_batch_ext_add_token(batch.get(), seq_id, id);
+    if (idx < 0) {
+        GGML_ABORT("%s: failed to add token %d to the batch (error %d, n_tokens = %d)\n", __func__, id, idx, size());
+    }
+    llama_batch_ext_set_pos(batch.get(), idx, &pos);
+    if (output) {
+        llama_batch_ext_set_output_logits(batch.get(), idx, true);
+    }
+    tokens.push_back({ id, { pos, 0, 0, 0 }, seq_id, output, { nullptr, 0, 0 } });
+    return idx;
+}
+
+bool common_batch::set_output(int32_t idx, bool value) {
+    if (idx < 0 || idx >= (int32_t) tokens.size()) {
+        return false;
+    }
+    tokens[idx].output = value;
+    return llama_batch_ext_set_output_logits(batch.get(), idx, value);
+}
+
+bool common_batch::set_embd(int32_t idx, llama_embd embd) {
+    if (idx < 0 || idx >= (int32_t) tokens.size()) {
+        return false;
+    }
+    if (!llama_batch_ext_set_embd_token(batch.get(), idx, embd)) {
+        return false;
+    }
+    tokens[idx].embd = embd;
+    return true;
+}
+
+int32_t common_batch::add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output) {
+    const int32_t idx = llama_batch_ext_add_embd(batch.get(), seq_id, embd);
+    if (idx < 0) {
+        GGML_ABORT("%s: failed to add embedding to the batch (error %d, n_tokens = %d)\n", __func__, idx, size());
+    }
+    llama_batch_ext_set_pos(batch.get(), idx, pos);
+    if (output) {
+        llama_batch_ext_set_output_logits(batch.get(), idx, true);
+    }
+    token t = { LLAMA_TOKEN_NULL, { 0, 0, 0, 0 }, seq_id, output, embd };
+    for (int32_t j = 0; j < n_pos; ++j) {
+        t.pos[j] = pos[j];
+    }
+    tokens.push_back(t);
+    return idx;
+}
+
+common_batch common_batch_from_llama_batch(llama_context * ctx, const llama_batch & batch) {
+    common_batch res(ctx);
+
+    const bool has_token = batch.token != nullptr;
+    const bool has_embd  = batch.embd  != nullptr;
+
+    const size_t n_embd = llama_model_n_embd_inp(llama_get_model(ctx));
+
+    // positions continue from the memory when none are given
+    auto * mem = llama_get_memory(ctx);
+    std::vector<llama_pos> pos_next(llama_n_seq_max(ctx));
+    for (llama_seq_id s = 0; s < (llama_seq_id) pos_next.size(); ++s) {
+        pos_next[s] = llama_memory_seq_pos_max(mem, s) + 1;
+    }
+
+    for (int32_t i = 0; i < batch.n_tokens; ++i) {
+        const int32_t      n_sid  = batch.n_seq_id ? batch.n_seq_id[i]  : 1;
+        const llama_seq_id seq_id = batch.seq_id   ? batch.seq_id[i][0] : 0;
+
+        llama_pos pos[GGML_MROPE_SECTIONS] = { 0, 0, 0, 0 };
+        if (!batch.pos) {
+            pos[0] = pos_next[seq_id]++;
+        } else if (has_token) {
+            pos[0] = batch.pos[i];
+        } else {
+            // embedding batch: section-major layout pos[j*n_tokens + i]
+            for (int32_t j = 0; j < res.n_pos; ++j) {
+                pos[j] = batch.pos[j * batch.n_tokens + i];
+            }
+        }
+
+        const bool output = batch.logits ? batch.logits[i] != 0 : i == batch.n_tokens - 1;
+
+        const llama_embd embd = { has_embd ? batch.embd + (size_t) i * n_embd : nullptr, 1, n_embd };
+
+        int32_t idx;
+        if (has_token) {
+            idx = res.add(batch.token[i], pos[0], seq_id, output);
+            if (has_embd) {
+                res.set_embd(idx, embd);
+            }
+        } else {
+            idx = res.add_embd(embd, pos, seq_id, output);
+        }
+
+        for (int32_t s = 1; s < n_sid; ++s) {
+            llama_batch_ext_add_seq(res.get(), idx, batch.seq_id[i][s]);
+        }
+    }
+
+    return res;
+}
+
+common_batch common_batch_get_one(llama_context * ctx, const llama_tokens & tokens) {
+    common_batch batch(ctx);

     auto mem = llama_get_memory(ctx);
-    llama_pos pos = mem ? llama_memory_seq_pos_max(mem, 0) + 1 : 0;
+    llama_pos pos = llama_memory_seq_pos_max(mem, 0) + 1; // -1 + 1 == 0 when the memory is empty

     for (size_t i = 0; i < tokens.size(); ++i) {
-        const int32_t idx = llama_batch_ext_add_token(batch.get(), 0, tokens[i]);
-        llama_batch_ext_set_pos(batch.get(), idx, &pos);
+        const bool output = i == tokens.size() - 1;
+        batch.add(tokens[i], pos, 0, output);
         pos++;
     }

-    if (!tokens.empty()) {
-        llama_batch_ext_set_output_logits(batch.get(), (int32_t) tokens.size() - 1, true);
-    }
-
     return batch;
 }

@@ -2205,7 +2293,7 @@ bool common_prompt_batch_decode(
         // memory, so we can't just remove the last token from the memory and replay the last token which
         // is the reason for this logic.
         llama_tokens prefix_tokens(all_tokens.begin() + offset, all_tokens.begin() + offset + n_tokens_before_last);
-        llama_batch_ext_ptr batch_prefix = common_batch_ext_get_one(ctx, prefix_tokens);
+        common_batch batch_prefix = common_batch_get_one(ctx, prefix_tokens);
         if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_prefix.get())) {
             COM_ERR("%s", "failed to eval\n");
             return false;
@@ -2215,10 +2303,8 @@ bool common_prompt_batch_decode(
         llama_state_save_file(ctx, state_path.data(), all_tokens.data(), all_tokens.size());
         COM_INF("saved session before last token to %s, n_new = %zu\n", state_path.data(), all_tokens.size());

-        llama_token last_token = all_tokens.back();
-        llama_batch_ext_ptr batch_last = common_batch_ext_get_one(ctx, { last_token });
-        llama_pos pos = n_past;
-        llama_batch_ext_set_pos(batch_last.get(), 0, &pos);
+        common_batch batch_last(ctx);
+        batch_last.add(all_tokens.back(), n_past, 0, true);

         if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_last.get())) {
             COM_ERR("%s", "failed to eval last token\n");
@@ -2227,7 +2313,7 @@ bool common_prompt_batch_decode(
         n_past++;
     } else {
         llama_tokens new_tokens(all_tokens.begin() + offset, all_tokens.begin() + offset + n_new);
-        llama_batch_ext_ptr batch = common_batch_ext_get_one(ctx, new_tokens);
+        common_batch batch = common_batch_get_one(ctx, new_tokens);
         if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
             COM_ERR("%s", "failed to eval\n");
             return false;
diff --git a/common/common.h b/common/common.h
index 329469410..45b15def7 100644
--- a/common/common.h
+++ b/common/common.h
@@ -8,6 +8,7 @@
 #include "ggml.h"
 #include "llama.h"

+#include <array>
 #include <list>
 #include <set>
 #include <sstream>
@@ -879,7 +880,6 @@ void string_process_escapes(std::string & input);
 std::string string_from(bool value);
 std::string string_from(const std::vector<int> & values);
 std::string string_from(const struct llama_context * ctx, const std::vector<llama_token> & tokens);
-std::string string_from(const struct llama_context * ctx, const struct llama_batch & batch);

 bool glob_match(const std::string & pattern, const std::string & str);

@@ -1039,9 +1039,54 @@ void common_batch_add(
     const std::vector<llama_seq_id> & seq_ids,
                                bool   logits);

+// wrapper around llama_batch_ext that provide getter functions for downstream code
+struct common_batch {
+    struct token {
+        llama_token  id;
+        std::array<llama_pos, GGML_MROPE_SECTIONS> pos; // only pos[0] is used for text tokens
+        llama_seq_id seq_id;
+        bool         output;
+        llama_embd   embd; // non-owning view of the data passed to add_embd()/set_embd(), data == NULL if none
+    };
+
+    std::vector<token> tokens; // mirror of the entries, tokens[i] describes batch index i
+    llama_batch_ext_ptr batch;
+
+    int32_t n_pos = 1; // positions per embedding entry, GGML_MROPE_SECTIONS for MROPE/IMROPE
+
+    common_batch() = default;
+    common_batch(struct llama_context * ctx);
+
+    llama_batch_ext * get() const { return batch.get(); }
+
+    // content type of the batch, all entries carry the same combination
+    bool has_token() const { return !tokens.empty() && tokens[0].id != LLAMA_TOKEN_NULL; }
+    bool has_embd () const { return !tokens.empty() && tokens[0].embd.data != nullptr; }
+
+    void clear();
+
+    // returns the batch index (>= 0), aborts if the entry cannot be added (batch full, invalid token or seq id)
+    int32_t add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output);
+
+    bool set_output(int32_t idx, bool value);
+
+    // 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);
+
+    // add an embedding-only entry (no token id), aborts like add() on failure
+    // pos points to n_pos positions
+    int32_t add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output);
+
+    int32_t size() const { return (int32_t) tokens.size(); }
+};
+
 // create a single-sequence batch from a list of tokens
 // last token always have output_logits set to true
-llama_batch_ext_ptr common_batch_ext_get_one(struct llama_context * ctx, const llama_tokens & tokens);
+common_batch common_batch_get_one(struct llama_context * ctx, const llama_tokens & tokens);
+
+// convert a legacy llama_batch, applying its defaults: seq 0, positions continue from memory, last token is output
+// the embd rows are read at the model input width
+common_batch common_batch_from_llama_batch(struct llama_context * ctx, const llama_batch & batch);

 // decodes a single batch of tokens for a prompt and manages session tokens
 //
diff --git a/common/speculative.cpp b/common/speculative.cpp
index 6fdfa4dc3..82e9e9223 100644
--- a/common/speculative.cpp
+++ b/common/speculative.cpp
@@ -165,7 +165,7 @@ struct common_speculative_impl {

     virtual void begin(llama_seq_id seq_id, const llama_tokens & prompt) = 0;

-    virtual bool process(const llama_batch & batch) = 0;
+    virtual bool process(const common_batch & batch) = 0;

     virtual void draft(common_speculative_draft_params_vec & dparams) = 0;

@@ -179,7 +179,11 @@ struct common_speculative_impl {
 struct common_speculative_impl_draft_simple : public common_speculative_impl {
     common_params_speculative_draft params;

-    llama_batch batch;
+    common_batch batch;
+
+    // zero row at the draft input width, stands in for target embeddings the draft cannot read
+    std::vector<float> zeros;
+    bool zeros_warned = false; // the substitution is reported once

     std::vector<common_sampler_ptr> smpls;

@@ -194,6 +198,8 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
             throw std::runtime_error("draft-simple requires a draft context");
         }

+        zeros.assign(llama_model_n_embd_inp(llama_get_model(ctx_dft)), 0.0f);
+
         SPC_TRC("%s", "adding speculative implementation 'draft-simple'\n");
         SPC_TRC("- n_max=%d, n_min=%d, p_min=%f\n", this->params.n_max, this->params.n_min, this->params.p_min);
         SPC_TRC("- gpu_layers=%d, cache_k=%s, cache_v=%s, ctx_tgt=%s, ctx_dft=%s, devices=[%s]\n",
@@ -204,7 +210,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
                 ctx_dft ? "yes" : "no",
                 common_speculative_get_devices_str(this->params.devices).c_str());

-        batch = llama_batch_init(llama_n_batch(ctx_dft), 0, 1);
+        batch = common_batch(ctx_dft);

         // TODO: optimize or pass from outside?
         // {
@@ -249,21 +255,46 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
         }
     }

-    ~common_speculative_impl_draft_simple() override {
-        llama_batch_free(batch);
-    }
-
     void begin(llama_seq_id /*seq_id*/, const llama_tokens & /*prompt*/) override {
         // noop
     }

-    bool process(const llama_batch & batch) override {
+    bool process(const common_batch & batch_in) override {
         auto * ctx_dft = params.ctx_dft;

-        llama_batch batch_dft = batch;
-        batch_dft.logits = nullptr;
+        // copy the entries to a batch owned by the draft context, only the last token is output
+        batch.clear();
+        const int32_t n_tokens = batch_in.size();
+        for (int32_t k = 0; k < n_tokens; ++k) {
+            const auto & t = batch_in.tokens[k];
+            const bool output = k == n_tokens - 1;
+            if (t.id != LLAMA_TOKEN_NULL) {
+                const int32_t idx = batch.add(t.id, t.pos[0], t.seq_id, output);
+                if (t.embd.data) {
+                    batch.set_embd(idx, t.embd);
+                }
+            } else {
+                // mtmd input is projected by the target encoder, a draft with a different width cannot read it
+                // it gets zeros instead, keeping its positions contiguous
+                // ref: https://github.com/ggml-org/llama.cpp/pull/29385#discussion_r4124743243
+                const size_t n_embd = t.embd.n_rows * t.embd.n_embd;
+                const bool same_width = n_embd == zeros.size();
+                if (!same_width && !zeros_warned) {
+                    SPC_WRN("target embeddings of size %zu do not fit the draft input width %zu, "
+                            "the draft receives zero rows for them and drafts after multimodal input will be poor\n",
+                            n_embd, zeros.size());
+                    zeros_warned = true;
+                }
+                const llama_embd embd = same_width ? t.embd : llama_embd{ zeros.data(), 1, zeros.size() };
+                batch.add_embd(embd, t.pos.data(), t.seq_id, output);
+            }
+        }

-        const int ret = llama_decode(ctx_dft, batch_dft);
+        if (batch.size() == 0) {
+            return true;
+        }
+
+        const int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());

         if (ret != 0) {
             SPC_ERR("failed to decode draft batch, ret = %d\n", ret);
@@ -277,7 +308,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
     void draft(common_speculative_draft_params_vec & dparams) override {
         auto & ctx_dft = params.ctx_dft;

-        common_batch_clear(batch);
+        batch.clear();

         // keep track of which sequences are still drafting
         int n_drafting = 0;
@@ -294,12 +325,12 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
             drafting[seq_id] = true;
             common_sampler_reset(smpls[seq_id].get());

-            common_batch_add(batch, dp.id_last, dp.pos0, { seq_id }, true);
+            batch.add(dp.id_last, dp.pos0, seq_id, true);
         }

-        int ret = llama_decode(ctx_dft, batch);
+        int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
         if (ret != 0) {
-            SPC_ERR("llama_decode returned %d\n", ret);
+            SPC_ERR("llama_process returned %d\n", ret);
             return;
         }

@@ -308,7 +339,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
         while (n_drafting > 0) {
             int i_batch = 0;

-            common_batch_clear(batch);
+            batch.clear();

             for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
                 if (!drafting[seq_id]) {
@@ -353,17 +384,17 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
                     continue;
                 }

-                common_batch_add(batch, id, dp.pos0 + i + 1, { seq_id }, true);
+                batch.add(id, dp.pos0 + i + 1, seq_id, true);
             }

-            if (batch.n_tokens == 0) {
+            if (batch.size() == 0) {
                 break;
             }

             // evaluate the drafted tokens on the draft model
-            ret = llama_decode(ctx_dft, batch);
+            ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
             if (ret != 0) {
-                SPC_ERR("llama_decode[%d] returned %d\n", i, ret);
+                SPC_ERR("llama_process[%d] returned %d\n", i, ret);
                 break;
             }

@@ -423,7 +454,8 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
 //      encoder+decoder on n_accepted+1 rows).
 struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
     common_params_speculative_draft params;
-    llama_batch batch;
+    common_batch batch;     // decoder input, (token, g_embd) pairs
+    common_batch batch_enc; // encoder input, built from the extracted target features

     std::vector<common_sampler_ptr> smpls;

@@ -477,11 +509,8 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
         n_embd_enc = (int32_t) target_layer_ids_n * n_embd_tgt;
         n_layer_tgt = llama_model_n_layer(model_tgt);

-        const int32_t n_b = (int32_t) llama_n_batch(ctx_dft);
-        batch = llama_batch_init(/*n_tokens=*/ n_b, /*embd=*/ n_embd_dec, /*n_seq_max=*/ 1);
-        // llama_batch_init allocates only one of token/embd; eagle3 decoder needs both.
-        // TODO: fix, how to call without malloc
-        batch.token = (llama_token *) malloc(sizeof(llama_token) * n_b);
+        batch     = common_batch(ctx_dft);
+        batch_enc = common_batch(ctx_dft);

         smpls.resize(n_seq);
         for (auto & s : smpls) {
@@ -543,12 +572,6 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
             llama_sampler_free(backend_chains[seq_id]);
         }
         backend_chains.clear();
-
-        if (batch.token != nullptr) {
-            free(batch.token);
-            batch.token = nullptr;
-        }
-        llama_batch_free(batch);
     }

     void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {
@@ -567,16 +590,16 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
         }
     }

-    bool process(const llama_batch & batch_in) override {
-        if (batch_in.n_tokens <= 0) {
+    bool process(const common_batch & batch_in) override {
+        if (batch_in.size() <= 0) {
             return true;
         }

-        if (batch_in.token == nullptr || batch_in.embd != nullptr) {
+        if (!batch_in.has_token() || batch_in.has_embd()) {
             return true;
         }

-        const int32_t n_tokens = batch_in.n_tokens;
+        const int32_t n_tokens = batch_in.size();

         // i_batch_beg[seq] / i_batch_end[seq]: inclusive batch indices of this seq's
         // first/last token in batch_in. Assumes per-seq tokens are contiguous within
@@ -584,8 +607,7 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
         std::vector<int32_t> i_batch_beg(n_seq, -1);
         std::vector<int32_t> i_batch_end(n_seq, -1);
         for (int k = 0; k < n_tokens; ++k) {
-            GGML_ASSERT(batch_in.n_seq_id[k] == 1);
-            const llama_seq_id seq_id = batch_in.seq_id[k][0];
+            const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
             if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {
                 continue;
             }
@@ -619,24 +641,23 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {

         g_embd_buf.resize((size_t) n_tokens * n_embd_dec);

-        // llama_encode() requires the full encoder batch to fit in n_ubatch.
+        // llama_process() requires the full encoder batch to fit in n_ubatch.
         // Allow batch > ubatch: eagle3's per-token encoder can be chunked safely.
         const int32_t n_ubatch_dft = (int32_t) llama_n_ubatch(ctx_dft);
         for (int32_t i = 0; i < n_tokens; i += n_ubatch_dft) {
             const int32_t n_chunk = std::min(n_ubatch_dft, n_tokens - i);

-            llama_batch enc_batch = {
-                /*.n_tokens =*/ n_chunk,
-                /*.token    =*/ nullptr,
-                /*.embd     =*/ features_buf.data() + (size_t) i * n_embd_enc,
-                /*.pos      =*/ nullptr,
-                /*.n_seq_id =*/ nullptr,
-                /*.seq_id   =*/ nullptr,
-                /*.logits   =*/ nullptr,
-            };
-            const int32_t rc = llama_encode(ctx_dft, enc_batch);
+            // the per-token encoder does not use positions, generate placeholder ones from the memory state
+            batch_enc.clear();
+            llama_pos pos = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), 0) + 1;
+            for (int32_t j = 0; j < n_chunk; ++j) {
+                batch_enc.add_embd({ features_buf.data() + (size_t) (i + j) * n_embd_enc, 1, (size_t) n_embd_enc }, &pos, 0, true);
+                pos++;
+            }
+
+            const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_ENCODE, batch_enc.get());
             if (rc != 0) {
-                SPC_ERR("llama_encode(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
+                SPC_ERR("llama_process(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
                         rc, (int) n_chunk, (int) i);
                 return false;
             }
@@ -664,7 +685,7 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
         //       deferred boundary, completed by the next process() or draft() call.
         //   (c) refresh deferred state — stash this ubatch's full g_embd into verify_g,
         //       update pending_g_last / pending_pos_last to the last row.
-        common_batch_clear(batch);
+        batch.clear();

         for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
             const int32_t beg = i_batch_beg[seq_id];
@@ -679,36 +700,34 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
             //   2) pending_pos_last + 1 == pos[beg]
             //   3) pending_pos_last > dft_pos_max // TODO: is this check needed?
             const llama_pos pending_pos = pending_pos_last[seq_id];
-            if (pending_pos >= 0 && pending_pos + 1 == batch_in.pos[beg]) {
+            if (pending_pos >= 0 && pending_pos + 1 == batch_in.tokens[beg].pos[0]) {
                 const llama_pos dft_pos_max = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), seq_id);
                 if (pending_pos > dft_pos_max) {
-                    common_batch_add(batch, batch_in.token[beg], pending_pos, { seq_id }, /*logits=*/ false);
-                    std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,
-                                pending_g_last[seq_id].data(), row_bytes);
+                    const int32_t idx = batch.add(batch_in.tokens[beg].id, pending_pos, seq_id, /*output=*/ false);
+                    batch.set_embd(idx, { pending_g_last[seq_id].data(), 1, (size_t) n_embd_dec });
                 }
             }

             for (int32_t k = beg; k < end; ++k) {
-                common_batch_add(batch, batch_in.token[k + 1], batch_in.pos[k], { seq_id }, /*logits=*/ false);
-                std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,
-                            g_embd + (size_t) k * n_embd_dec, row_bytes);
+                const int32_t idx = batch.add(batch_in.tokens[k + 1].id, batch_in.tokens[k].pos[0], seq_id, /*output=*/ false);
+                batch.set_embd(idx, { g_embd + (size_t) k * n_embd_dec, 1, (size_t) n_embd_dec });
             }

             // refresh deferred state
             const int32_t n_rows = end - beg + 1;
-            verify_pos_first[seq_id] = batch_in.pos[beg];
-            pending_pos_last[seq_id] = batch_in.pos[end];
+            verify_pos_first[seq_id] = batch_in.tokens[beg].pos[0];
+            pending_pos_last[seq_id] = batch_in.tokens[end].pos[0];
             verify_g_rows[seq_id]    = n_rows;
             verify_g[seq_id].resize((size_t) n_rows * n_embd_dec, 0.0f);
             std::memcpy(verify_g[seq_id].data(),       g_embd + (size_t) beg * n_embd_dec, row_bytes * n_rows);
             std::memcpy(pending_g_last[seq_id].data(), g_embd + (size_t) end * n_embd_dec, row_bytes);
         }

-        if (batch.n_tokens > 0) {
-            const int32_t rc = llama_decode(ctx_dft, batch);
+        if (batch.size() > 0) {
+            const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
             if (rc != 0) {
-                SPC_ERR("llama_decode(ctx_dft) failed rc=%d (n_tokens=%d, ubatch_pos[0]=%d)\n",
-                        rc, (int) batch.n_tokens, (int) batch_in.pos[0]);
+                SPC_ERR("llama_process(ctx_dft) failed rc=%d (n_tokens=%d, ubatch_pos[0]=%d)\n",
+                        rc, (int) batch.size(), (int) batch_in.tokens[0].pos[0]);
                 return false;
             }
         }
@@ -719,14 +738,12 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
     void draft(common_speculative_draft_params_vec & dparams) override {
         auto & ctx_dft = params.ctx_dft;

-        common_batch_clear(batch);
+        batch.clear();

         // keep track of which sequences are still drafting
         int n_drafting = 0;
         std::vector<bool> drafting(n_seq);

-        const size_t row_bytes = (size_t) n_embd_dec * sizeof(float);
-
         // Complete the deferred boundary pair (dp.id_last, pending_g_last) at memory
         // pos pending_pos_last. dp.id_last is target's freshest sample (= corrected
         // token after verify, or first generated token after prefill), matching the
@@ -747,19 +764,17 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {

             llama_memory_seq_rm(llama_get_memory(ctx_dft), seq_id, pending_pos_last[seq_id], -1);

-            common_batch_add(batch, dp.id_last, pending_pos_last[seq_id], { seq_id }, true);
-            std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,
-                        pending_g_last[seq_id].data(),
-                        row_bytes);
+            const int32_t idx = batch.add(dp.id_last, pending_pos_last[seq_id], seq_id, true);
+            batch.set_embd(idx, { pending_g_last[seq_id].data(), 1, (size_t) n_embd_dec });
         }

-        if (batch.n_tokens == 0) {
+        if (batch.size() == 0) {
             return;
         }

-        int ret = llama_decode(ctx_dft, batch);
+        int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
         if (ret != 0) {
-            SPC_ERR("llama_decode returned %d\n", ret);
+            SPC_ERR("llama_process returned %d\n", ret);
             return;
         }

@@ -768,7 +783,7 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
         while (n_drafting > 0) {
             int i_batch = 0;

-            common_batch_clear(batch);
+            batch.clear();

             for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
                 if (!drafting[seq_id]) {
@@ -814,17 +829,17 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
                     continue;
                 }

-                common_batch_add(batch, id, pending_pos_last[seq_id] + (i + 1), { seq_id }, true);
-                std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec, prenorm, row_bytes);
+                const int32_t idx = batch.add(id, pending_pos_last[seq_id] + (i + 1), seq_id, true);
+                batch.set_embd(idx, { prenorm, 1, (size_t) n_embd_dec });
             }

-            if (batch.n_tokens == 0) {
+            if (batch.size() == 0) {
                 break;
             }

-            ret = llama_decode(ctx_dft, batch);
+            ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
             if (ret != 0) {
-                SPC_ERR("llama_decode[%d] returned %d\n", i, ret);
+                SPC_ERR("llama_process[%d] returned %d\n", i, ret);
                 break;
             }

@@ -908,8 +923,10 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
 struct common_speculative_impl_draft_dflash : public common_speculative_impl {
     common_params_speculative_draft params;

-    llama_batch batch;        // noise tokens
-    llama_batch batch_inject; // target features for KV cache injection
+    common_batch batch;        // noise tokens
+    common_batch batch_inject; // target features for KV cache injection
+
+    std::vector<float> features_buf; // [n_chunk, n_embd_enc] gathered target features

     std::vector<common_sampler_ptr> smpls;

@@ -1005,15 +1022,11 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
         }
         this->n_max = this->params.n_max;

-        batch        = llama_batch_init(llama_n_batch(ctx_dft), 0,          n_seq);
-        batch_inject = llama_batch_init(llama_n_ubatch(ctx_dft), n_embd_enc, n_seq);
+        batch        = common_batch(ctx_dft);
+        batch_inject = common_batch(ctx_dft);

-        // embd batches on an M-RoPE draft need 4 position rows per token
+        // embd batches on an M-RoPE draft carry 4 position rows per token
         is_mrope = llama_model_rope_type(model_dft) == LLAMA_ROPE_TYPE_MROPE;
-        if (is_mrope) {
-            free(batch_inject.pos);
-            batch_inject.pos = (llama_pos *) malloc(sizeof(llama_pos) * 4 * llama_n_batch(ctx_dft));
-        }

         smpls.resize(n_seq);
         for (auto & s : smpls) {
@@ -1062,9 +1075,6 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
             llama_sampler_free(backend_chains[seq_id]);
         }
         backend_chains.clear();
-
-        llama_batch_free(batch);
-        llama_batch_free(batch_inject);
     }

     void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {
@@ -1085,8 +1095,8 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
         }
     }

-    bool process(const llama_batch & batch_in) override {
-        if (batch_in.n_tokens <= 0) {
+    bool process(const common_batch & batch_in) override {
+        if (batch_in.size() <= 0) {
             return true;
         }

@@ -1094,20 +1104,19 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
         // produce the target-layer features used to seed the draft KV cache, so
         // embeddings are injected too, except the pinned ones skipped below.
         // TODO: revisit after https://github.com/ggml-org/llama.cpp/pull/24669 is merged
-        const bool has_tokens     = batch_in.token != nullptr;
-        const bool has_embeddings = batch_in.embd  != nullptr;
+        const bool has_tokens     = batch_in.has_token();
+        const bool has_embeddings = batch_in.has_embd();
         if (has_tokens == has_embeddings) {
             return true;
         }

-        const int32_t n_tokens = batch_in.n_tokens;
+        const int32_t n_tokens = batch_in.size();

         // per-seq inclusive batch range (assumes each seq's tokens are contiguous in the batch)
         std::vector<int32_t> i_batch_beg(n_seq, -1);
         std::vector<int32_t> i_batch_end(n_seq, -1);
         for (int32_t k = 0; k < n_tokens; ++k) {
-            GGML_ASSERT(batch_in.n_seq_id[k] == 1);
-            const llama_seq_id seq_id = batch_in.seq_id[k][0];
+            const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
             if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {
                 continue;
             }
@@ -1130,7 +1139,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {

             // an M-RoPE image pins all its rows to one position, so a windowed draft
             // cache cannot free cells for it - skip it, the draft can jump over the gap
-            const bool pos_pinned = batch_in.pos[i_batch_beg[seq_id]] == batch_in.pos[i_batch_end[seq_id]];
+            const bool pos_pinned = batch_in.tokens[i_batch_beg[seq_id]].pos[0] == batch_in.tokens[i_batch_end[seq_id]].pos[0];
             if (has_embeddings && n_rows > 1 && pos_pinned) {
                 continue;
             }
@@ -1140,34 +1149,28 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {

                 // gather target features per extract layer; the fused decode encodes and
                 // injects them into the K/V cache at the target positions
-                batch_inject.n_tokens = n_chunk;
+                features_buf.resize((size_t) n_chunk * n_embd_enc);
                 for (uint32_t k = 0; k < target_layer_ids_n; ++k) {
                     const float * layer = llama_get_embeddings_layer_inp(ctx_tgt, (uint32_t) target_layer_ids[k]);
                     if (!layer) {
                         GGML_ABORT("DFlash: target layer %d input not extracted.", target_layer_ids[k]);
                     }
                     for (int32_t i = 0; i < n_chunk; ++i) {
-                        float       * dst = batch_inject.embd + (size_t) i * n_embd_enc + k * (size_t) n_embd_tgt;
+                        float       * dst = features_buf.data() + (size_t) i * n_embd_enc + k * (size_t) n_embd_tgt;
                         const float * src = layer + (size_t) (i_batch_beg[seq_id] + offset + i) * n_embd_tgt;
                         std::memcpy(dst, src, (size_t) n_embd_tgt * sizeof(float));
                     }
                 }

+                batch_inject.clear();
                 for (int32_t i = 0; i < n_chunk; ++i) {
-                    const llama_pos p = batch_in.pos[i_batch_beg[seq_id] + offset + i];
-                    batch_inject.pos[i] = p;
-                    if (is_mrope) {
-                        batch_inject.pos[1 * n_chunk + i] = p;
-                        batch_inject.pos[2 * n_chunk + i] = p;
-                        batch_inject.pos[3 * n_chunk + i] = 0;
-                    }
-                    batch_inject.n_seq_id[i]  = 1;
-                    batch_inject.seq_id[i][0] = seq_id;
-                    batch_inject.logits[i]    = false;
+                    const llama_pos p = batch_in.tokens[i_batch_beg[seq_id] + offset + i].pos[0];
+                    const llama_pos pos_arr[4] = { p, p, p, 0 };
+                    batch_inject.add_embd({ features_buf.data() + (size_t) i * n_embd_enc, 1, (size_t) n_embd_enc }, pos_arr, seq_id, false);
                 }
-                const int32_t rc = llama_decode(ctx_dft, batch_inject);
+                const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch_inject.get());
                 if (rc != 0) {
-                    LOG_ERR("%s: llama_decode(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
+                    LOG_ERR("%s: llama_process(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
                             __func__, rc, (int) n_chunk, (int) offset);
                     return false;
                 }
@@ -1180,7 +1183,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
     void draft(common_speculative_draft_params_vec & dparams) override {
         auto & ctx_dft = params.ctx_dft;

-        common_batch_clear(batch);
+        batch.clear();

         // build one batch holding every drafting sequence's noise block into a single decode)
         // record where each block starts and its size
@@ -1200,21 +1203,21 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
             const int32_t n_draft = params.n_max;

             const int32_t n_block_tokens = n_draft + (is_dspark && sample_from_anchor ? 0 : 1);
-            i_block_beg[seq_id] = batch.n_tokens;
+            i_block_beg[seq_id] = batch.size();
             n_block    [seq_id] = n_block_tokens;
             for (int32_t i = 0; i < n_block_tokens; ++i) {
-                common_batch_add(batch, i == 0 ? dp.id_last : mask_token_id, n + i, { seq_id }, !is_dflash2);
+                batch.add(i == 0 ? dp.id_last : mask_token_id, n + i, seq_id, !is_dflash2);
             }
         }

-        if (batch.n_tokens == 0) {
+        if (batch.size() == 0) {
             return;
         }

         // decode all sequence's noise block in a single batch
-        int ret = llama_decode(ctx_dft, batch);
+        int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
         if (ret != 0) {
-            LOG_WRN("%s: llama_decode returned %d\n", __func__, ret);
+            LOG_WRN("%s: llama_process returned %d\n", __func__, ret);
             return;
         }

@@ -1328,7 +1331,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
 struct common_speculative_impl_draft_mtp : public common_speculative_impl {
     common_params_speculative_draft params; // reuses the draft-model params slot (ctx_tgt/ctx_dft)

-    llama_batch batch;
+    common_batch batch;

     std::vector<common_sampler_ptr> smpls;

@@ -1384,11 +1387,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
                 ctx_dft ? "yes" : "no",
                 common_speculative_get_devices_str(this->params.devices).c_str());

-        const int32_t n_b = (int32_t) llama_n_batch(ctx_dft);
-        batch = llama_batch_init(/*n_tokens=*/ n_b, /*embd=*/ n_embd, /*n_seq_max=*/ 1);
-        // llama_batch_init allocates only one of token/embd; MTP needs both.
-        // TODO: fix, how to call without malloc
-        batch.token = (llama_token *) malloc(sizeof(llama_token) * n_b);
+        batch = common_batch(ctx_dft);

         smpls.resize(n_seq);
         for (auto & s : smpls) {
@@ -1453,12 +1452,6 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
             llama_sampler_free(backend_chains[seq_id]);
         }
         backend_chains.clear();
-
-        if (batch.token != nullptr) {
-            free(batch.token);
-            batch.token = nullptr;
-        }
-        llama_batch_free(batch);
     }

     void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {
@@ -1473,23 +1466,23 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
         if (pos_max < N - 1 && !is_mem_shared) {
             SPC_WRN("ctx_dft pos_max=%d < N-1=%d - "
                     "process() hook may not have run on every prefill ubatch "
-                    "(need_embd / logits=1 on every prompt position?). "
+                    "(need_embd / output flag on every prompt position?). "
                     "Drafts may degrade.\n",
                     (int) pos_max, N - 1);
         }
     }

-    bool process(const llama_batch & batch_in) override {
-        if (batch_in.n_tokens <= 0) {
+    bool process(const common_batch & batch_in) override {
+        if (batch_in.size() <= 0) {
             return true;
         }

         // TODO: how to make it work with vision tokens?
-        if (batch_in.token == nullptr || batch_in.embd != nullptr) {
+        if (!batch_in.has_token() || batch_in.has_embd()) {
             return true;
         }

-        const int32_t n_tokens = batch_in.n_tokens;
+        const int32_t n_tokens = batch_in.size();

         // remember the first and last batch index for each sequence
         std::fill(i_batch_beg.begin(), i_batch_beg.end(), -1);
@@ -1497,9 +1490,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {

         for (int k = 0; k < n_tokens; ++k) {
             for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
-                GGML_ASSERT(batch_in.n_seq_id[k] == 1);
-
-                if (batch_in.seq_id[k][0] == seq_id) {
+                if (batch_in.tokens[k].seq_id == seq_id) {
                     i_batch_end[seq_id] = k;
                     if (i_batch_beg[seq_id] < 0) {
                         i_batch_beg[seq_id] = k;
@@ -1515,33 +1506,26 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {

         // if kv is shared with target (e.g Gemma4), then we can skip this catch-up decode
         if (!is_mem_shared) {
-            common_batch_clear(batch);
-
-            for (int k = 0; k < n_tokens; ++k) {
-                common_batch_add(batch, batch_in.token[k], batch_in.pos[k], { batch_in.seq_id[k][0] }, 0);
-            }
+            batch.clear();

-            // shift the tgt embeddings to the right by one position
+            // pair each token with the tgt embedding shifted right by one position, and
+            // the first token of each sequence with the pending embedding from a previous run
             // assumes that the tokens in the batch are sequential for each sequence
             // i.e. we cannot have seq_id like this: [0, 0, 0, 1, 1, 0, 1, 1]
             //                                                       ^--- this is a problem
             // TODO:this is generally true, but would be nice to assert it
-            {
-                const float * h_tgt = llama_get_embeddings_nextn(ctx_tgt);
-                std::memcpy(batch.embd + (size_t) 1 * n_embd, h_tgt, row_bytes * (n_tokens-1));
-            }
+            const float * h_tgt = llama_get_embeddings_nextn(ctx_tgt);

-            // fill the pending embeddings from a previous run
-            auto set_h = [&](int idx, const float * h_row) {
-                std::memcpy(batch.embd + (size_t) idx * n_embd, h_row, row_bytes);
-            };
+            for (int k = 0; k < n_tokens; ++k) {
+                const llama_seq_id seq_id = batch_in.tokens[k].seq_id;

-            for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
-                if (i_batch_beg[seq_id] < 0) {
-                    continue;
-                }
+                const int32_t idx = batch.add(batch_in.tokens[k].id, batch_in.tokens[k].pos[0], 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;

-                set_h(i_batch_beg[seq_id], pending_h[seq_id].data());
+                batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
             }

             auto * mem_dft = llama_get_memory(ctx_dft);
@@ -1554,15 +1538,15 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
                         if (i_batch_beg[seq_id] < 0) {
                             continue;
                         }
-                        llama_memory_seq_rm(mem_dft, seq_id, batch_in.pos[i_batch_beg[seq_id]], -1);
+                        llama_memory_seq_rm(mem_dft, seq_id, batch_in.tokens[i_batch_beg[seq_id]].pos[0], -1);
                     }
                     llama_set_nextn_layer_offset(ctx_dft, head);
                 }

-                const int32_t rc = llama_decode(ctx_dft, batch);
+                const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
                 if (rc != 0) {
-                    SPC_ERR("llama_decode(ctx_dft) head=%d failed rc=%d (pos=%d)\n",
-                            head, (int) rc, (int) batch_in.pos[0]);
+                    SPC_ERR("llama_process(ctx_dft) head=%d failed rc=%d (pos=%d)\n",
+                            head, (int) rc, (int) batch_in.tokens[0].pos[0]);
                     ok = false;
                     break;
                 }
@@ -1600,14 +1584,12 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
     void draft(common_speculative_draft_params_vec & dparams) override {
         auto & ctx_dft = params.ctx_dft;

-        common_batch_clear(batch);
+        batch.clear();

         // keep track of which sequences are still drafting
         int n_drafting = 0;
         std::vector<bool> drafting(n_seq);

-        const size_t row_bytes = (size_t) n_embd * sizeof(float);
-
         for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
             auto & dp = dparams[seq_id];

@@ -1619,10 +1601,10 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
             drafting[seq_id] = true;
             common_sampler_reset(smpls[seq_id].get());

-            common_batch_add(batch, dp.id_last, dp.pos0, { seq_id }, true);
-            std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, pending_h[seq_id].data(), row_bytes);
+            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 });

-            i_last[seq_id] = batch.n_tokens - 1;
+            i_last[seq_id] = idx;

             if (chain_heads) {
                 chain_h[seq_id].assign(pending_h[seq_id].begin(), pending_h[seq_id].end());
@@ -1648,16 +1630,16 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
                 llama_set_nextn_layer_offset(ctx_dft, i);
             }

-            int ret = llama_decode(ctx_dft, batch);
+            int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
             if (ret != 0) {
-                SPC_ERR("llama_decode[%d] returned %d\n", i, ret);
+                SPC_ERR("llama_process[%d] returned %d\n", i, ret);
                 break;
             }

             // rebuild the batch for the next step: the growing-KV paths re-add only the
             // new token (the KV already holds the prefix), while chained heads re-add the
             // whole prefix at the next head. dropped sequences are simply not re-added.
-            common_batch_clear(batch);
+            batch.clear();

             for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
                 if (!drafting[seq_id]) {
@@ -1708,24 +1690,24 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
                     const int n_rows = (int) result.size() + 1; // id_last + tokens drafted so far
                     for (int t = 0; t < n_rows; ++t) {
                         const llama_token tok = (t == 0) ? dp.id_last : result[t - 1];
-                        common_batch_add(batch, tok, dp.pos0 + t, { seq_id }, t == n_rows - 1);
-                        std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd,
-                                    chain_h[seq_id].data() + (size_t) t * n_embd, row_bytes);
+                        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 });
+                        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
-                    common_batch_add(batch, id, dp.pos0, { seq_id }, true);
-                    std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, h_row, row_bytes);
+                    const int32_t idx = batch.add(id, dp.pos0, seq_id, true);
+                    batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
+                    i_last[seq_id] = idx;
                 } else {
-                    common_batch_add(batch, id, dp.pos0 + i + 1, { seq_id }, true);
-                    std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, h_row, row_bytes);
+                    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 });
+                    i_last[seq_id] = idx;
                 }
-
-                i_last[seq_id] = batch.n_tokens - 1;
             }

-            if (batch.n_tokens == 0) {
+            if (batch.size() == 0) {
                 break;
             }

@@ -1787,7 +1769,7 @@ struct common_speculative_impl_ngram_simple : public common_speculative_impl {
         // noop
     }

-    bool process(const llama_batch & /*batch*/) override {
+    bool process(const common_batch & /*batch*/) override {
         // TODO: implement
         return true;
     }
@@ -1835,7 +1817,7 @@ struct common_speculative_impl_ngram_map_k : public common_speculative_impl {
         common_ngram_map_begin(config[seq_id], prompt);
     }

-    bool process(const llama_batch & /*batch*/) override {
+    bool process(const common_batch & /*batch*/) override {
         // TODO: implement
         return true;
     }
@@ -1993,7 +1975,7 @@ struct common_speculative_impl_ngram_mod : public common_speculative_impl {
         sinfo.n_draft_last = result.size();
     }

-    bool process(const llama_batch & /*batch*/) override {
+    bool process(const common_batch & /*batch*/) override {
         // TODO: implement
         return true;
     }
@@ -2155,7 +2137,7 @@ struct common_speculative_impl_ngram_cache : public common_speculative_impl {
         }
     }

-    bool process(const llama_batch & /*batch*/) override {
+    bool process(const common_batch & /*batch*/) override {
         // TODO: implement
         return true;
     }
@@ -2181,6 +2163,9 @@ struct common_speculative_impl_ngram_cache : public common_speculative_impl {
 struct common_speculative {
     common_speculative_draft_params_vec dparams;

+    // the target context, used to convert legacy llama_batch inputs
+    llama_context * ctx_tgt = nullptr;
+
     // list of implementations to use and their states
     std::vector<std::unique_ptr<common_speculative_impl>> impls;

@@ -2726,6 +2711,7 @@ common_speculative * common_speculative_init(common_params_speculative & params,

     common_speculative_ptr result(new common_speculative {
         /* .dparams     = */ common_speculative_draft_params_vec(n_seq),
+        /* .ctx_tgt     = */ params.draft.ctx_tgt,
         /* .impls       = */ std::move(impls),
         /* .impl_last   = */ std::vector<common_speculative_impl *>(n_seq, nullptr),
         /* .synth_probs = */ {},
@@ -2789,6 +2775,17 @@ void common_speculative_begin(common_speculative * spec, llama_seq_id seq_id, co
 }

 bool common_speculative_process(common_speculative * spec, const llama_batch & batch) {
+    if (spec == nullptr) {
+        return true;
+    }
+
+    // ngram-only setups have no target context, they do not read the batch anyway
+    const common_batch tmp = spec->ctx_tgt ? common_batch_from_llama_batch(spec->ctx_tgt, batch) : common_batch();
+
+    return common_speculative_process(spec, tmp);
+}
+
+bool common_speculative_process(common_speculative * spec, const common_batch & batch) {
     bool result = true;

     if (spec == nullptr) {
diff --git a/common/speculative.h b/common/speculative.h
index c968750e2..211fcdabd 100644
--- a/common/speculative.h
+++ b/common/speculative.h
@@ -77,6 +77,9 @@ common_speculative_draft_params & common_speculative_get_draft_params(common_spe
 void common_speculative_begin(common_speculative * spec, llama_seq_id seq_id, const llama_tokens & prompt);

 // process the batch and update the internal state of the speculative context
+bool common_speculative_process(common_speculative * spec, const common_batch & batch);
+
+// legacy llama_batch input, converted with common_batch_from_llama_batch()
 bool common_speculative_process(common_speculative * spec, const llama_batch & batch);

 // generate drafts for the sequences specified with `common_speculative_get_draft_params`
diff --git a/examples/speculative-simple/speculative-simple.cpp b/examples/speculative-simple/speculative-simple.cpp
index 863af5a2c..81aa106f1 100644
--- a/examples/speculative-simple/speculative-simple.cpp
+++ b/examples/speculative-simple/speculative-simple.cpp
@@ -228,7 +228,6 @@ int main(int argc, char ** argv) {
                 common_batch_add(batch_tgt, draft[i], n_past + i, { seq_id }, true);
             }

-            //LOG_DBG("target batch: %s\n", string_from(ctx_tgt, batch_tgt).c_str());

             llama_decode(ctx_tgt, batch_tgt);
         }
diff --git a/tools/mtmd/mtmd-cli.cpp b/tools/mtmd/mtmd-cli.cpp
index ba18b3e32..6fe058fd8 100644
--- a/tools/mtmd/mtmd-cli.cpp
+++ b/tools/mtmd/mtmd-cli.cpp
@@ -81,7 +81,7 @@ struct mtmd_cli_context {
     llama_context     * lctx;
     const llama_vocab * vocab;
     common_sampler    * smpl;
-    llama_batch         batch;
+    common_batch        batch;
     int                 n_batch;

     mtmd::bitmaps bitmaps;
@@ -115,7 +115,7 @@ struct mtmd_cli_context {
         vocab = llama_model_get_vocab(model);
         smpl = common_sampler_init(model, params.sampling);
         n_threads = params.cpuparams.n_threads;
-        batch = llama_batch_init(1, 0, 1); // batch for next token generation
+        batch = common_batch(lctx); // batch for next token generation
         n_batch = params.n_batch;

         init_vision_context(params);
@@ -148,7 +148,6 @@ struct mtmd_cli_context {
     }

     ~mtmd_cli_context() {
-        llama_batch_free(batch);
         common_sampler_free(smpl);
     }

@@ -230,9 +229,9 @@ static int generate_response(mtmd_cli_context & ctx, int n_predict) {
         }

         // eval the token
-        common_batch_clear(ctx.batch);
-        common_batch_add(ctx.batch, token_id, ctx.n_past++, {0}, true);
-        if (llama_decode(ctx.lctx, ctx.batch)) {
+        ctx.batch.clear();
+        ctx.batch.add(token_id, ctx.n_past++, 0, true);
+        if (llama_process(ctx.lctx, LLAMA_PROCESS_TYPE_DECODE, ctx.batch.get())) {
             LOG_ERR("failed to decode token\n");
             return 1;
         }
diff --git a/tools/mtmd/mtmd-helper-common.h b/tools/mtmd/mtmd-helper-common.h
index f907346c7..bc68ed4b9 100644
--- a/tools/mtmd/mtmd-helper-common.h
+++ b/tools/mtmd/mtmd-helper-common.h
@@ -6,6 +6,7 @@

 #include "ggml.h"
 #include "llama.h"
+#include "llama-cpp.h"
 #include "mtmd.h"

 #include <cstdarg>
@@ -73,112 +74,99 @@ inline mtmd_helper_logger g_logger;
 struct decode_embd_batch {
     int n_pos_per_embd;
     int n_mmproj_embd;
-    std::vector<llama_pos>      pos;
-    std::vector<llama_pos>      pos_view; // used by mrope
-    std::vector<int32_t>        n_seq_id;
-    std::vector<llama_seq_id>   seq_id_0;
-    std::vector<llama_seq_id *> seq_ids;
-    std::vector<int8_t>         logits;
-    llama_batch batch;
-    decode_embd_batch(float * embd, int32_t n_tokens, int n_pos_per_embd, int n_mmproj_embd) : n_pos_per_embd(n_pos_per_embd), n_mmproj_embd(n_mmproj_embd) {
+    int32_t n_tokens;
+    const float * embd;              // [n_tokens, n_mmproj_embd], not owned
+    std::vector<llama_pos> pos;      // [n_pos_per_embd, n_tokens], section-major
+    std::vector<llama_pos> pos_view; // sliced positions of the last get_view()
+    std::vector<int8_t>    logits;
+    llama_seq_id seq_id = 0;
+
+    llama_batch_ext_ptr batch; // rendered sub-batch, see render()
+
+    decode_embd_batch(const float * embd, int32_t n_tokens, int n_pos_per_embd, int n_mmproj_embd)
+            : n_pos_per_embd(n_pos_per_embd), n_mmproj_embd(n_mmproj_embd), n_tokens(n_tokens), embd(embd) {
         GGML_ASSERT(n_tokens > 0 && n_pos_per_embd > 0 && n_mmproj_embd > 0);
-        pos     .resize((size_t) n_tokens * (size_t) n_pos_per_embd);
-        n_seq_id.resize(n_tokens);
-        seq_ids .resize(n_tokens + 1);
-        logits  .resize(n_tokens);
-        seq_id_0.resize(1);
-        seq_ids [n_tokens] = nullptr;
-        batch = {
-            /*n_tokens       =*/ n_tokens,
-            /*tokens         =*/ nullptr,
-            /*embd           =*/ embd,
-            /*pos            =*/ pos.data(),
-            /*n_seq_id       =*/ n_seq_id.data(),
-            /*seq_id         =*/ seq_ids.data(),
-            /*logits         =*/ logits.data(),
-        };
+        pos   .resize((size_t) n_tokens * (size_t) n_pos_per_embd);
+        logits.resize(n_tokens, 0);
     }

     void set_position_normal(llama_pos pos_0, llama_seq_id seq_id) {
-        seq_id_0[0] = seq_id;
-        for (int i = 0; i < batch.n_tokens; i++) {
-            batch.pos     [i] = pos_0 + i;
-            batch.n_seq_id[i] = 1;
-            batch.seq_id  [i] = seq_id_0.data();
-            batch.logits  [i] = false;
+        this->seq_id = seq_id;
+        for (int i = 0; i < n_tokens; i++) {
+            pos[i] = pos_0 + i;
         }
     }

     // M-RoPE for image
     void set_position_mrope_2d(const std::vector<mtmd_decoder_pos> & rel_pos, llama_seq_id seq_id) {
         GGML_ASSERT(n_pos_per_embd == 4);
-        GGML_ASSERT(!rel_pos.empty() && (int32_t)rel_pos.size() == batch.n_tokens);
-        seq_id_0[0] = seq_id;
-        for (int32_t i = 0; i < batch.n_tokens; i++) {
+        GGML_ASSERT(!rel_pos.empty() && (int32_t)rel_pos.size() == n_tokens);
+        this->seq_id = seq_id;
+        for (int32_t i = 0; i < n_tokens; i++) {
             const size_t idx = (size_t) i;
-            const size_t n_tokens = (size_t) batch.n_tokens;
-            pos[idx                    ] = rel_pos[i].t;
-            pos[idx + n_tokens         ] = rel_pos[i].y;
-            pos[idx + n_tokens * 2     ] = rel_pos[i].x;
-            pos[idx + n_tokens * 3     ] = rel_pos[i].z;
-        }
-        for (int i = 0; i < batch.n_tokens; i++) {
-            batch.n_seq_id[i] = 1;
-            batch.seq_id  [i] = seq_id_0.data();
-            batch.logits  [i] = false;
+            const size_t n   = (size_t) n_tokens;
+            pos[idx        ] = rel_pos[i].t;
+            pos[idx + n    ] = rel_pos[i].y;
+            pos[idx + n * 2] = rel_pos[i].x;
+            pos[idx + n * 3] = rel_pos[i].z;
         }
     }

     // M-RoPE for audio
     void set_position_mrope_1d(llama_pos pos_0, llama_seq_id seq_id) {
         GGML_ASSERT(n_pos_per_embd == 4);
-        seq_id_0[0] = seq_id;
-        for (int i = 0; i < batch.n_tokens; i++) {
+        this->seq_id = seq_id;
+        for (int i = 0; i < n_tokens; i++) {
             const size_t idx = (size_t) i;
-            const size_t n_tokens = (size_t) batch.n_tokens;
-            pos[idx                    ] = pos_0 + i;
-            pos[idx + n_tokens         ] = pos_0 + i;
-            pos[idx + n_tokens * 2     ] = pos_0 + i;
-            pos[idx + n_tokens * 3     ] = pos_0 + i;
-        }
-        for (int i = 0; i < batch.n_tokens; i++) {
-            batch.n_seq_id[i] = 1;
-            batch.seq_id  [i] = seq_id_0.data();
-            batch.logits  [i] = false;
+            const size_t n   = (size_t) n_tokens;
+            pos[idx        ] = pos_0 + i;
+            pos[idx + n    ] = pos_0 + i;
+            pos[idx + n * 2] = pos_0 + i;
+            pos[idx + n * 3] = pos_0 + i;
         }
     }

-    llama_batch get_view(int offset, int n_tokens) {
-        GGML_ASSERT(offset >= 0 && n_tokens > 0 && offset + n_tokens <= batch.n_tokens);
-        llama_pos * pos_ptr;
+    // describe the entries [offset, offset + n) with section-major positions
+    mtmd_helper_embd_batch get_view(int offset, int n) {
+        GGML_ASSERT(offset >= 0 && n > 0 && offset + n <= n_tokens);
         pos_view.clear();
-        pos_view.reserve((size_t) n_tokens * (size_t) n_pos_per_embd);
-        if (n_pos_per_embd > 1) {
-            // mrope
-            // for example, with layout of src: 1234...1234...1234...1234...
-            //       offset 2 will give us dst: 34...34...34...34...
-            for (int i = 0; i < n_pos_per_embd; i++) {
-                // assume n_tokens is less than or equal to batch.n_tokens
-                // batch.n_tokens is number of **total** tokens
-                // n_tokens is number of viewed token
-                size_t src_idx = (size_t) i * (size_t) batch.n_tokens + (size_t) offset;
-                pos_view.insert(pos_view.end(),
-                    pos.data() + src_idx,
-                    pos.data() + src_idx + n_tokens);
-            }
-            pos_ptr = pos_view.data();
-        } else {
-            // normal
-            pos_ptr = pos.data() + offset;
+        pos_view.reserve((size_t) n * (size_t) n_pos_per_embd);
+        for (int j = 0; j < n_pos_per_embd; j++) {
+            const size_t src = (size_t) j * (size_t) n_tokens + (size_t) offset;
+            pos_view.insert(pos_view.end(), pos.data() + src, pos.data() + src + n);
         }
         return {
-            /*n_tokens       =*/ n_tokens,
-            /*tokens         =*/ nullptr,
-            /*embd           =*/ batch.embd     + offset * n_mmproj_embd,
-            /*pos            =*/ pos_ptr,
-            /*n_seq_id       =*/ batch.n_seq_id + offset,
-            /*seq_id         =*/ batch.seq_id   + offset,
-            /*logits         =*/ batch.logits   + offset,
+            /*n_tokens =*/ n,
+            /*embd     =*/ embd + (size_t) offset * n_mmproj_embd,
+            /*n_embd   =*/ n_mmproj_embd,
+            /*pos      =*/ pos_view.data(),
+            /*n_pos    =*/ n_pos_per_embd,
+            /*seq_id   =*/ seq_id,
         };
     }
+
+    // render the entries [offset, offset + n) into a batch owned by this object, ready for llama_process()
+    llama_batch_ext * render(llama_context * lctx, int offset, int n) {
+        GGML_ASSERT(offset >= 0 && n > 0 && offset + n <= n_tokens);
+        if (!batch) {
+            batch.reset(llama_batch_ext_init(lctx));
+        }
+        llama_batch_ext_clear(batch.get());
+        for (int i = offset; i < offset + n; i++) {
+            const llama_embd e = { embd + (size_t) i * n_mmproj_embd, 1, (size_t) n_mmproj_embd };
+            const int32_t idx = llama_batch_ext_add_embd(batch.get(), seq_id, e);
+            GGML_ASSERT(idx >= 0);
+
+            llama_pos p[GGML_MROPE_SECTIONS] = { 0, 0, 0, 0 };
+            for (int j = 0; j < n_pos_per_embd; j++) {
+                p[j] = pos[(size_t) j * (size_t) n_tokens + (size_t) i];
+            }
+            llama_batch_ext_set_pos(batch.get(), idx, p);
+
+            if (logits[i]) {
+                llama_batch_ext_set_output_logits(batch.get(), idx, true);
+            }
+        }
+        return batch.get();
+    }
 };
diff --git a/tools/mtmd/mtmd-helper-gen.cpp b/tools/mtmd/mtmd-helper-gen.cpp
index 1c58d3ae1..5fb7ea9a0 100644
--- a/tools/mtmd/mtmd-helper-gen.cpp
+++ b/tools/mtmd/mtmd-helper-gen.cpp
@@ -222,14 +222,13 @@ public:
             return 0;
         }
         const int32_t n_tokens_batch = std::min(n_batch, n_prompt - prompt_pos);
-        llama_batch batch_view = prompt_batch->get_view(prompt_pos, n_tokens_batch);

         const bool is_last_batch = (prompt_pos + n_tokens_batch) == n_prompt;
         if (is_last_batch) {
-            batch_view.logits[n_tokens_batch - 1] = 1;
+            prompt_batch->logits[prompt_pos + n_tokens_batch - 1] = 1;
         }

-        if (llama_decode(lctx, batch_view) != 0) {
+        if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, prompt_batch->render(lctx, prompt_pos, n_tokens_batch)) != 0) {
             LOG_ERR("mtmd_helper_gen_audio: prompt decode failed\n");
             return -1;
         }
@@ -286,10 +285,10 @@ public:
         decode_embd_batch batch_embd(fb.data(), 1, n_pos_per_embd, n_embd);
         if (mrope) batch_embd.set_position_mrope_1d(pos, seq_id);
         else       batch_embd.set_position_normal  (pos, seq_id);
-        batch_embd.batch.logits[0] = 1;
+        batch_embd.logits[0] = 1;
         pos++;

-        if (llama_decode(lctx, batch_embd.batch) != 0) {
+        if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch_embd.render(lctx, 0, 1)) != 0) {
             LOG_ERR("mtmd_helper_gen_audio: decode failed\n");
             return 1;
         }
@@ -586,13 +585,12 @@ public:
             return 0;
         }
         const int32_t n_tokens_batch = std::min(n_batch, n_prompt - prompt_pos);
-        llama_batch batch_view = prompt_batch->get_view(prompt_pos, n_tokens_batch);

         if ((prompt_pos + n_tokens_batch) == n_prompt) {
-            batch_view.logits[n_tokens_batch - 1] = 1;
+            prompt_batch->logits[prompt_pos + n_tokens_batch - 1] = 1;
         }

-        if (llama_decode(lctx, batch_view) != 0) {
+        if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, prompt_batch->render(lctx, prompt_pos, n_tokens_batch)) != 0) {
             LOG_ERR("mtmd_helper_gen_audio: prompt decode failed\n");
             return -1;
         }
@@ -646,12 +644,12 @@ public:
             }
         }

-        decode_embd_batch batch_embd(const_cast<float *>(out.embd), 1, 1, n_embd);
+        decode_embd_batch batch_embd(out.embd, 1, 1, n_embd);
         batch_embd.set_position_normal(pos, seq_id);
-        batch_embd.batch.logits[0] = 1;
+        batch_embd.logits[0] = 1;
         pos++;

-        if (llama_decode(lctx, batch_embd.batch) != 0) {
+        if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch_embd.render(lctx, 0, 1)) != 0) {
             LOG_ERR("mtmd_helper_gen_audio: decode failed\n");
             return 1;
         }
@@ -842,8 +840,8 @@ private:
         GGML_ASSERT(n_rows > 0);
         decode_embd_batch batch(prompt_embd_buf.data(), n_rows, 1, n_e);
         batch.set_position_normal(pos, seq_id);
-        batch.batch.logits[n_rows - 1] = 1;
-        if (llama_decode(lctx, batch.batch) != 0) {
+        batch.logits[n_rows - 1] = 1;
+        if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.render(lctx, 0, n_rows)) != 0) {
             LOG_ERR("mtmd_helper_gen_audio: chunk prompt decode failed\n");
             return 1;
         }
diff --git a/tools/mtmd/mtmd-helper.cpp b/tools/mtmd/mtmd-helper.cpp
index bdf8bf6fe..cd05ade2d 100644
--- a/tools/mtmd/mtmd-helper.cpp
+++ b/tools/mtmd/mtmd-helper.cpp
@@ -169,19 +169,19 @@ int32_t mtmd_helper_decode_image_chunk(
     while (i_batch < n_img_batches) { // split into batches
         int pos_offset = i_batch*n_batch;
         int n_tokens_batch = std::min(n_batch, n_tokens - pos_offset);
-        llama_batch batch_embd_view = batch_embd.get_view(pos_offset, n_tokens_batch);

         LOG_INF("decoding %s batch %d/%d, n_tokens_batch = %d\n", name, i_batch+1, n_img_batches, n_tokens_batch);

         int64_t t1 = ggml_time_ms();
-        int32_t ret = llama_decode(lctx, batch_embd_view);
+        int32_t ret = llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch_embd.render(lctx, pos_offset, n_tokens_batch));
         if (ret != 0) {
             LOG_ERR("failed to decode %s\n", name);
             return ret;
         }

         if (callback != nullptr) {
-            ret = callback(batch_embd_view, user_data);
+            const mtmd_helper_embd_batch view = batch_embd.get_view(pos_offset, n_tokens_batch);
+            ret = callback(&view, user_data);
             if (ret != 0) {
                 LOG_ERR("post-decode callback failed\n");
                 return ret;
@@ -209,37 +209,35 @@ int32_t mtmd_helper_eval_chunk_single(mtmd_context * ctx,
         llama_pos * new_n_past) {
     GGML_ASSERT(n_batch > 0);
     int32_t ret;
-    llama_batch text_batch = llama_batch_init(n_batch, 0, 1);
     auto chunk_type = mtmd_input_chunk_get_type(chunk);

     if (chunk_type == MTMD_INPUT_CHUNK_TYPE_TEXT) {
         size_t n_tokens;
         const auto tokens = mtmd_input_chunk_get_tokens_text(chunk, &n_tokens);
         // LOG_INF("decoding text chunk, n_tokens = %zu\n", n_tokens);
+        llama_batch_ext_ptr text_batch(llama_batch_ext_init(lctx));
         size_t i = 0;
         while (i < n_tokens) { // split into batches
-            text_batch.n_tokens = 0; // clear the batch
-            for (; i < n_tokens && text_batch.n_tokens < n_batch; i++) {
-                int32_t j = text_batch.n_tokens;
-                text_batch.token   [j]    = tokens[i];
-                text_batch.pos     [j]    = n_past++;
-                text_batch.n_seq_id[j]    = 1;
-                text_batch.seq_id  [j][0] = seq_id;
-                text_batch.logits  [j]    = false;
-
-                text_batch.n_tokens++;
+            llama_batch_ext_clear(text_batch.get());
+            int32_t n_added = 0;
+            int32_t idx     = -1;
+            for (; i < n_tokens && n_added < n_batch; i++) {
+                idx = llama_batch_ext_add_token(text_batch.get(), seq_id, tokens[i]);
+                GGML_ASSERT(idx >= 0);
+                llama_pos pos = n_past++;
+                llama_batch_ext_set_pos(text_batch.get(), idx, &pos);
+                n_added++;
             }
             bool is_last_token = (i == n_tokens);
             if (logits_last && is_last_token) {
-                text_batch.logits[text_batch.n_tokens - 1] = true;
+                llama_batch_ext_set_output_logits(text_batch.get(), idx, true);
             }
-            ret = llama_decode(lctx, text_batch);
+            ret = llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, text_batch.get());
             if (ret != 0) {
                 LOG_ERR("failed to decode text\n");
-                llama_batch_free(text_batch);
                 return ret;
             }
-            *new_n_past += text_batch.n_tokens;
+            *new_n_past += n_added;
         }

     } else if (chunk_type == MTMD_INPUT_CHUNK_TYPE_IMAGE || chunk_type == MTMD_INPUT_CHUNK_TYPE_AUDIO) {
@@ -251,7 +249,6 @@ int32_t mtmd_helper_eval_chunk_single(mtmd_context * ctx,
         ret = mtmd_encode_chunk(ctx, chunk);
         if (ret != 0) {
             LOG_ERR("failed to encode %s slice\n", name);
-            llama_batch_free(text_batch);
             return ret;
         }

@@ -261,14 +258,12 @@ int32_t mtmd_helper_eval_chunk_single(mtmd_context * ctx,
         ret = mtmd_helper_decode_image_chunk(ctx, lctx, chunk, embd, n_past, seq_id, n_batch, new_n_past, nullptr, nullptr);
         if (ret != 0) {
             LOG_ERR("failed to decode %s\n", name);
-            llama_batch_free(text_batch);
             return ret;
         }
     } else {
         GGML_ABORT("chunk type not supported");
     }

-    llama_batch_free(text_batch);
     return 0;
 }

diff --git a/tools/mtmd/mtmd-helper.h b/tools/mtmd/mtmd-helper.h
index 10f2171c0..7436230f0 100644
--- a/tools/mtmd/mtmd-helper.h
+++ b/tools/mtmd/mtmd-helper.h
@@ -92,9 +92,9 @@ MTMD_API llama_pos mtmd_helper_get_n_pos(const mtmd_input_chunks * chunks);
 MTMD_API void mtmd_helper_image_get_decoder_pos(const mtmd_image_tokens * image, llama_pos pos_0, struct mtmd_decoder_pos * out_pos);

 // helper function that automatically:
-// 1. run llama_decode() on text chunks
-// 2. run mtmd_encode_chunk() on image chunks, then mtmd_get_output_embd() and then llama_decode()
-// if any of the mtmd_encode_chunk() or llama_decode() calls return non-zero, stop and forward the error
+// 1. decode text chunks
+// 2. run mtmd_encode_chunk() on image chunks, then mtmd_get_output_embd() and then decode the embeddings
+// if any of the mtmd_encode_chunk() or decode calls return non-zero, stop and forward the error
 // otherwise, returns 0 on success
 // this function is NOT thread-safe
 MTMD_API int32_t mtmd_helper_eval_chunks(mtmd_context * ctx,
@@ -117,7 +117,17 @@ MTMD_API int32_t mtmd_helper_eval_chunk_single(mtmd_context * ctx,
                                                bool logits_last,
                                                llama_pos * new_n_past);

-typedef int32_t (*mtmd_helper_post_decode_callback)(struct llama_batch batch, void * user_data);
+// one decoded sub-batch of embeddings, passed to mtmd_helper_post_decode_callback
+struct mtmd_helper_embd_batch {
+    int32_t n_tokens;
+    const float     * embd;   // [n_tokens, n_embd]
+    int32_t           n_embd;
+    const llama_pos * pos;    // [n_pos, n_tokens], section-major
+    int32_t           n_pos;  // 4 for M-RoPE models, 1 otherwise
+    llama_seq_id      seq_id;
+};
+
+typedef int32_t (*mtmd_helper_post_decode_callback)(const struct mtmd_helper_embd_batch * batch, void * user_data);

 // helper function to decode an image whose embeddings have already been calculated
 // this helper will handle batching and pre/post decoding setup (for ex. gemma 3 requires non-causal attention)
diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp
index 611e82a6a..efff54997 100644
--- a/tools/server/server-context.cpp
+++ b/tools/server/server-context.cpp
@@ -109,8 +109,7 @@ enum slot_state {
 struct server_slot; // forward declaration

 struct server_batch {
-    llama_batch batch;
-    bool batch_rendered = false;
+    common_batch view; // the rendered sub-batch [off, off + n_tokens), see render()

     struct token {
         int32_t id_slot;
@@ -126,36 +125,21 @@ struct server_batch {
     // track if given slot can be batched with slots already in the batch
     server_slot * slot_batched = nullptr;

-    // in embd mode, we temporarily swap out the tokens arr and restore it on clear()
     bool has_embd = false;
-    llama_token * tokens_ptr = nullptr;
     std::vector<float> embd;

     float  alora_scale       = -1.0f;
     size_t alora_disabled_id = 0;

-    server_batch() {
-        batch.pos = nullptr; // sentinel: uninitialized batch
-    }
-
-    ~server_batch() {
-        if (batch.pos != nullptr) {
-            clear();
-            llama_batch_free(batch);
-        }
-    }
-
-    void init(int32_t n_tokens_alloc, int32_t n_embd) {
+    void init(llama_context * ctx, int32_t n_tokens_alloc, int32_t n_embd) {
         this->n_tokens_alloc = n_tokens_alloc;
         this->n_embd = n_embd;
-        batch = llama_batch_init(n_tokens_alloc, 0, 1);
-        tokens_ptr = batch.token;
+        view = common_batch(ctx);
         tokens.reserve(n_tokens_alloc);
     }

     bool add(int32_t id_slot, llama_token token, llama_pos pos, bool output, bool is_prompt) {
         GGML_ASSERT(!has_embd); // cannot mix tokens + embd in same batch
-        GGML_ASSERT(batch.pos != nullptr);
         if ((int32_t)tokens.size() >= n_tokens_alloc) {
             return false;
         }
@@ -164,7 +148,6 @@ struct server_batch {
     }

     bool add(int32_t id_slot, const std::vector<float> & embd_in, llama_pos pos, bool output, bool is_prompt) {
-        GGML_ASSERT(batch.pos != nullptr);
         if ((int32_t)tokens.size() >= n_tokens_alloc) {
             return false;
         }
@@ -177,16 +160,11 @@ struct server_batch {
     void clear() {
         tokens.clear();
         embd.clear();
-        common_batch_clear(batch);
+        view.clear();
         slot_batched      = nullptr;
         alora_scale       = -1.0f;
         alora_disabled_id = 0;
-        batch_rendered    = false;
         has_embd          = false;
-        if (batch.token == nullptr) {
-            batch.token = tokens_ptr;
-            batch.embd  = nullptr;
-        }
     }

     int32_t size() const {
@@ -198,41 +176,22 @@ struct server_batch {
         tokens[idx].output = output;
     }

-    void render() {
-        GGML_ASSERT(!batch_rendered);
-        GGML_ASSERT(batch.pos != nullptr);
-        common_batch_clear(batch);
-        for (int32_t i = 0; i < size(); i++) {
-            const auto & t = tokens[i];
-            common_batch_add(batch, t.token, t.pos, { t.id_slot }, t.output);
-        }
-        if (has_embd) {
-            batch.token = nullptr; // will be restored on clear()
-            batch.embd  = embd.data();
-        }
-        batch_rendered = true;
-    }
-
-    llama_batch get_view(int32_t off, int32_t n_tokens) const {
-        GGML_ASSERT(batch.pos != nullptr);
-        GGML_ASSERT(batch_rendered);
+    // render the sub-batch [off, off + n_tokens) into view, index i in view is index off + i here
+    void render(int32_t off, int32_t n_tokens) {
         GGML_ASSERT(off >= 0 && off < size());
         GGML_ASSERT(n_tokens > 0 && off + n_tokens <= size());

-        auto * token = batch.token ? batch.token + off          : nullptr;
-        auto * embd  = batch.embd  ? batch.embd  + off * n_embd : nullptr;
-
-        llama_batch view = {
-            n_tokens,
-            token,
-            embd,
-            batch.pos      + off,
-            batch.n_seq_id + off,
-            batch.seq_id   + off,
-            batch.logits   + off,
-        };
-
-        return view;
+        view.clear();
+        for (int32_t i = off; i < off + n_tokens; i++) {
+            const auto & t = tokens[i];
+            if (has_embd) {
+                // text embeddings broadcast the same position across the M-RoPE sections
+                const llama_pos pos[GGML_MROPE_SECTIONS] = { t.pos, t.pos, t.pos, 0 };
+                view.add_embd({ embd.data() + (size_t) i * n_embd, 1, (size_t) n_embd }, pos, t.id_slot, t.output);
+            } else {
+                view.add(t.token, t.pos, t.id_slot, t.output);
+            }
+        }
     }
 };

@@ -761,13 +720,24 @@ static int process_mtmd_chunk(const server_slot & slot, mtmd::batch_ptr & mbatch
         if (mbatch) {
             float * embd = mtmd_batch_get_output_embd(mbatch.get(), chunk.get());
             if (embd) {
-                void * cb_data = slot.spec;
-                static auto cb = [](llama_batch batch, void * user_data) {
-                    common_speculative * spec = static_cast<common_speculative *>(user_data);
-                    if (!common_speculative_process(spec, batch)) {
-                        return 1;
+                struct cb_data_t {
+                    common_speculative * spec;
+                    llama_context * ctx;
+                } cb_data = { slot.spec, slot.ctx_tgt };
+
+                static auto cb = [](const mtmd_helper_embd_batch * b, void * user_data) {
+                    const auto * data = static_cast<cb_data_t *>(user_data);
+
+                    common_batch batch(data->ctx);
+                    for (int32_t i = 0; i < b->n_tokens; ++i) {
+                        llama_pos pos[GGML_MROPE_SECTIONS] = { 0, 0, 0, 0 };
+                        for (int32_t j = 0; j < b->n_pos; ++j) {
+                            pos[j] = b->pos[j * b->n_tokens + i];
+                        }
+                        batch.add_embd({ b->embd + (size_t) i * b->n_embd, 1, (size_t) b->n_embd }, pos, b->seq_id, false);
                     }
-                    return 0;
+
+                    return common_speculative_process(data->spec, batch) ? 0 : 1;
                 };

                 llama_pos new_n_past; // unused for now
@@ -781,7 +751,7 @@ static int process_mtmd_chunk(const server_slot & slot, mtmd::batch_ptr & mbatch
                     llama_n_batch(slot.ctx_tgt),
                     &new_n_past,
                     cb,
-                    cb_data
+                    &cb_data
                 );
                 if (res != 0) {
                     SLT_ERR(slot, "failed to decode mtmd chunk, idx = %zu, res = %d\n", idx, res);
@@ -1356,7 +1326,7 @@ private:
         {
             const int32_t n_batch = llama_n_batch(ctx_tgt);
             const int32_t n_embd  = llama_model_n_embd_inp(model_tgt);
-            batch.init(std::max(n_batch, params_base.n_parallel), n_embd);
+            batch.init(ctx_tgt, std::max(n_batch, params_base.n_parallel), n_embd);
         }

         if (params_base.cache_ram_mib != 0) {
@@ -2160,7 +2130,7 @@ private:
         queue_results.send(std::move(res));
     }

-    void send_embedding(const server_slot & slot, const llama_batch & batch) {
+    void send_embedding(const server_slot & slot, const common_batch & batch) {
         auto res = std::make_unique<server_task_result_embd>();
         res->id        = slot.task->id;
         res->index     = slot.task->index;
@@ -2171,8 +2141,8 @@ private:

         std::vector<float> embd_res(n_embd_out, 0.0f);

-        for (int i = 0; i < batch.n_tokens; ++i) {
-            if (!batch.logits[i] || batch.seq_id[i][0] != slot.id) {
+        for (int i = 0; i < batch.size(); ++i) {
+            if (!batch.tokens[i].output || batch.tokens[i].seq_id != slot.id) {
                 continue;
             }

@@ -2180,11 +2150,11 @@ private:
             if (llama_pooling_type(slot.ctx_tgt) == LLAMA_POOLING_TYPE_NONE) {
                 embd = llama_get_embeddings_ith(slot.ctx_tgt, i);
             } else {
-                embd = llama_get_embeddings_seq(slot.ctx_tgt, batch.seq_id[i][0]);
+                embd = llama_get_embeddings_seq(slot.ctx_tgt, batch.tokens[i].seq_id);
             }

             if (embd == nullptr) {
-                SLT_ERR(slot, "failed to get embeddings, token = %d, seq_id = %d\n", batch.token[i], batch.seq_id[i][0]);
+                SLT_ERR(slot, "failed to get embeddings, token = %d, seq_id = %d\n", batch.tokens[i].id, batch.tokens[i].seq_id);

                 res->embedding.push_back(std::vector<float>(n_embd_out, 0.0f));
                 continue;
@@ -2205,24 +2175,24 @@ private:
         queue_results.send(std::move(res));
     }

-    void send_rerank(const server_slot & slot, const llama_batch & batch) {
+    void send_rerank(const server_slot & slot, const common_batch & batch) {
         auto res = std::make_unique<server_task_result_rerank>();
         res->id       = slot.task->id;
         res->index    = slot.task->index;
         res->n_tokens = slot.task->n_tokens();

-        for (int i = 0; i < batch.n_tokens; ++i) {
-            if (!batch.logits[i] || batch.seq_id[i][0] != slot.id) {
+        for (int i = 0; i < batch.size(); ++i) {
+            if (!batch.tokens[i].output || batch.tokens[i].seq_id != slot.id) {
                 continue;
             }

-            const float * embd = llama_get_embeddings_seq(ctx_tgt, batch.seq_id[i][0]);
+            const float * embd = llama_get_embeddings_seq(ctx_tgt, batch.tokens[i].seq_id);
             if (embd == NULL) {
                 embd = llama_get_embeddings_ith(ctx_tgt, i);
             }

             if (embd == NULL) {
-                SLT_ERR(slot, "failed to get embeddings, token = %d, seq_id = %d\n", batch.token[i], batch.seq_id[i][0]);
+                SLT_ERR(slot, "failed to get embeddings, token = %d, seq_id = %d\n", batch.tokens[i].id, batch.tokens[i].seq_id);

                 res->score = -1e6;
                 continue;
@@ -2845,7 +2815,6 @@ private:
         try {
             scoped_timer t(t_pre_decode, n_pre_decode);
             pre_decode();
-            batch.render();
         } catch (const std::exception & e) {
             SRV_ERR("pre_decode() failed: %s\n", e.what());
             abort_all_slots("pre_decode() failed: " + std::string(e.what()));
@@ -2875,7 +2844,6 @@ private:
             llama_set_embeddings(ctx_tgt, slot_batched->need_embd());
         }

-        llama_batch batch_view;
         int32_t off_next = 0;
         int32_t n_batch = llama_n_batch(ctx_tgt);
         for (int32_t off = 0; off < batch.size(); off = off_next) {
@@ -2884,8 +2852,8 @@ private:
                 scoped_timer t(t_decode, n_decode);
                 // TODO @ngxson : maybe handle n_batch == 1 here instead of inside decode()

-                batch_view = batch.get_view(off, n_tokens);
-                bool ok = decode(n_batch, off, batch_view);
+                batch.render(off, n_tokens);
+                bool ok = decode(n_batch, off);
 #ifdef DEBUG_TIMINGS
                 llama_synchronize(ctx_tgt);
 #endif
@@ -2908,7 +2876,7 @@ private:

             try {
                 scoped_timer t(t_post_decode, n_post_decode);
-                post_decode(n_tokens, off, batch_view);
+                post_decode(n_tokens, off);
             } catch (const std::exception & e) {
                 SRV_ERR("post_decode() failed: %s\n", e.what());
                 abort_all_slots("post_decode() failed: " + std::string(e.what()));
@@ -3655,7 +3623,7 @@ private:

     // returns true = success ; false = retry with smaller batch size
     // throw std::runtime_error on fatal error
-    bool decode(int32_t & n_batch, int32_t off, llama_batch & batch_view) {
+    bool decode(int32_t & n_batch, int32_t off) {
         SRV_DBG("n_batch (effective) = %d, off = %d\n", n_batch, off);

         metrics_pre_decode();
@@ -3682,7 +3650,7 @@ private:
         }

         bool has_output = false;
-        for (int i = off; i < off + batch_view.n_tokens; ++i) {
+        for (int i = off; i < off + batch.view.size(); ++i) {
             has_output |= batch.tokens[i].output;
         }

@@ -3690,7 +3658,7 @@ private:
         // note: the sync is done here too, so that the wait is also covered by the yield
         int ret = 0;
         queue_tasks.yield_to_queue([&]() {
-            ret = llama_decode(ctx_tgt, batch_view);
+            ret = llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, batch.view.get());
             if (ret == 0 && has_output) {
                 llama_synchronize(ctx_tgt);
             }
@@ -3746,7 +3714,7 @@ private:
             return false; // retry with the updated n_batch
         } else {
             // success, apply batch metrics
-            metrics_post_decode(off, batch_view.n_tokens, has_output);
+            metrics_post_decode(off, batch.view.size(), has_output);
         }

         // TODO: avoid restoring the draft context and re-evaluating the drafted tokens when not needed [TAG_SPEC_AVOID_DRAFT_REEVAL]
@@ -3755,7 +3723,7 @@ private:
         if (spec) {
             bool ok = true;
             queue_tasks.yield_to_queue([&]() {
-                ok = common_speculative_process(spec.get(), batch_view);
+                ok = common_speculative_process(spec.get(), batch.view);
             });

             if (!ok) {
@@ -3792,8 +3760,8 @@ private:
         return true;
     }

-    void post_decode(int32_t n_batch_tokens, int32_t off, llama_batch & batch_view) {
-        // for checking if a given batch index is inside batch_view
+    void post_decode(int32_t n_batch_tokens, int32_t off) {
+        // for checking if a given batch index is inside the current sub-batch
         auto is_inside_view = [&](int32_t idx) {
             return idx >= off && idx < off + n_batch_tokens;
         };
@@ -3829,14 +3797,14 @@ private:
             if (slot.state == SLOT_STATE_DONE_PROMPT) {
                 if (slot.task->type == SERVER_TASK_TYPE_EMBEDDING) {
                     // prompt evaluated for embedding
-                    send_embedding(slot, batch_view);
+                    send_embedding(slot, batch.view);
                     slot.release();
                     slot.i_batch = -1;
                     return;
                 }

                 if (slot.task->type == SERVER_TASK_TYPE_RERANK) {
-                    send_rerank(slot, batch_view);
+                    send_rerank(slot, batch.view);
                     slot.release();
                     slot.i_batch = -1;
                     return;