Commit 60e9cf7a7 for llama.cpp

commit 60e9cf7a7be81e4a1be7f15de6e8470d0e5b8b29
Author: Xuan-Son Nguyen <son@huggingface.co>
Date:   Wed Sep 30 18:08:43 2026 +0200

    batch: migrate the rest of examples to llama_batch_ext (#29601)

    * migrate the rest

    * test-thread-safety

    * rm common_batch_staged

diff --git a/common/common.cpp b/common/common.cpp
index 598a97d10..401de1dc2 100644
--- a/common/common.cpp
+++ b/common/common.cpp
@@ -1738,33 +1738,6 @@ void common_threadpools::init(llama_context * ctx, const common_params & params)
     llama_attach_threadpool(ctx, threadpool, threadpool_batch);
 }

-//
-// Batch utils
-//
-
-void common_batch_clear(struct llama_batch & batch) {
-    batch.n_tokens = 0;
-}
-
-void common_batch_add(
-                 struct llama_batch & batch,
-                        llama_token   id,
-                          llama_pos   pos,
-    const std::vector<llama_seq_id> & seq_ids,
-                               bool   logits) {
-    GGML_ASSERT(batch.seq_id[batch.n_tokens] && "llama_batch size exceeded");
-
-    batch.token   [batch.n_tokens] = id;
-    batch.pos     [batch.n_tokens] = pos;
-    batch.n_seq_id[batch.n_tokens] = seq_ids.size();
-    for (size_t i = 0; i < seq_ids.size(); ++i) {
-        batch.seq_id[batch.n_tokens][i] = seq_ids[i];
-    }
-    batch.logits  [batch.n_tokens] = logits;
-
-    batch.n_tokens++;
-}
-
 //
 // Vocab utils
 //
@@ -2118,35 +2091,41 @@ common_batch::common_batch(llama_context * ctx) : batch(llama_batch_ext_init(ctx

 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 size() - 1;
+}
+
+int32_t common_batch::add(llama_token id, llama_pos pos, const std::vector<llama_seq_id> & seq_ids, bool output) {
+    GGML_ASSERT(!seq_ids.empty());
+
+    const int32_t idx = add(id, pos, seq_ids[0], output);
+    for (size_t s = 1; s < seq_ids.size(); ++s) {
+        add_seq(idx, seq_ids[s]);
     }
-    tokens.push_back({ id, { pos, 0, 0, 0 }, seq_id, output, { nullptr, 0, 0 } });
     return idx;
 }

+bool common_batch::add_seq(int32_t idx, llama_seq_id seq_id) {
+    if (idx < 0 || idx >= size()) {
+        return false;
+    }
+    tokens[idx].seq_ids_extra.push_back(seq_id);
+    return true;
+}
+
 bool common_batch::set_output(int32_t idx, bool value) {
-    if (idx < 0 || idx >= (int32_t) tokens.size()) {
+    if (idx < 0 || idx >= size()) {
         return false;
     }
     tokens[idx].output = value;
-    return llama_batch_ext_set_output_logits(batch.get(), idx, value);
+    return true;
 }

 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)) {
+    if (idx < 0 || idx >= size() || tokens[idx].embd.data != nullptr) {
         return false;
     }
     tokens[idx].embd = embd;
@@ -2154,83 +2133,64 @@ bool common_batch::set_embd(int32_t idx, llama_embd embd) {
 }

 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 };
+    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;
+    return size() - 1;
 }

-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;
+llama_batch_ext * common_batch::get_sub_batch(int32_t off, int32_t n) {
+    GGML_ASSERT(batch && "common_batch was not initialized with a context");
+    GGML_ASSERT(off >= 0 && n >= 0 && off + n <= size());

-    const size_t n_embd = llama_model_n_embd_inp(llama_get_model(ctx));
+    llama_batch_ext * res = batch.get();
+    llama_batch_ext_clear(res);

-    // 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;
+    for (int32_t i = off; i < off + n; ++i) {
+        const token & t = tokens[i];

-        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];
+        int32_t idx;
+        if (t.id != LLAMA_TOKEN_NULL) {
+            idx = llama_batch_ext_add_token(res, t.seq_id, t.id);
+            if (idx < 0) {
+                GGML_ABORT("%s: failed to add token %d at index %d (error %d, n = %d)\n", __func__, t.id, i, idx, n);
+            }
+            llama_batch_ext_set_pos(res, idx, t.pos.data());
+            if (t.embd.data && !llama_batch_ext_set_embd_token(res, idx, t.embd)) {
+                GGML_ABORT("%s: failed to set the embedding of token %d at index %d\n", __func__, t.id, 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];
+            idx = llama_batch_ext_add_embd(res, t.seq_id, t.embd);
+            if (idx < 0) {
+                GGML_ABORT("%s: failed to add embedding at index %d (error %d, n = %d)\n", __func__, i, idx, n);
             }
+            llama_batch_ext_set_pos(res, idx, t.pos.data());
         }
+        GGML_ASSERT(idx == i - off);

-        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);
+        for (const llama_seq_id seq_id : t.seq_ids_extra) {
+            if (!llama_batch_ext_add_seq(res, idx, seq_id)) {
+                GGML_ABORT("%s: failed to add seq %d to the entry at index %d\n", __func__, seq_id, i);
             }
-        } 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]);
+        if (t.output) {
+            llama_batch_ext_set_output_logits(res, idx, true);
         }
     }

     return res;
 }

-common_batch common_batch_get_one(llama_context * ctx, const llama_tokens & tokens) {
+common_batch common_batch_get_one(llama_context * ctx, const llama_token * tokens, int32_t n_tokens) {
     common_batch batch(ctx);

     auto mem = llama_get_memory(ctx);
     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 bool output = i == tokens.size() - 1;
+    for (int32_t i = 0; i < n_tokens; ++i) {
+        const bool output = i == n_tokens - 1;
         batch.add(tokens[i], pos, 0, output);
         pos++;
     }
@@ -2238,6 +2198,10 @@ common_batch common_batch_get_one(llama_context * ctx, const llama_tokens & toke
     return batch;
 }

+common_batch common_batch_get_one(llama_context * ctx, const llama_tokens & tokens) {
+    return common_batch_get_one(ctx, tokens.data(), (int32_t) tokens.size());
+}
+
 bool common_prompt_batch_decode(
               struct llama_context * ctx,
                 const llama_tokens & all_tokens,
diff --git a/common/common.h b/common/common.h
index e95eb2fd0..dcc5ec1ae 100644
--- a/common/common.h
+++ b/common/common.h
@@ -1031,23 +1031,16 @@ struct common_memory {
 // Batch utils
 //

-void common_batch_clear(struct llama_batch & batch);
-
-void common_batch_add(
-                 struct llama_batch & batch,
-                        llama_token   id,
-                          llama_pos   pos,
-    const std::vector<llama_seq_id> & seq_ids,
-                               bool   logits);
-
 // wrapper around llama_batch_ext that provide getter functions for downstream code
+// entries can exceed n_batch, use get_sub_batch() to decode them in chunks
 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;
+        llama_seq_id seq_id; // the first sequence id, see add_seq()
         bool         output;
         llama_embd   embd; // non-owning view of the data passed to add_embd()/set_embd(), data == NULL if none
+        std::vector<llama_seq_id> seq_ids_extra; // see add_seq()
     };

     std::vector<token> tokens; // mirror of the entries, tokens[i] describes batch index i
@@ -1058,7 +1051,10 @@ struct common_batch {
     common_batch() = default;
     common_batch(struct llama_context * ctx);

-    llama_batch_ext * get() const { return batch.get(); }
+    llama_batch_ext * get() { return get_sub_batch(0, size()); }
+
+    // render entries [off, off + n) into batch, the result is overwritten by the next call
+    llama_batch_ext * get_sub_batch(int32_t off, int32_t n);

     // content type of the batch, all entries carry the same combination
     bool has_token() const { return !tokens.empty() && tokens[0].id != LLAMA_TOKEN_NULL; }
@@ -1066,15 +1062,21 @@ struct common_batch {

     void clear();

-    // returns the batch index (>= 0), aborts if the entry cannot be added (batch full, invalid token or seq id)
+    // returns the batch index
     int32_t add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output);

+    // same, with the entry shared by all seq_ids (must not be empty)
+    int32_t add(llama_token id, llama_pos pos, const std::vector<llama_seq_id> & seq_ids, bool output);
+
+    // add the entry at idx to another sequence, tokens[idx].seq_id keeps the first one
+    bool add_seq(int32_t idx, llama_seq_id seq_id);
+
     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
+    // add an embedding-only entry (no token id)
     // pos points to n_pos positions
     int32_t add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output);

@@ -1082,13 +1084,10 @@ struct common_batch {
 };

 // create a single-sequence batch from a list of tokens
-// last token always have output_logits set to true
+// positions continue from the memory, last token always have output_logits set to true
+common_batch common_batch_get_one(struct llama_context * ctx, const llama_token * tokens, int32_t n_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
 //
 // Note: We save state before the last token so that we can replay it to ensure
diff --git a/common/speculative.cpp b/common/speculative.cpp
index 82e9e9223..b1244d9a6 100644
--- a/common/speculative.cpp
+++ b/common/speculative.cpp
@@ -2163,9 +2163,6 @@ 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;

@@ -2711,7 +2708,6 @@ 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 = */ {},
@@ -2774,17 +2770,6 @@ 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;

diff --git a/common/speculative.h b/common/speculative.h
index 211fcdabd..d46b21eb7 100644
--- a/common/speculative.h
+++ b/common/speculative.h
@@ -79,9 +79,6 @@ void common_speculative_begin(common_speculative * spec, llama_seq_id seq_id, co
 // 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`
 void common_speculative_draft(common_speculative * spec);

diff --git a/examples/batched/batched.cpp b/examples/batched/batched.cpp
index 830e45f5a..9da844c7e 100644
--- a/examples/batched/batched.cpp
+++ b/examples/batched/batched.cpp
@@ -117,7 +117,7 @@ int main(int argc, char ** argv) {

     // create a llama_batch
     // we use this object to submit token data for decoding
-    llama_batch batch = llama_batch_init(std::max(tokens_list.size(), (size_t) n_parallel), 0, n_parallel);
+    common_batch batch(ctx);

     std::vector<llama_seq_id> seq_ids(n_parallel, 0);
     for (int32_t i = 0; i < n_parallel; ++i) {
@@ -126,12 +126,12 @@ int main(int argc, char ** argv) {

     // evaluate the initial prompt
     for (size_t i = 0; i < tokens_list.size(); ++i) {
-        common_batch_add(batch, tokens_list[i], i, seq_ids, false);
+        batch.add(tokens_list[i], i, seq_ids, false);
     }
-    GGML_ASSERT(batch.n_tokens == (int) tokens_list.size());
+    GGML_ASSERT(batch.size() == (int) tokens_list.size());

     if (llama_model_has_encoder(model)) {
-        if (llama_encode(ctx, batch)) {
+        if (llama_process(ctx, LLAMA_PROCESS_TYPE_ENCODE, batch.get())) {
             LOG_ERR("%s : failed to eval\n", __func__);
             return 1;
         }
@@ -141,14 +141,14 @@ int main(int argc, char ** argv) {
             decoder_start_token_id = llama_vocab_bos(vocab);
         }

-        common_batch_clear(batch);
-        common_batch_add(batch, decoder_start_token_id, 0, seq_ids, false);
+        batch.clear();
+        batch.add(decoder_start_token_id, 0, seq_ids, false);
     }

     // llama_decode will output logits only for the last token of the prompt
-    batch.logits[batch.n_tokens - 1] = true;
+    batch.set_output(batch.size() - 1, true);

-    if (llama_decode(ctx, batch) != 0) {
+    if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
         LOG_ERR("%s: llama_decode() failed\n", __func__);
         return 1;
     }
@@ -170,16 +170,16 @@ int main(int argc, char ** argv) {

     // remember the batch index of the last token for each parallel sequence
     // we need this to determine which logits to sample from
-    std::vector<int32_t> i_batch(n_parallel, batch.n_tokens - 1);
+    std::vector<int32_t> i_batch(n_parallel, batch.size() - 1);

-    int n_cur    = batch.n_tokens;
+    int n_cur    = batch.size();
     int n_decode = 0;

     const auto t_main_start = ggml_time_us();

     while (n_cur <= n_predict) {
         // prepare the next batch
-        common_batch_clear(batch);
+        batch.clear();

         // sample the next token for each parallel sequence / stream
         for (int32_t i = 0; i < n_parallel; ++i) {
@@ -208,23 +208,23 @@ int main(int argc, char ** argv) {

             streams[i] += common_token_to_piece(ctx, new_token_id);

-            i_batch[i] = batch.n_tokens;
+            i_batch[i] = batch.size();

             // push this new token for next evaluation
-            common_batch_add(batch, new_token_id, n_cur, { i }, true);
+            batch.add(new_token_id, n_cur, i, true);

             n_decode += 1;
         }

         // all streams are finished
-        if (batch.n_tokens == 0) {
+        if (batch.size() == 0) {
             break;
         }

         n_cur += 1;

         // evaluate the current batch with the transformer model
-        if (llama_decode(ctx, batch)) {
+        if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
             LOG_ERR("%s : failed to eval, return code %d\n", __func__, 1);
             return 1;
         }
@@ -249,7 +249,6 @@ int main(int argc, char ** argv) {

     fprintf(stderr, "\n");

-    llama_batch_free(batch);

     for (auto & sampler_config : sampler_configs) {
         llama_sampler_free(sampler_config.sampler);
diff --git a/examples/debug/debug.cpp b/examples/debug/debug.cpp
index 761e7a2db..18a264b63 100644
--- a/examples/debug/debug.cpp
+++ b/examples/debug/debug.cpp
@@ -194,7 +194,8 @@ static bool run(llama_context * ctx, const common_params & params) {
         return false;
     }

-    if (llama_decode(ctx, llama_batch_get_one(tokens.data(), tokens.size()))) {
+    common_batch batch = common_batch_get_one(ctx, tokens);
+    if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
         LOG_ERR("%s : failed to eval\n", __func__);
         return false;
     }
diff --git a/examples/diffusion/diffusion.cpp b/examples/diffusion/diffusion.cpp
index 97d6b6944..0ff57d953 100644
--- a/examples/diffusion/diffusion.cpp
+++ b/examples/diffusion/diffusion.cpp
@@ -1,5 +1,7 @@
 #include "diffusion.h"

+#include "common.h"
+
 #include "log.h"

 #include <algorithm>
@@ -144,8 +146,7 @@ void diffusion_generate(llama_context *          ctx,

     struct llama_sampler * dist_sampler = llama_sampler_init_dist(params.seed);

-    llama_batch batch = llama_batch_init(params.max_length, 0, 1);
-    batch.n_tokens    = params.max_length;
+    common_batch batch(ctx);

     // Pre-allocate buffers for CFG if needed
     int32_t                  logits_size = n_vocab * params.max_length;
@@ -202,18 +203,15 @@ void diffusion_generate(llama_context *          ctx,
             }

             // Setup batch
+            batch.clear();
             for (int32_t i = 0; i < params.max_length; i++) {
-                batch.token[i]     = output_tokens[i];
-                batch.pos[i]       = i;
-                batch.n_seq_id[i]  = 1;
-                batch.seq_id[i][0] = 0;
-                batch.logits[i]    = 1;
+                batch.add(output_tokens[i], i, 0, true);
             }

             float * logits = nullptr;

             if (params.cfg_scale > 0.0f) {
-                int ret = llama_decode(ctx, batch);
+                int ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
                 if (ret != 0) {
                     LOG_ERR("Failed to generate conditional");
                     break;
@@ -227,10 +225,11 @@ void diffusion_generate(llama_context *          ctx,
                     un_x_buffer[i] = params.mask_token_id;
                 }

+                batch.clear();
                 for (int32_t i = 0; i < params.max_length; i++) {
-                    batch.token[i] = un_x_buffer[i];
+                    batch.add(un_x_buffer[i], i, 0, true);
                 }
-                ret = llama_decode(ctx, batch);
+                ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
                 if (ret != 0) {
                     LOG_ERR("Failed to generate unconditional");
                     break;
@@ -244,7 +243,7 @@ void diffusion_generate(llama_context *          ctx,
                 }
                 logits = cond_logits_buffer.data();
             } else {
-                int ret = llama_decode(ctx, batch);
+                int ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
                 if (ret != 0) {
                     LOG_ERR("%s: failed to decode at step %d, ret = %d\n", __func__, global_step, ret);
                     break;
@@ -400,7 +399,6 @@ void diffusion_generate(llama_context *          ctx,
             total_time / 1000.0 / params.steps,
             total_sampling_time / 1000.0 / params.steps);

-    llama_batch_free(batch);
     llama_sampler_free(sampler);
     llama_sampler_free(dist_sampler);

diff --git a/examples/embedding/embedding.cpp b/examples/embedding/embedding.cpp
index f6a20ef9d..a59f04cae 100644
--- a/examples/embedding/embedding.cpp
+++ b/examples/embedding/embedding.cpp
@@ -27,27 +27,27 @@ static std::vector<std::string> split_lines(const std::string & s, const std::st
     return lines;
 }

-static void batch_add_seq(llama_batch & batch, const std::vector<int32_t> & tokens, llama_seq_id seq_id) {
+static void batch_add_seq(common_batch & batch, const std::vector<int32_t> & tokens, llama_seq_id seq_id) {
     size_t n_tokens = tokens.size();
     for (size_t i = 0; i < n_tokens; i++) {
-        common_batch_add(batch, tokens[i], i, { seq_id }, true);
+        batch.add(tokens[i], i, seq_id, true);
     }
 }

-static void batch_decode(llama_context * ctx, llama_batch & batch, float * output, int n_seq, int n_embd_out, int embd_norm) {
+static void batch_decode(llama_context * ctx, common_batch & batch, float * output, int n_seq, int n_embd_out, int embd_norm) {
     const enum llama_pooling_type pooling_type = llama_pooling_type(ctx);

     // clear previous kv_cache values (irrelevant for embeddings)
     llama_memory_clear(llama_get_memory(ctx), true);

     // run model
-    LOG_INF("%s: n_tokens = %d, n_seq = %d\n", __func__, batch.n_tokens, n_seq);
-    if (llama_decode(ctx, batch) < 0) {
+    LOG_INF("%s: n_tokens = %d, n_seq = %d\n", __func__, batch.size(), n_seq);
+    if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) < 0) {
         LOG_ERR("%s : failed to process\n", __func__);
     }

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

@@ -61,8 +61,8 @@ static void batch_decode(llama_context * ctx, llama_batch & batch, float * outpu
             GGML_ASSERT(embd != NULL && "failed to get token embeddings");
         } else {
             // try to get sequence embeddings - supported only when pooling_type is not NONE
-            embd = llama_get_embeddings_seq(ctx, batch.seq_id[i][0]);
-            embd_pos = batch.seq_id[i][0];
+            embd = llama_get_embeddings_seq(ctx, batch.tokens[i].seq_id);
+            embd_pos = batch.tokens[i].seq_id;
             GGML_ASSERT(embd != NULL && "failed to get sequence embeddings");
         }

@@ -242,7 +242,7 @@ int main(int argc, char ** argv) {

     // initialize batch
     const int n_prompts = prompts.size();
-    struct llama_batch batch = llama_batch_init(n_batch, 0, 1);
+    common_batch batch(ctx);

     // count number of embeddings
     int n_embd_count = 0;
@@ -269,12 +269,12 @@ int main(int argc, char ** argv) {
         const uint64_t n_toks = inp.size();

         // encode if at capacity
-        if (batch.n_tokens + n_toks > n_batch || s >= n_seq_max) {
+        if (batch.size() + n_toks > n_batch || s >= n_seq_max) {
             float * out = emb + e * n_embd_out;
             batch_decode(ctx, batch, out, s, n_embd_out, params.embd_normalize);
-            e += pooling_type == LLAMA_POOLING_TYPE_NONE ? batch.n_tokens : s;
+            e += pooling_type == LLAMA_POOLING_TYPE_NONE ? batch.size() : s;
             s = 0;
-            common_batch_clear(batch);
+            batch.clear();
         }

         // add to batch
@@ -407,7 +407,6 @@ int main(int argc, char ** argv) {
     llama_perf_context_print(ctx);

     // clean up
-    llama_batch_free(batch);
     llama_backend_free();

     return 0;
diff --git a/examples/eval-callback/eval-callback.cpp b/examples/eval-callback/eval-callback.cpp
index 4ce8d600b..703ce130b 100644
--- a/examples/eval-callback/eval-callback.cpp
+++ b/examples/eval-callback/eval-callback.cpp
@@ -26,7 +26,8 @@ static bool run(llama_context * ctx, const common_params & params) {
         LOG_INF("  %d\n", tokens[i]);
     }

-    if (llama_decode(ctx, llama_batch_get_one(tokens.data(), tokens.size()))) {
+    common_batch batch = common_batch_get_one(ctx, tokens);
+    if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
         LOG_ERR("%s : failed to eval\n", __func__);
         return false;
     }
diff --git a/examples/idle/idle.cpp b/examples/idle/idle.cpp
index 409fd25c1..ddbda7993 100644
--- a/examples/idle/idle.cpp
+++ b/examples/idle/idle.cpp
@@ -57,12 +57,13 @@ int main(int argc, char ** argv) {
         return 1;
     }

-    llama_batch batch = llama_batch_get_one(prompt_tokens.data(), prompt_tokens.size());
-
     const int n_iters = 3;

     // warm-up
-    llama_decode(ctx, batch);
+    {
+        common_batch batch = common_batch_get_one(ctx, prompt_tokens);
+        llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+    }
     llama_memory_clear(llama_get_memory(ctx), true);
     llama_synchronize(ctx);

@@ -71,13 +72,16 @@ int main(int argc, char ** argv) {
         double t_sum2_us = 0.0;

         for (int i = 0; i < n_iters; i++) {
+            // positions continue from the memory
+            common_batch batch = common_batch_get_one(ctx, prompt_tokens);
+
             // this pause is important - it simulates "idle GPU"
             std::this_thread::sleep_for(std::chrono::milliseconds(t_pause_ms));

             const int64_t t_start_us = llama_time_us();

             // this should take constant time
-            llama_decode(ctx, batch);
+            llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
             llama_synchronize(ctx);

             const int64_t t_end_us = llama_time_us();
diff --git a/examples/llama.android/lib/src/main/cpp/ai_chat.cpp b/examples/llama.android/lib/src/main/cpp/ai_chat.cpp
index 03ab96cfd..ca6f7172d 100644
--- a/examples/llama.android/lib/src/main/cpp/ai_chat.cpp
+++ b/examples/llama.android/lib/src/main/cpp/ai_chat.cpp
@@ -35,7 +35,7 @@ constexpr float DEFAULT_SAMPLER_TEMP    = 0.3f;

 static llama_model                      * g_model;
 static llama_context                    * g_context;
-static llama_batch                        g_batch;
+static common_batch                       g_batch;
 static common_chat_templates_ptr          g_chat_templates;
 static common_sampler                   * g_sampler;

@@ -116,7 +116,7 @@ Java_com_arm_aichat_internal_InferenceEngineImpl_prepare(JNIEnv * /*env*/, jobje
     auto *context = init_context(g_model);
     if (!context) { return 1; }
     g_context = context;
-    g_batch = llama_batch_init(BATCH_SIZE, 0, 1);
+    g_batch = common_batch(context);
     g_chat_templates = common_chat_templates_init(g_model, "");
     g_sampler = new_sampler(DEFAULT_SAMPLER_TEMP);
     return 0;
@@ -164,18 +164,18 @@ Java_com_arm_aichat_internal_InferenceEngineImpl_benchModel(JNIEnv *env, jobject
     for (nri = 0; nri < nr; nri++) {
         LOGi("Benchmark prompt processing (pp = %d)", pp);

-        common_batch_clear(g_batch);
+        common_batch batch(context);

         const int n_tokens = pp;
         for (i = 0; i < n_tokens; i++) {
-            common_batch_add(g_batch, 0, i, {0}, false);
+            batch.add(0, i, 0, false);
         }

-        g_batch.logits[g_batch.n_tokens - 1] = true;
+        batch.set_output(batch.size() - 1, true);
         llama_memory_clear(llama_get_memory(context), false);

         const auto t_pp_start = ggml_time_us();
-        if (llama_decode(context, g_batch) != 0) {
+        if (llama_process(context, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
             LOGe("llama_decode() failed during prompt processing");
         }
         const auto t_pp_end = ggml_time_us();
@@ -187,12 +187,12 @@ Java_com_arm_aichat_internal_InferenceEngineImpl_benchModel(JNIEnv *env, jobject
         llama_memory_clear(llama_get_memory(context), false);
         const auto t_tg_start = ggml_time_us();
         for (i = 0; i < tg; i++) {
-            common_batch_clear(g_batch);
+            batch.clear();
             for (j = 0; j < pl; j++) {
-                common_batch_add(g_batch, 0, i, {j}, true);
+                batch.add(0, i, j, true);
             }

-            if (llama_decode(context, g_batch) != 0) {
+            if (llama_process(context, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
                 LOGe("llama_decode() failed during text generation");
             }
         }
@@ -315,7 +315,7 @@ static void reset_short_term_states() {

 static int decode_tokens_in_batches(
         llama_context *context,
-        llama_batch &batch,
+        common_batch &batch,
         const llama_tokens &tokens,
         const llama_pos start_pos,
         const bool compute_last_logit = false) {
@@ -323,7 +323,7 @@ static int decode_tokens_in_batches(
     LOGd("%s: Decode %d tokens starting at position %d", __func__, (int) tokens.size(), start_pos);
     for (int i = 0; i < (int) tokens.size(); i += BATCH_SIZE) {
         const int cur_batch_size = std::min((int) tokens.size() - i, BATCH_SIZE);
-        common_batch_clear(batch);
+        batch.clear();
         LOGv("%s: Preparing a batch size of %d starting at: %d", __func__, cur_batch_size, i);

         // Shift context if current batch cannot fit into the context
@@ -337,11 +337,11 @@ static int decode_tokens_in_batches(
             const llama_token token_id = tokens[i + j];
             const llama_pos position = start_pos + i + j;
             const bool want_logit = compute_last_logit && (i + j == tokens.size() - 1);
-            common_batch_add(batch, token_id, position, {0}, want_logit);
+            batch.add(token_id, position, 0, want_logit);
         }

         // Decode this batch
-        const int decode_result = llama_decode(context, batch);
+        const int decode_result = llama_process(context, LLAMA_PROCESS_TYPE_DECODE, batch.get());
         if (decode_result) {
             LOGe("%s: llama_decode failed w/ %d", __func__, decode_result);
             return 1;
@@ -506,9 +506,9 @@ Java_com_arm_aichat_internal_InferenceEngineImpl_generateNextToken(
     common_sampler_accept(g_sampler, new_token_id, true);

     // Populate the batch with new token, then decode
-    common_batch_clear(g_batch);
-    common_batch_add(g_batch, new_token_id, current_position, {0}, true);
-    if (llama_decode(g_context, g_batch) != 0) {
+    g_batch.clear();
+    g_batch.add(new_token_id, current_position, 0, true);
+    if (llama_process(g_context, LLAMA_PROCESS_TYPE_DECODE, g_batch.get()) != 0) {
         LOGe("%s: llama_decode() failed for generated token", __func__);
         return nullptr;
     }
@@ -553,7 +553,7 @@ Java_com_arm_aichat_internal_InferenceEngineImpl_unload(JNIEnv * /*unused*/, job
     // Free up resources
     common_sampler_free(g_sampler);
     g_chat_templates.reset();
-    llama_batch_free(g_batch);
+    g_batch = common_batch();
     llama_free(g_context);
     llama_model_free(g_model);
 }
diff --git a/examples/lookahead/lookahead.cpp b/examples/lookahead/lookahead.cpp
index b7f5c6de8..62772814f 100644
--- a/examples/lookahead/lookahead.cpp
+++ b/examples/lookahead/lookahead.cpp
@@ -101,8 +101,13 @@ int main(int argc, char ** argv) {
     const auto t_enc_start = ggml_time_us();

     // eval the prompt
-    llama_decode(ctx, llama_batch_get_one( inp.data(), n_input - 1));
-    llama_decode(ctx, llama_batch_get_one(&inp.back(),           1));
+    {
+        common_batch batch = common_batch_get_one(ctx, inp.data(), n_input - 1);
+        llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+
+        batch = common_batch_get_one(ctx, &inp.back(), 1);
+        llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+    }

     for (int s = 1; s < W + G + 1; ++s) {
         llama_memory_seq_cp(mem, 0, s, -1, -1);
@@ -124,7 +129,7 @@ int main(int argc, char ** argv) {
     // seq_id == 0           : the current input token
     // seq_id [1, W]         : tokens from the past N - 1 Jacobi iterations
     // seq_id [W + 1, W + G] : verification n-grams
-    llama_batch batch = llama_batch_init(llama_n_ctx(ctx), 0, W + G + 1);
+    common_batch batch(ctx);

     // target model sampling context
     struct common_sampler * smpl = common_sampler_init(model, params.sampling);
@@ -204,10 +209,10 @@ int main(int argc, char ** argv) {
         //                                                      V  V  V  V  V  V
         //                                                             id
         {
-            common_batch_clear(batch);
+            batch.clear();

             // current token - first token of the first level
-            common_batch_add(batch, id, n_past, seq_id_all, true);
+            batch.add(id, n_past, seq_id_all, true);

             // verification n-grams - queue this before the lookahead tokens for less KV cache fragmentation
             {
@@ -230,9 +235,9 @@ int main(int argc, char ** argv) {
                         const llama_token t = ngrams_observed.tokens[idx + j];

                         ngrams_cur[g].tokens [j + 1] = t;
-                        ngrams_cur[g].i_batch[j + 1] = batch.n_tokens;
+                        ngrams_cur[g].i_batch[j + 1] = batch.size();

-                        common_batch_add(batch, t, n_past + j + 1, { W + 1 + g }, true);
+                        batch.add(t, n_past + j + 1, W + 1 + g, true);
                     }
                 }
             }
@@ -244,18 +249,18 @@ int main(int argc, char ** argv) {
                     seq_id_look[j] = i + j + 1;
                 }

-                common_batch_add(batch, tokens_j[0][i], n_past + i, seq_id_look, false);
+                batch.add(tokens_j[0][i], n_past + i, seq_id_look, false);
             }

             // fill the rest of the levels
             for (int j = 1; j < N - 1; j++) {
                 for (int i = 0; i < W; i++) {
-                    common_batch_add(batch, tokens_j[j][i], n_past + j + i, { i + 1 }, j == N - 2);
+                    batch.add(tokens_j[j][i], n_past + j + i, i + 1, j == N - 2);
                 }
             }
         }

-        if (llama_decode(ctx, batch) != 0) {
+        if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
             LOG_ERR("\n\n%s: llama_decode failed - increase KV cache size\n", __func__);
             return 1;
         }
@@ -473,7 +478,6 @@ int main(int argc, char ** argv) {

     common_sampler_free(smpl);

-    llama_batch_free(batch);

     llama_backend_free();

diff --git a/examples/lookup/lookup.cpp b/examples/lookup/lookup.cpp
index 662105865..004062c57 100644
--- a/examples/lookup/lookup.cpp
+++ b/examples/lookup/lookup.cpp
@@ -98,8 +98,13 @@ int main(int argc, char ** argv){

     const auto t_enc_start = ggml_time_us();

-    llama_decode(ctx, llama_batch_get_one( inp.data(), n_input - 1));
-    llama_decode(ctx, llama_batch_get_one(&inp.back(),           1));
+    {
+        common_batch batch = common_batch_get_one(ctx, inp.data(), n_input - 1);
+        llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+
+        batch = common_batch_get_one(ctx, &inp.back(), 1);
+        llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+    }

     const auto t_enc_end = ggml_time_us();

@@ -115,7 +120,7 @@ int main(int argc, char ** argv){

     std::vector<llama_token> draft;

-    llama_batch batch_tgt = llama_batch_init(llama_n_ctx(ctx), 0, 1);
+    common_batch batch_tgt(ctx);

     const auto t_dec_start = ggml_time_us();

@@ -192,8 +197,8 @@ int main(int argc, char ** argv){
         // clean the cache of draft tokens that weren't accepted
         llama_memory_seq_rm(llama_get_memory(ctx), 0, n_past, -1);

-        common_batch_clear(batch_tgt);
-        common_batch_add(batch_tgt, draft[0], n_past, { 0 }, true);
+        batch_tgt.clear();
+        batch_tgt.add(draft[0], n_past, 0, true);

         // Draft already contains a single token sampled from the model:
         GGML_ASSERT(draft.size() == 1);
@@ -203,13 +208,13 @@ int main(int argc, char ** argv){
         common_ngram_cache_draft(inp, draft, n_draft, LLAMA_NGRAM_MIN, LLAMA_NGRAM_MAX, ngram_cache_context, ngram_cache_dynamic, ngram_cache_static);

         for (size_t i = 1; i < draft.size(); ++i) {
-            common_batch_add(batch_tgt, draft[i], n_past + i, { 0 }, true);
+            batch_tgt.add(draft[i], n_past + i, 0, true);
         }

         t_draft_us += ggml_time_us() - t_start_draft_us;
         n_drafted += draft.size() - 1;

-        llama_decode(ctx, batch_tgt);
+        llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_tgt.get());
         ++n_past;

         draft.erase(draft.begin());
@@ -241,7 +246,6 @@ int main(int argc, char ** argv){

     common_sampler_free(smpl);

-    llama_batch_free(batch_tgt);

     llama_backend_free();

diff --git a/examples/parallel/parallel.cpp b/examples/parallel/parallel.cpp
index 4b74540f0..de877ca36 100644
--- a/examples/parallel/parallel.cpp
+++ b/examples/parallel/parallel.cpp
@@ -224,8 +224,6 @@ int main(int argc, char ** argv) {

     LOG_INF("\n\n");

-    const int n_ctx = llama_n_ctx(ctx);
-
     if (sseed >= 0) {
         LOG_INF("%s: initializing all samplers with the same RNG seed: %d (use a negative seed to have different seeds)\n", __func__, sseed);
     } else {
@@ -252,7 +250,7 @@ int main(int argc, char ** argv) {

     // the max batch size is as large as the context to handle cases where we get very long input prompt from multiple
     // users. regardless of the size, the main loop will chunk the batch into a maximum of params.n_batch tokens at a time
-    llama_batch batch = llama_batch_init(n_ctx, 0, 1);
+    common_batch batch(ctx);

     int32_t n_total_prompt = 0;
     int32_t n_total_gen    = 0;
@@ -268,10 +266,10 @@ int main(int argc, char ** argv) {
         LOG_INF("%s: Evaluating the system prompt ...\n", __func__);

         for (int32_t i = 0; i < n_tokens_system; ++i) {
-            common_batch_add(batch, tokens_system[i], i, { 0 }, false);
+            batch.add(tokens_system[i], i, 0, false);
         }

-        if (llama_decode(ctx, batch) != 0) {
+        if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
             LOG_ERR("%s: llama_decode() failed\n", __func__);
             return 1;
         }
@@ -287,7 +285,7 @@ int main(int argc, char ** argv) {
     LOG_INF("Processing requests ...\n\n");

     while (true) {
-        common_batch_clear(batch);
+        batch.clear();

         // decode any currently ongoing sequences
         for (auto & client : clients) {
@@ -295,14 +293,14 @@ int main(int argc, char ** argv) {
                 continue;
             }

-            client.i_batch = batch.n_tokens;
+            client.i_batch = batch.size();

-            common_batch_add(batch, client.sampled, client.n_past++, { client.id + 1 }, true);
+            batch.add(client.sampled, client.n_past++, client.id + 1, true);

             client.n_decoded += 1;
         }

-        if (batch.n_tokens == 0) {
+        if (batch.size() == 0) {
             // all sequences have ended - clear the entire KV cache
             for (int i = 1; i <= n_clients; ++i) {
                 llama_memory_seq_rm(mem, i, -1, -1);
@@ -314,7 +312,7 @@ int main(int argc, char ** argv) {
         }

         // insert new sequences for decoding
-        if (cont_batching || batch.n_tokens == 0) {
+        if (cont_batching || batch.size() == 0) {
             for (auto & client : clients) {
                 if (client.seq_id == -1 && g_seq_id < n_seq) {
                     client.seq_id = g_seq_id;
@@ -350,17 +348,17 @@ int main(int argc, char ** argv) {
                     tokens_prompt = common_tokenize(ctx, client.prompt, false);

                     for (size_t i = 0; i < tokens_prompt.size(); ++i) {
-                        common_batch_add(batch, tokens_prompt[i], client.n_past++, { client.id + 1 }, false);
+                        batch.add(tokens_prompt[i], client.n_past++, client.id + 1, false);
                     }

                     // extract the logits only for the last token
-                    if (batch.n_tokens > 0) {
-                        batch.logits[batch.n_tokens - 1] = true;
+                    if (batch.size() > 0) {
+                        batch.set_output(batch.size() - 1, true);
                     }

                     client.n_prompt  = tokens_prompt.size();
                     client.n_decoded = 0;
-                    client.i_batch   = batch.n_tokens - 1;
+                    client.i_batch   = batch.size() - 1;

                     LOG_INF("\033[31mClient %3d, seq %4d, junk = %4d, prompt = %d, started decoding ...\033[0m\n", client.id, client.seq_id, n_junk_cur, client.n_prompt);

@@ -374,7 +372,7 @@ int main(int argc, char ** argv) {
             }
         }

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

@@ -383,27 +381,17 @@ int main(int argc, char ** argv) {

         int32_t i_next = 0;

-        for (int32_t i = 0; i < batch.n_tokens; i = i_next) {
+        for (int32_t i = 0; i < batch.size(); i = i_next) {
             // experiment: process in powers of 2
-            //if (i + n_batch > (int32_t) batch.n_tokens && n_batch > 32) {
+            //if (i + n_batch > (int32_t) batch.size() && n_batch > 32) {
             //    n_batch /= 2;
             //    i -= n_batch;
             //    continue;
             //}

-            const int32_t n_tokens = std::min(n_batch, batch.n_tokens - i);
-
-            llama_batch batch_view = {
-                n_tokens,
-                batch.token    + i,
-                nullptr,
-                batch.pos      + i,
-                batch.n_seq_id + i,
-                batch.seq_id   + i,
-                batch.logits   + i,
-            };
+            const int32_t n_tokens = std::min(n_batch, batch.size() - i);

-            const int ret = llama_decode(ctx, batch_view);
+            const int ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get_sub_batch(i, n_tokens));
             if (ret != 0) {
                 if (n_batch == 1 || ret < 0) {
                     // if you get here, it means the KV cache is full - try increasing it via the context size
@@ -511,7 +499,6 @@ int main(int argc, char ** argv) {
     // TODO: print sampling/grammar timings for all clients
     llama_perf_context_print(ctx);

-    llama_batch_free(batch);

     llama_backend_free();

diff --git a/examples/passkey/passkey.cpp b/examples/passkey/passkey.cpp
index 8440a2bf7..9ac8a0170 100644
--- a/examples/passkey/passkey.cpp
+++ b/examples/passkey/passkey.cpp
@@ -125,7 +125,7 @@ int main(int argc, char ** argv) {
     LOG_INF("prompt tokens: %d\n", n_tokens_all);
     //LOG_INF("prompt: %s\n", params.prompt.c_str());

-    llama_batch batch = llama_batch_init(params.n_batch, 0, 1);
+    common_batch batch(ctx);

     int n_past = 0;

@@ -144,17 +144,17 @@ int main(int argc, char ** argv) {
             n_past = llama_memory_seq_pos_max(mem, 0) + 1;
         }

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

         for (int j = 0; j < n_batch && i + j < n_tokens_all; j++) {
-            common_batch_add(batch, tokens_list[i + j], n_past++, { 0 }, false);
+            batch.add(tokens_list[i + j], n_past++, 0, false);
         }

         if (i + n_batch >= n_tokens_all) {
-            batch.logits[batch.n_tokens - 1] = true;
+            batch.set_output(batch.size() - 1, true);
         }

-        if (llama_decode(ctx, batch) != 0) {
+        if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
             LOG_INF("%s: llama_decode() failed\n", __func__);
             return 1;
         }
@@ -176,17 +176,17 @@ int main(int argc, char ** argv) {

         n_past = llama_memory_seq_pos_max(mem, 0) + 1;

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

         for (int j = 0; j < n_batch && i + j < n_tokens_all; j++) {
-            common_batch_add(batch, tokens_list[i + j], n_past++, { 0 }, false);
+            batch.add(tokens_list[i + j], n_past++, 0, false);
         }

         if (i + n_batch >= n_tokens_all) {
-            batch.logits[batch.n_tokens - 1] = true;
+            batch.set_output(batch.size() - 1, true);
         }

-        if (llama_decode(ctx, batch) != 0) {
+        if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
             LOG_ERR("%s: llama_decode() failed\n", __func__);
             return 1;
         }
@@ -223,7 +223,7 @@ int main(int argc, char ** argv) {
     while (n_cur <= n_len) {
         // sample the next token
         {
-            const llama_token new_token_id = llama_sampler_sample(smpl, ctx, batch.n_tokens - 1);
+            const llama_token new_token_id = llama_sampler_sample(smpl, ctx, batch.size() - 1);

             // is it an end of generation?
             if (llama_vocab_is_eog(vocab, new_token_id) || n_cur == n_len) {
@@ -237,16 +237,16 @@ int main(int argc, char ** argv) {
             n_decode += 1;

             // prepare the next batch
-            common_batch_clear(batch);
+            batch.clear();

             // push this new token for next evaluation
-            common_batch_add(batch, new_token_id, n_past++, { 0 }, true);
+            batch.add(new_token_id, n_past++, 0, true);
         }

         n_cur += 1;

         // evaluate the current batch with the transformer model
-        if (llama_decode(ctx, batch)) {
+        if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
             LOG_ERR("%s : failed to eval, return code %d\n", __func__, 1);
             return 1;
         }
@@ -266,7 +266,6 @@ int main(int argc, char ** argv) {

     llama_sampler_free(smpl);

-    llama_batch_free(batch);

     llama_free(ctx);
     llama_model_free(model);
diff --git a/examples/retrieval/retrieval.cpp b/examples/retrieval/retrieval.cpp
index 7d93ab117..8793e5751 100644
--- a/examples/retrieval/retrieval.cpp
+++ b/examples/retrieval/retrieval.cpp
@@ -75,30 +75,30 @@ static std::vector<chunk> chunk_file(const std::string & filename, int chunk_siz
     return chunks;
 }

-static void batch_add_seq(llama_batch & batch, const std::vector<int32_t> & tokens, llama_seq_id seq_id) {
+static void batch_add_seq(common_batch & batch, const std::vector<int32_t> & tokens, llama_seq_id seq_id) {
     size_t n_tokens = tokens.size();
     for (size_t i = 0; i < n_tokens; i++) {
-        common_batch_add(batch, tokens[i], i, { seq_id }, true);
+        batch.add(tokens[i], i, seq_id, true);
     }
 }

-static void batch_process(llama_context * ctx, llama_batch & batch, float * output, int n_seq, int n_embd) {
+static void batch_process(llama_context * ctx, common_batch & batch, float * output, int n_seq, int n_embd) {
     // clear previous kv_cache values (irrelevant for embeddings)
     llama_memory_clear(llama_get_memory(ctx), false);

     // run model
-    LOG_INF("%s: n_tokens = %d, n_seq = %d\n", __func__, batch.n_tokens, n_seq);
-    if (llama_decode(ctx, batch) < 0) {
+    LOG_INF("%s: n_tokens = %d, n_seq = %d\n", __func__, batch.size(), n_seq);
+    if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) < 0) {
         LOG_ERR("%s : failed to process\n", __func__);
     }

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

         // try to get sequence embeddings - supported only when pooling_type is not NONE
-        const float * embd = llama_get_embeddings_seq(ctx, batch.seq_id[i][0]);
+        const float * embd = llama_get_embeddings_seq(ctx, batch.tokens[i].seq_id);
         if (embd == NULL) {
             embd = llama_get_embeddings_ith(ctx, i);
             if (embd == NULL) {
@@ -107,7 +107,7 @@ static void batch_process(llama_context * ctx, llama_batch & batch, float * outp
             }
         }

-        float * out = output + batch.seq_id[i][0] * n_embd;
+        float * out = output + batch.tokens[i].seq_id * n_embd;
         common_embd_normalize(embd, out, n_embd, 2);
     }
 }
@@ -217,7 +217,7 @@ int main(int argc, char ** argv) {

     // initialize batch
     const int n_chunks = chunks.size();
-    struct llama_batch batch = llama_batch_init(n_batch, 0, 1);
+    common_batch batch(ctx);

     // allocate output
     const int n_embd_out = llama_model_n_embd_out(model);
@@ -234,10 +234,10 @@ int main(int argc, char ** argv) {
         const uint64_t n_toks = inp.size();

         // encode if at capacity
-        if (batch.n_tokens + n_toks > n_batch || s >= llama_n_seq_max(ctx)) {
+        if (batch.size() + n_toks > n_batch || s >= llama_n_seq_max(ctx)) {
             float * out = emb + p * n_embd_out;
             batch_process(ctx, batch, out, s, n_embd_out);
-            common_batch_clear(batch);
+            batch.clear();
             p += s;
             s = 0;
         }
@@ -258,7 +258,7 @@ int main(int argc, char ** argv) {
         chunks[i].tokens.clear();
     }

-    struct llama_batch query_batch = llama_batch_init(n_batch, 0, 1);
+    common_batch query_batch(ctx);

     // start loop, receive query and return top k similar chunks based on cosine similarity
     std::string query;
@@ -272,7 +272,7 @@ int main(int argc, char ** argv) {
         std::vector<float> query_emb(n_embd_out, 0);
         batch_process(ctx, query_batch, query_emb.data(), 1, n_embd_out);

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

         // compute cosine similarities
         {
@@ -302,6 +302,5 @@ int main(int argc, char ** argv) {
     llama_perf_context_print(ctx);

     // clean up
-    llama_batch_free(query_batch);
     llama_backend_free();
 }
diff --git a/examples/simple-chat/simple-chat.cpp b/examples/simple-chat/simple-chat.cpp
index 30a0966e0..0cad9652c 100644
--- a/examples/simple-chat/simple-chat.cpp
+++ b/examples/simple-chat/simple-chat.cpp
@@ -6,6 +6,17 @@
 #include <string>
 #include <vector>

+// fill the batch with tokens at consecutive positions starting from pos_0, output logits only for the last one
+static void batch_set_tokens(llama_batch_ext * batch, const llama_token * tokens, int32_t n_tokens, llama_pos pos_0) {
+    llama_batch_ext_clear(batch);
+    for (int32_t i = 0; i < n_tokens; ++i) {
+        const int32_t idx = llama_batch_ext_add_token(batch, 0, tokens[i]);
+        const llama_pos pos = pos_0 + i;
+        llama_batch_ext_set_pos(batch, idx, &pos);
+    }
+    llama_batch_ext_set_output_logits(batch, n_tokens - 1, true);
+}
+
 static void print_usage(int, char ** argv) {
     printf("\nexample usage:\n");
     printf("\n    %s -m model.gguf [-c context_size] [-ngl n_gpu_layers]\n", argv[0]);
@@ -96,6 +107,8 @@ int main(int argc, char ** argv) {
     llama_sampler_chain_add(smpl, llama_sampler_init_temp(0.8f));
     llama_sampler_chain_add(smpl, llama_sampler_init_dist(LLAMA_DEFAULT_SEED));

+    llama_batch_ext * batch = llama_batch_ext_init(ctx);
+
     // helper function to evaluate a prompt and generate a response
     auto generate = [&](const std::string & prompt) {
         std::string response;
@@ -109,20 +122,25 @@ int main(int argc, char ** argv) {
             GGML_ABORT("failed to tokenize the prompt\n");
         }

-        // prepare a batch for the prompt
-        llama_batch batch = llama_batch_get_one(prompt_tokens.data(), prompt_tokens.size());
+        // the tokens to evaluate next: the prompt, then the sampled token
+        const llama_token * tokens = prompt_tokens.data();
+        int n_tokens = prompt_tokens.size();
+
         llama_token new_token_id;
         while (true) {
             // check if we have enough space in the context to evaluate this batch
             int n_ctx = llama_n_ctx(ctx);
             int n_ctx_used = llama_memory_seq_pos_max(llama_get_memory(ctx), 0) + 1;
-            if (n_ctx_used + batch.n_tokens > n_ctx) {
+            if (n_ctx_used + n_tokens > n_ctx) {
                 printf("\033[0m\n");
                 fprintf(stderr, "context size exceeded\n");
                 exit(0);
             }

-            int ret = llama_decode(ctx, batch);
+            // positions continue from the memory
+            batch_set_tokens(batch, tokens, n_tokens, n_ctx_used);
+
+            int ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch);
             if (ret != 0) {
                 GGML_ABORT("failed to decode, ret = %d\n", ret);
             }
@@ -147,7 +165,8 @@ int main(int argc, char ** argv) {
             response += piece;

             // prepare the next batch with the sampled token
-            batch = llama_batch_get_one(&new_token_id, 1);
+            tokens   = &new_token_id;
+            n_tokens = 1;
         }

         return response;
@@ -201,6 +220,7 @@ int main(int argc, char ** argv) {
     for (auto & msg : messages) {
         free(const_cast<char *>(msg.content));
     }
+    llama_batch_ext_free(batch);
     llama_sampler_free(smpl);
     llama_free(ctx);
     llama_model_free(model);
diff --git a/examples/simple/simple.cpp b/examples/simple/simple.cpp
index 982a4d860..ebb1969fb 100644
--- a/examples/simple/simple.cpp
+++ b/examples/simple/simple.cpp
@@ -5,6 +5,17 @@
 #include <string>
 #include <vector>

+// fill the batch with tokens at consecutive positions starting from pos_0, output logits only for the last one
+static void batch_set_tokens(llama_batch_ext * batch, const llama_token * tokens, int32_t n_tokens, llama_pos pos_0) {
+    llama_batch_ext_clear(batch);
+    for (int32_t i = 0; i < n_tokens; ++i) {
+        const int32_t idx = llama_batch_ext_add_token(batch, 0, tokens[i]);
+        const llama_pos pos = pos_0 + i;
+        llama_batch_ext_set_pos(batch, idx, &pos);
+    }
+    llama_batch_ext_set_output_logits(batch, n_tokens - 1, true);
+}
+
 static void print_usage(int, char ** argv) {
     printf("\nexample usage:\n");
     printf("\n    %s -m model.gguf [-n n_predict] [-ngl n_gpu_layers] [prompt]\n", argv[0]);
@@ -144,10 +155,13 @@ int main(int argc, char ** argv) {

     // prepare a batch for the prompt

-    llama_batch batch = llama_batch_get_one(prompt_tokens.data(), prompt_tokens.size());
+    llama_batch_ext * batch = llama_batch_ext_init(ctx);
+    int n_tokens = n_prompt; // number of tokens in the current batch
+
+    batch_set_tokens(batch, prompt_tokens.data(), n_prompt, 0);

     if (llama_model_has_encoder(model)) {
-        if (llama_encode(ctx, batch)) {
+        if (llama_process(ctx, LLAMA_PROCESS_TYPE_ENCODE, batch)) {
             fprintf(stderr, "%s : failed to eval\n", __func__);
             return 1;
         }
@@ -157,7 +171,8 @@ int main(int argc, char ** argv) {
             decoder_start_token_id = llama_vocab_bos(vocab);
         }

-        batch = llama_batch_get_one(&decoder_start_token_id, 1);
+        batch_set_tokens(batch, &decoder_start_token_id, 1, 0);
+        n_tokens = 1;
     }

     // main loop
@@ -166,14 +181,14 @@ int main(int argc, char ** argv) {
     int n_decode = 0;
     llama_token new_token_id;

-    for (int n_pos = 0; n_pos + batch.n_tokens < n_prompt + n_predict; ) {
+    for (int n_pos = 0; n_pos + n_tokens < n_prompt + n_predict; ) {
         // evaluate the current batch with the transformer model
-        if (llama_decode(ctx, batch)) {
+        if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch)) {
             fprintf(stderr, "%s : failed to eval, return code %d\n", __func__, 1);
             return 1;
         }

-        n_pos += batch.n_tokens;
+        n_pos += n_tokens;

         // sample the next token
         {
@@ -195,7 +210,8 @@ int main(int argc, char ** argv) {
             fflush(stdout);

             // prepare the next batch with the sampled token
-            batch = llama_batch_get_one(&new_token_id, 1);
+            batch_set_tokens(batch, &new_token_id, 1, n_pos);
+            n_tokens = 1;

             n_decode += 1;
         }
@@ -213,6 +229,7 @@ int main(int argc, char ** argv) {
     llama_perf_context_print(ctx);
     fprintf(stderr, "\n");

+    llama_batch_ext_free(batch);
     llama_sampler_free(smpl);
     llama_free(ctx);
     llama_model_free(model);
diff --git a/examples/speculative-simple/speculative-simple.cpp b/examples/speculative-simple/speculative-simple.cpp
index 81aa106f1..08a1f2a88 100644
--- a/examples/speculative-simple/speculative-simple.cpp
+++ b/examples/speculative-simple/speculative-simple.cpp
@@ -125,12 +125,12 @@ int main(int argc, char ** argv) {

     // eval the prompt on the target and feed it to the speculative implementation(s)
     {
-        llama_batch batch_prompt = llama_batch_init(inp.size(), 0, 1);
+        common_batch batch_prompt(ctx_tgt);
         for (size_t i = 0; i < inp.size() - 1; ++i) {
-            common_batch_add(batch_prompt, inp[i], i, { seq_id }, false);
+            batch_prompt.add(inp[i], i, seq_id, false);
         }

-        llama_decode(ctx_tgt, batch_prompt);
+        llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, batch_prompt.get());

         if (!common_speculative_process(spec, batch_prompt)) {
             LOG_ERR("%s", "failed to process speculative prompt\n");
@@ -149,7 +149,7 @@ int main(int argc, char ** argv) {

     common_speculative_begin(spec, seq_id, prompt_tgt);

-    llama_batch batch_tgt = llama_batch_init(llama_n_batch(ctx_tgt), 0, 1);
+    common_batch batch_tgt(ctx_tgt);

     llama_tokens draft;

@@ -219,17 +219,17 @@ int main(int argc, char ** argv) {
         }

         // always have a token to evaluate from before - id_last
-        common_batch_clear(batch_tgt);
-        common_batch_add  (batch_tgt, id_last, n_past++, { seq_id }, true);
+        batch_tgt.clear();
+        batch_tgt.add(id_last, n_past++, seq_id, true);

         // evaluate the target model on [id_last, draft0, draft1, ..., draftN-1]
         {
             for (size_t i = 0; i < draft.size(); ++i) {
-                common_batch_add(batch_tgt, draft[i], n_past + i, { seq_id }, true);
+                batch_tgt.add(draft[i], n_past + i, seq_id, true);
             }


-            llama_decode(ctx_tgt, batch_tgt);
+            llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, batch_tgt.get());
         }

         // feed the batch to the speculative implementation(s) - this drives the draft model, MTP, Eagle3, etc.
@@ -364,7 +364,6 @@ int main(int argc, char ** argv) {
     LOG_INF("target:\n\n");
     common_perf_print(ctx_tgt, smpl.get());

-    llama_batch_free(batch_tgt);

     common_speculative_free(spec);

diff --git a/examples/speculative/speculative.cpp b/examples/speculative/speculative.cpp
index 17071aa05..1bb47594b 100644
--- a/examples/speculative/speculative.cpp
+++ b/examples/speculative/speculative.cpp
@@ -190,9 +190,16 @@ int main(int argc, char ** argv) {
     const auto t_enc_start = ggml_time_us();

     // eval the prompt with both models
-    llama_decode(ctx_tgt, llama_batch_get_one( inp.data(), n_input - 1));
-    llama_decode(ctx_tgt, llama_batch_get_one(&inp.back(),           1));
-    llama_decode(ctx_dft, llama_batch_get_one( inp.data(), n_input));
+    {
+        common_batch batch = common_batch_get_one(ctx_tgt, inp.data(), n_input - 1);
+        llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+
+        batch = common_batch_get_one(ctx_tgt, &inp.back(), 1);
+        llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+
+        batch = common_batch_get_one(ctx_dft, inp.data(), n_input);
+        llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+    }

     const auto t_enc_end = ggml_time_us();

@@ -223,8 +230,8 @@ int main(int argc, char ** argv) {
         drafts[s].smpl = common_sampler_init(model_dft, params.sampling);
     }

-    llama_batch batch_dft = llama_batch_init(llama_n_batch(ctx_dft), 0, 1);
-    llama_batch batch_tgt = llama_batch_init(llama_n_batch(ctx_tgt), 0, n_seq_dft);
+    common_batch batch_dft(ctx_dft);
+    common_batch batch_tgt(ctx_tgt);

     const auto t_dec_start = ggml_time_us();

@@ -465,12 +472,12 @@ int main(int argc, char ** argv) {
             drafts[0].dists.push_back(std::vector<llama_token_data>());
             drafts[0].i_batch_tgt.push_back(0);

-            common_batch_clear(batch_dft);
-            common_batch_add  (batch_dft, token_id, n_past_dft, { 0 }, true);
+            batch_dft.clear();
+            batch_dft.add(token_id, n_past_dft, 0, true);

             llama_memory_seq_rm(mem_dft, 0, n_past_dft, -1);
             // LOG_DBG("dft batch: %s\n", LOG_BATCH_TOSTR_PRETTY(ctx_dft, batch_dft).c_str());
-            llama_decode(ctx_dft, batch_dft);
+            llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch_dft.get());

             ++n_past_dft;
         }
@@ -495,12 +502,12 @@ int main(int argc, char ** argv) {
         drafts[0].drafting    = true;
         drafts[0].i_batch_dft = 0;

-        common_batch_clear(batch_tgt);
-        common_batch_add  (batch_tgt, drafts[0].tokens[0], n_past_tgt, { 0 }, true);
+        batch_tgt.clear();
+        batch_tgt.add(drafts[0].tokens[0], n_past_tgt, 0, true);

         // sample n_draft tokens from the draft model using tree-based sampling
         for (int i = 0; i < n_draft; ++i) {
-            batch_dft.n_tokens = 0;
+            batch_dft.clear();

             for (int s = 0; s < n_seq_dft; ++s) {
                 drafts[s].skip = false;
@@ -531,14 +538,8 @@ int main(int argc, char ** argv) {
                         llama_memory_seq_cp(mem_dft, s, n_seq_cur, -1, -1);

                         // all previous tokens from this branch are now also part of the new branch
-                        for (int t = 0; t < batch_tgt.n_tokens; ++t) {
-                            for (int p = 0; p < batch_tgt.n_seq_id[t]; ++p) {
-                                if (batch_tgt.seq_id[t][p] == s) {
-                                    batch_tgt.seq_id[t][batch_tgt.n_seq_id[t]] = n_seq_cur;
-                                    batch_tgt.n_seq_id[t]++;
-                                    break;
-                                }
-                            }
+                        for (int t : drafts[s].i_batch_tgt) {
+                            batch_tgt.add_seq(t, n_seq_cur);
                         }

                         // copy the draft state
@@ -577,32 +578,32 @@ int main(int argc, char ** argv) {
                     drafts[s].dists.push_back({cur_p->data, cur_p->data + cur_p->size});

                     // add unique drafted tokens to the target batch
-                    drafts[s].i_batch_tgt.push_back(batch_tgt.n_tokens);
+                    drafts[s].i_batch_tgt.push_back(batch_tgt.size());

-                    common_batch_add(batch_tgt, id, n_past_tgt + i + 1, { s }, true);
+                    batch_tgt.add(id, n_past_tgt + i + 1, s, true);

                     // add the token to the batch for batched decoding with the draft model
-                    drafts[s].i_batch_dft = batch_dft.n_tokens;
+                    drafts[s].i_batch_dft = batch_dft.size();

-                    common_batch_add(batch_dft, id, n_past_cur, { s }, true);
+                    batch_dft.add(id, n_past_cur, s, true);

-                    if (batch_tgt.n_tokens > n_draft) {
+                    if (batch_tgt.size() > n_draft) {
                         drafts[s].drafting = false;
                     }
                 }
             }

             // no sequence is drafting anymore
-            if (batch_dft.n_tokens == 0) {
+            if (batch_dft.size() == 0) {
                 break;
             }

             // evaluate the drafted tokens on the draft model
-            llama_decode(ctx_dft, batch_dft);
+            llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch_dft.get());
             ++n_past_cur;
             ++n_drafted;

-            if (batch_tgt.n_tokens > n_draft) {
+            if (batch_tgt.size() > n_draft) {
                 break;
             }
         }
@@ -615,7 +616,7 @@ int main(int argc, char ** argv) {
             }

             // LOG_DBG("target batch: %s\n", LOG_BATCH_TOSTR_PRETTY(ctx_tgt, batch_tgt).c_str());
-            llama_decode(ctx_tgt, batch_tgt);
+            llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, batch_tgt.get());
             ++n_past_tgt;
         }

@@ -658,7 +659,6 @@ int main(int argc, char ** argv) {
         common_sampler_free(drafts[s].smpl);
     }

-    llama_batch_free(batch_dft);

     llama_backend_free();

diff --git a/tests/test-backend-sampler.cpp b/tests/test-backend-sampler.cpp
index c23e7248d..56736ac46 100644
--- a/tests/test-backend-sampler.cpp
+++ b/tests/test-backend-sampler.cpp
@@ -129,7 +129,7 @@ struct test_context {
         GGML_ASSERT(ctx);

         last_batch_info.clear();
-        llama_batch batch = llama_batch_init(512, 0, prompts.size());
+        common_batch batch(ctx.get());

         for (const auto & [seq_id, prompt] : prompts) {
             std::vector<llama_token> tokens;
@@ -141,7 +141,6 @@ struct test_context {
                                            false, false);
             if (n_tokens < 0) {
                 fprintf(stderr, "Warning: tokenization failed for seq_id %d\n", seq_id);
-                llama_batch_free(batch);
                 return false;
             }

@@ -155,7 +154,7 @@ struct test_context {

             int32_t start_pos = seq_positions[seq_id];
             for (size_t i = 0; i < tokens.size(); i++) {
-                common_batch_add(batch, tokens[i], start_pos + i, { seq_id }, i == tokens.size() - 1);
+                batch.add(tokens[i], start_pos + i, seq_id, i == tokens.size() - 1);
             }

             seq_positions[seq_id] = start_pos + tokens.size();
@@ -163,31 +162,18 @@ struct test_context {


         printf("Batch contents:\n");
-        printf("n_tokens: %d\n", batch.n_tokens);
-        for (int i = 0; i < batch.n_tokens; i++) {
-            printf("token[%d]: tok=%-5d, pos=%d, n_seq_id=%d, seq_ids=[", i, batch.token[i], batch.pos[i], batch.n_seq_id[i]);
-
-            for (int j = 0; j < batch.n_seq_id[i]; j++) {
-                printf("%d%s", batch.seq_id[i][j], j < batch.n_seq_id[i]-1 ? ", " : "");
-            }
-            printf("], logits=%d\n", batch.logits[i]);
+        printf("n_tokens: %d\n", batch.size());
+        for (int i = 0; i < batch.size(); i++) {
+            const auto & t = batch.tokens[i];
+            printf("token[%d]: tok=%-5d, pos=%d, seq_id=%d, logits=%d\n", i, t.id, t.pos[0], t.seq_id, t.output);
         }

-        if (llama_decode(ctx.get(), batch) != 0) {
+        if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
             fprintf(stderr, "Warning: llama_decode failed\n");
-            llama_batch_free(batch);
             return false;
         }

-        // Build mapping from seq id to batch token idx
-        for (int i = 0; i < batch.n_tokens; i++) {
-            if (batch.logits[i]) {
-                llama_seq_id seq_id = batch.seq_id[i][0];
-                last_batch_info[seq_id] = i;
-            }
-        }
-
-        llama_batch_free(batch);
+        update_batch_info(batch);
         return true;
     }

@@ -200,11 +186,12 @@ struct test_context {
         return it->second;
     }

-    void update_batch_info(const llama_batch & batch) {
+    // build mapping from seq id to batch token idx
+    void update_batch_info(const common_batch & batch) {
         last_batch_info.clear();
-        for (int i = 0; i < batch.n_tokens; i++) {
-            if (batch.logits[i]) {
-                llama_seq_id cur_seq = batch.seq_id[i][0];
+        for (int i = 0; i < batch.size(); i++) {
+            if (batch.tokens[i].output) {
+                llama_seq_id cur_seq = batch.tokens[i].seq_id;
                 last_batch_info[cur_seq] = i;
             }
         }
@@ -213,20 +200,18 @@ struct test_context {
     bool decode_token(llama_token token, llama_seq_id seq_id = 0) {
         GGML_ASSERT(ctx);

-        llama_batch batch = llama_batch_init(1, 0, 1);
+        common_batch batch(ctx.get());
         int32_t pos = seq_positions[seq_id];
-        common_batch_add(batch, token, pos, { seq_id }, true);
+        batch.add(token, pos, seq_id, true);

-        if (llama_decode(ctx.get(), batch) != 0) {
+        if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
             fprintf(stderr, "Warning: llama_decode failed for token %d in seq %d\n", token, seq_id);
-            llama_batch_free(batch);
             return false;
         }

         update_batch_info(batch);

         seq_positions[seq_id]++;
-        llama_batch_free(batch);

         return true;
     }
@@ -234,16 +219,15 @@ struct test_context {
     bool decode_tokens(const std::map<llama_seq_id, llama_token> & seq_tokens) {
         GGML_ASSERT(ctx);

-        llama_batch batch = llama_batch_init(seq_tokens.size(), 0, seq_tokens.size());
+        common_batch batch(ctx.get());

         for (const auto & [seq_id, token] : seq_tokens) {
             int32_t pos = seq_positions[seq_id];
-            common_batch_add(batch, token, pos, { seq_id }, true);
+            batch.add(token, pos, seq_id, true);
         }

-        if (llama_decode(ctx.get(), batch) != 0) {
+        if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
             fprintf(stderr, "Warning: llama_decode failed for batch tokens\n");
-            llama_batch_free(batch);
             return false;
         }

@@ -253,8 +237,6 @@ struct test_context {

         update_batch_info(batch);

-        llama_batch_free(batch);
-
         return true;
     }

@@ -1607,18 +1589,16 @@ static void test_backend_multi_output_limit(const test_params & params) {
     std::vector<llama_sampler_seq_config> configs = {{ seq_id, chain.get() }};
     test_context test_ctx(params, configs, 1, 3, 0, 2);

-    llama_batch batch = llama_batch_init(3, 0, 1);
+    common_batch batch(test_ctx.ctx.get());
     for (int i = 0; i < 3; ++i) {
-        common_batch_add(batch, llama_vocab_bos(test_ctx.vocab), i, { seq_id }, true);
+        batch.add(llama_vocab_bos(test_ctx.vocab), i, seq_id, true);
     }

     printf(">>> test_backend_multi_output_limit expected error start:\n");
-    const int ret = llama_decode(test_ctx.ctx.get(), batch);
+    const int ret = llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get());
     GGML_ASSERT(ret != 0 && "llama_decode should reject outputs above the per-sequence limit");
     printf("<<< test_backend_multi_output_limit expected error end.\n");

-    llama_batch_free(batch);
-
     printf("backend multi-output limit test PASSED\n");
 }

@@ -1649,14 +1629,22 @@ static void test_backend_multi_sequence_multi_output_dist(const test_params & pa
         { llama_vocab_eos(vocab), llama_vocab_bos(vocab) },
     };

-    llama_batch batch = llama_batch_init(4, 0, 1);
-    for (int pos = 0; pos < 2; ++pos) {
-        common_batch_add(batch, seq_tokens[0][pos], pos, { 0 }, true);
-        common_batch_add(batch, seq_tokens[1][pos], pos, { 1 }, true);
-    }
+    // a batch belongs to one context, so it is built per context
+    auto make_batch = [&](llama_context * ctx) {
+        common_batch batch(ctx);
+        for (int pos = 0; pos < 2; ++pos) {
+            batch.add(seq_tokens[0][pos], pos, 0, true);
+            batch.add(seq_tokens[1][pos], pos, 1, true);
+        }
+        return batch;
+    };

-    GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
-    GGML_ASSERT(llama_decode(reference_ctx.ctx.get(), batch) == 0);
+    common_batch batch = make_batch(test_ctx.ctx.get());
+    GGML_ASSERT(llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0);
+    {
+        common_batch batch_ref = make_batch(reference_ctx.ctx.get());
+        GGML_ASSERT(llama_process(reference_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch_ref.get()) == 0);
+    }

     std::mt19937 reference_rngs[] = {
         std::mt19937(seeds[0]),
@@ -1664,8 +1652,8 @@ static void test_backend_multi_sequence_multi_output_dist(const test_params & pa
     };
     std::uniform_real_distribution<double> reference_dist(0.0, 1.0);

-    for (int i = 0; i < batch.n_tokens; ++i) {
-        const llama_seq_id seq_id = batch.seq_id[i][0];
+    for (int i = 0; i < batch.size(); ++i) {
+        const llama_seq_id seq_id = batch.tokens[i].seq_id;
         GGML_ASSERT(seq_id == 0 || seq_id == 1);

         llama_sampler * chain = seq_id == 0 ? chain_0.get() : chain_1.get();
@@ -1706,8 +1694,6 @@ static void test_backend_multi_sequence_multi_output_dist(const test_params & pa
         GGML_ASSERT(rnd <= cumsum_sampled + 1e-4f);
     }

-    llama_batch_free(batch);
-
     printf("backend multi-sequence multi-output dist test PASSED\n");
 }

@@ -1750,33 +1736,28 @@ static void test_backend_multi_output_dist_transaction(const test_params & param

     int32_t pos = 0;
     auto decode = [&]() {
-        llama_batch batch = llama_batch_init(3, 0, 1);
+        common_batch batch(test_ctx.ctx.get());
         for (int32_t i = 0; i < 3; ++i) {
-            common_batch_add(batch, llama_vocab_bos(vocab), pos++, { seq_id }, true);
+            batch.add(llama_vocab_bos(vocab), pos++, seq_id, true);
         }
-        GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
-        return batch;
+        GGML_ASSERT(llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0);
     };

-    llama_batch batch = decode();
+    decode();
     verify_random(0, randoms[0], false);
-    llama_batch_free(batch);

-    batch = decode();
+    decode();
     verify_random(0, randoms[0]);
     verify_random(1, randoms[1]);
-    llama_batch_free(batch);

-    batch = decode();
+    decode();
     llama_sampler_ptr saved(llama_sampler_clone(chain.get()));
     verify_random(0, randoms[2]);
-    llama_batch_free(batch);

     llama_sampler_copy(saved.get(), chain.get());

-    batch = decode();
+    decode();
     verify_random(0, randoms[2]);
-    llama_batch_free(batch);

     printf("backend multi-output dist transaction test PASSED\n");
 }
@@ -1817,19 +1798,23 @@ static void test_backend_multi_output_sampling_chain(const test_params & params)
     llama_sampler_ptr reference_temp(llama_sampler_init_temp(temp));
     std::vector<llama_token_data> reference_data(n_vocab);

-    auto make_batch = [&](int32_t pos) {
-        llama_batch batch = llama_batch_init(2, 0, 1);
+    // a batch belongs to one context, so it is built per context
+    auto make_batch = [&](llama_context * ctx, int32_t pos) {
+        common_batch batch(ctx);
         for (int i = 0; i < 2; ++i) {
-            common_batch_add(batch, llama_vocab_bos(vocab), pos + i, { seq_id }, true);
+            batch.add(llama_vocab_bos(vocab), pos + i, seq_id, true);
         }
         return batch;
     };

-    llama_batch batch = make_batch(0);
-    GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
-    GGML_ASSERT(llama_decode(reference_ctx.ctx.get(), batch) == 0);
+    common_batch batch = make_batch(test_ctx.ctx.get(), 0);
+    GGML_ASSERT(llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0);
+    {
+        common_batch batch_ref = make_batch(reference_ctx.ctx.get(), 0);
+        GGML_ASSERT(llama_process(reference_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch_ref.get()) == 0);
+    }

-    for (int i = 0; i < batch.n_tokens; ++i) {
+    for (int i = 0; i < batch.size(); ++i) {
         const llama_token backend_token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), i);
         const float * sampled_logits = llama_get_sampled_logits_ith(test_ctx.ctx.get(), i);
         const float * sampled_probs = llama_get_sampled_probs_ith(test_ctx.ctx.get(), i);
@@ -1922,11 +1907,8 @@ static void test_backend_multi_output_sampling_chain(const test_params & params)
         GGML_ASSERT(std::fabs(prob_sum - 1.0f) <= 1e-3f);
     }

-    llama_batch_free(batch);
-
-    batch = make_batch(2);
-    GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
-    llama_batch_free(batch);
+    batch = make_batch(test_ctx.ctx.get(), 2);
+    GGML_ASSERT(llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0);

     printf("backend multi-output sampling chain test PASSED\n");
 }
@@ -1950,17 +1932,15 @@ static void test_backend_multi_output_cpu_suffix(const test_params & params) {
         std::vector<llama_sampler_seq_config> configs = {{ seq_id, chain.get() }};
         test_context test_ctx(params, configs, 1, 1, 0, 4);

-        llama_batch batch = llama_batch_init(1, 0, 1);
-        common_batch_add(batch, llama_vocab_bos(vocab), 0, { seq_id }, true);
-        GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
+        common_batch batch(test_ctx.ctx.get());
+        batch.add(llama_vocab_bos(vocab), 0, seq_id, true);
+        GGML_ASSERT(llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0);

         GGML_ASSERT(sampler_ctx->backend_initialized);
         GGML_ASSERT(sampler_ctx->backend_outputs_max_per_seq == 1);
         GGML_ASSERT(sampler_ctx->backend_apply_count > 0);
         GGML_ASSERT(sampler_ctx->apply_count == 0);
         GGML_ASSERT(llama_get_sampled_token_ith(test_ctx.ctx.get(), 0) != LLAMA_TOKEN_NULL);
-
-        llama_batch_free(batch);
     }

     {
@@ -1969,25 +1949,23 @@ static void test_backend_multi_output_cpu_suffix(const test_params & params) {
         std::vector<llama_sampler_seq_config> configs = {{ seq_id, chain.get() }};
         test_context test_ctx(params, configs, 1, 2, 0, 0);

-        llama_batch batch = llama_batch_init(2, 0, 1);
+        common_batch batch(test_ctx.ctx.get());
         for (int i = 0; i < 2; ++i) {
-            common_batch_add(batch, llama_vocab_bos(vocab), i, { seq_id }, true);
+            batch.add(llama_vocab_bos(vocab), i, seq_id, true);
         }
-        GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
+        GGML_ASSERT(llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0);

         GGML_ASSERT(!sampler_ctx->backend_initialized);
         GGML_ASSERT(sampler_ctx->backend_outputs_max_per_seq == 2);
         GGML_ASSERT(sampler_ctx->backend_apply_count == 0);
-        for (int i = 0; i < batch.n_tokens; ++i) {
+        for (int i = 0; i < batch.size(); ++i) {
             GGML_ASSERT(llama_get_sampled_token_ith(test_ctx.ctx.get(), i) == LLAMA_TOKEN_NULL);
             GGML_ASSERT(llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), i) == (uint32_t) k);
             GGML_ASSERT(llama_get_sampled_candidates_count_ith(test_ctx.ctx.get(), i) == (uint32_t) k);
             const llama_token token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), i);
             GGML_ASSERT(token >= 0 && token < llama_vocab_n_tokens(vocab));
         }
-        GGML_ASSERT(sampler_ctx->apply_count == batch.n_tokens);
-
-        llama_batch_free(batch);
+        GGML_ASSERT(sampler_ctx->apply_count == batch.size());
     }

     printf("backend multi-output CPU suffix test PASSED\n");
diff --git a/tests/test-fusion.cpp b/tests/test-fusion.cpp
index 65d444ecf..55e4f9df4 100644
--- a/tests/test-fusion.cpp
+++ b/tests/test-fusion.cpp
@@ -145,13 +145,11 @@ static llama_context_ptr create_ctx(llama_model * model, int n_ubatch) {
 // decode all tokens in one batch; returns the logits of every token
 static std::vector<float> decode_prefill(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));
-    llama_batch batch = llama_batch_init(tokens.size(), 0, 1);
+    common_batch batch(lctx);
     for (size_t i = 0; i < tokens.size(); i++) {
-        common_batch_add(batch, tokens[i], i, { 0 }, true);
+        batch.add(tokens[i], i, 0, true);
     }
-    batch.n_tokens = tokens.size();
-    if (llama_decode(lctx, batch)) {
-        llama_batch_free(batch);
+    if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
         throw std::runtime_error("prefill decode failed");
     }

@@ -163,20 +161,18 @@ static std::vector<float> decode_prefill(llama_model * model, llama_context * lc
             ret.push_back(logits_ith[j]);
         }
     }
-    llama_batch_free(batch);
     return ret;
 }

 // decode one token at a time; returns the logits of the last token of each step
 static std::vector<float> decode_gen(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));
-    llama_batch batch = llama_batch_init(1, 0, 1);
+    common_batch batch(lctx);
     std::vector<float> ret;
     for (size_t i = 0; i < tokens.size(); i++) {
-        common_batch_clear(batch);
-        common_batch_add(batch, tokens[i], i, { 0 }, true);
-        if (llama_decode(lctx, batch)) {
-            llama_batch_free(batch);
+        batch.clear();
+        batch.add(tokens[i], i, 0, true);
+        if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
             throw std::runtime_error("decode failed");
         }
         const float * logits = llama_get_logits_ith(lctx, 0);
@@ -184,7 +180,6 @@ static std::vector<float> decode_gen(llama_model * model, llama_context * lctx,
             ret.push_back(logits[j]);
         }
     }
-    llama_batch_free(batch);
     return ret;
 }

diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp
index 298073d06..015e3414b 100644
--- a/tests/test-llama-archs.cpp
+++ b/tests/test-llama-archs.cpp
@@ -510,20 +510,17 @@ static std::vector<float> get_logits(
     const uint32_t n_vocab  = llama_vocab_n_tokens(llama_model_get_vocab(model));
     const uint32_t n_ctx    = llama_n_ctx(lctx);
     const uint32_t n_tokens = tokens.size();
-    llama_batch batch = llama_batch_init(n_ctx, 0, 1);
+    common_batch batch(lctx);
     GGML_ASSERT(n_tokens <= n_ctx);
     for (uint32_t pos = 0; pos < n_tokens; pos++) {
-        common_batch_add(batch, tokens[pos], pos, {0}, true);
+        batch.add(tokens[pos], pos, 0, true);
     }
-    batch.n_tokens = n_tokens;
     if (encode) {
-        if (llama_encode(lctx, batch)) {
-            llama_batch_free(batch);
+        if (llama_process(lctx, LLAMA_PROCESS_TYPE_ENCODE, batch.get())) {
             throw std::runtime_error("failed to encode batch");
         }
     }
-    if (llama_decode(lctx, batch)) {
-        llama_batch_free(batch);
+    if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
         throw std::runtime_error("failed to decode batch");
     }

@@ -535,7 +532,6 @@ static std::vector<float> get_logits(
             ret.push_back(logits_ith[j]);
         }
     }
-    llama_batch_free(batch);
     return ret;
 }

diff --git a/tests/test-recurrent-state-rollback.cpp b/tests/test-recurrent-state-rollback.cpp
index 4ad0d6f9e..f8eda55c8 100644
--- a/tests/test-recurrent-state-rollback.cpp
+++ b/tests/test-recurrent-state-rollback.cpp
@@ -35,21 +35,17 @@ static const char * test_status_str(test_status status) {
 }

 static bool decode_tokens(llama_context * ctx, const std::vector<llama_token> & tokens, uint32_t count) {
-    llama_batch batch = llama_batch_init(count, 0, 1);
+    common_batch batch(ctx);
     for (uint32_t pos = 0; pos < count; ++pos) {
-        common_batch_add(batch, tokens[pos], pos, { 0 }, pos + 1 == count);
+        batch.add(tokens[pos], pos, 0, pos + 1 == count);
     }
-    const bool ok = llama_decode(ctx, batch) == 0;
-    llama_batch_free(batch);
-    return ok;
+    return llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0;
 }

-static bool decode_one(llama_context * ctx, llama_token tok, llama_pos pos) {
-    llama_batch batch = llama_batch_init(1, 0, 1);
-    common_batch_add(batch, tok, pos, { 0 }, true);
-    const bool ok = llama_decode(ctx, batch) == 0;
-    llama_batch_free(batch);
-    return ok;
+static bool decode_one(llama_context * ctx, llama_token tok, llama_pos pos, llama_seq_id seq = 0) {
+    common_batch batch(ctx);
+    batch.add(tok, pos, seq, true);
+    return llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0;
 }

 struct cache_buffer_collector : llama_io_write_i {
@@ -166,22 +162,22 @@ static test_status test_multi_seq_split_replay(const common_params & params, lla

     bool ok = true;

+    // decode tokens [p_begin, p_end) of seq s, a batch belongs to one context so it is built per call
+    const auto decode_range = [&](llama_context * ctx, uint32_t s, llama_pos p_begin, llama_pos p_end) {
+        common_batch batch(ctx);
+        for (llama_pos pos = p_begin; pos < p_end; ++pos) {
+            batch.add(tok(s, pos), pos, (llama_seq_id) s, false);
+        }
+        return llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0;
+    };
+
     // both contexts decode the identical [0, p0) prefill; only ctx_roll decodes
     // the tail, which is then rolled back so its restore is pending at replay
     for (uint32_t s = 0; s < n_seqs && ok; ++s) {
-        llama_batch batch = llama_batch_init(n_prompt, 0, 1);
-        for (llama_pos pos = 0; pos < (llama_pos) p0; ++pos) {
-            common_batch_add(batch, tok(s, pos), pos, { (llama_seq_id) s }, false);
-        }
-        ok = ok && llama_decode(ctx_roll.get(), batch) == 0;
-        ok = ok && llama_decode(ctx_ref.get(),  batch) == 0;
+        ok = ok && decode_range(ctx_roll.get(), s, 0, (llama_pos) p0);
+        ok = ok && decode_range(ctx_ref.get(),  s, 0, (llama_pos) p0);

-        common_batch_clear(batch);
-        for (llama_pos pos = p0; pos < (llama_pos) n_prompt; ++pos) {
-            common_batch_add(batch, tok(s, pos), pos, { (llama_seq_id) s }, false);
-        }
-        ok = ok && llama_decode(ctx_roll.get(), batch) == 0;
-        llama_batch_free(batch);
+        ok = ok && decode_range(ctx_roll.get(), s, (llama_pos) p0, (llama_pos) n_prompt);

         ok = ok && llama_memory_seq_rm(llama_get_memory(ctx_roll.get()), (llama_seq_id) s, p0, -1);

@@ -193,16 +189,19 @@ static test_status test_multi_seq_split_replay(const common_params & params, lla
         return test_status::FAIL;
     }

-    llama_batch batch = llama_batch_init(n_seqs*n_replay, 0, 1);
-    for (uint32_t s = 0; s < n_seqs; ++s) {
-        for (uint32_t i = 0; i < n_replay; ++i) {
-            const llama_pos pos = p0 + (llama_pos) i;
-            common_batch_add(batch, tok(s, pos), pos, { (llama_seq_id) s }, true);
+    // all seqs replay in a single batch
+    const auto decode_replay = [&](llama_context * ctx) {
+        common_batch batch(ctx);
+        for (uint32_t s = 0; s < n_seqs; ++s) {
+            for (uint32_t i = 0; i < n_replay; ++i) {
+                const llama_pos pos = p0 + (llama_pos) i;
+                batch.add(tok(s, pos), pos, (llama_seq_id) s, true);
+            }
         }
-    }
-    ok = llama_decode(ctx_roll.get(), batch) == 0;
-    ok = ok && llama_decode(ctx_ref.get(), batch) == 0;
-    llama_batch_free(batch);
+        return llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0;
+    };
+    ok = decode_replay(ctx_roll.get());
+    ok = ok && decode_replay(ctx_ref.get());
     if (!ok) {
         LOG_ERR("%s: multi-seq replay decode failed\n", __func__);
         return test_status::FAIL;
@@ -258,13 +257,12 @@ static test_status test_multi_seq_split_replay(const common_params & params, lla
     constexpr uint32_t n_tail = 4;

     {
-        llama_batch batch_tail = llama_batch_init(n_tail, 0, 1);
+        common_batch batch_tail(ctx_ref.get());
         for (uint32_t i = 0; i < n_tail; ++i) {
             const llama_pos pos = p0 + (llama_pos) (n_replay + i);
-            common_batch_add(batch_tail, tok(0, pos + 7), pos, { 0 }, false);
+            batch_tail.add(tok(0, pos + 7), pos, 0, false);
         }
-        ok = llama_decode(ctx_ref.get(), batch_tail) == 0;
-        llama_batch_free(batch_tail);
+        ok = llama_process(ctx_ref.get(), LLAMA_PROCESS_TYPE_DECODE, batch_tail.get()) == 0;
     }

     float diff_tail = 0.0f;
@@ -272,11 +270,8 @@ static test_status test_multi_seq_split_replay(const common_params & params, lla
     double nmse_tail_a0 = 0.0;
     for (uint32_t i = 0; i < n_tail && ok; ++i) {
         const llama_pos pos = p0 + (llama_pos) (n_replay + i);
-        llama_batch batch_one = llama_batch_init(1, 0, 1);
-        common_batch_add(batch_one, tok(1, pos), pos, { 1 }, true);
-        ok = llama_decode(ctx_roll.get(), batch_one) == 0;
-        ok = ok && llama_decode(ctx_ref.get(), batch_one) == 0;
-        llama_batch_free(batch_one);
+        ok = decode_one(ctx_roll.get(), tok(1, pos), pos, 1);
+        ok = ok && decode_one(ctx_ref.get(), tok(1, pos), pos, 1);
         if (!ok) {
             break;
         }
diff --git a/tests/test-save-load-state.cpp b/tests/test-save-load-state.cpp
index dee5e17be..02984d185 100644
--- a/tests/test-save-load-state.cpp
+++ b/tests/test-save-load-state.cpp
@@ -69,26 +69,9 @@ static bool get_current_logits(llama_context * ctx, std::vector<float> & out) {
     return true;
 }

-struct llama_batch_ptr {
-    llama_batch batch;
-
-    llama_batch_ptr(int32_t n_tokens, int32_t embd, int32_t n_seq_max)
-        : batch{llama_batch_init(n_tokens, embd, n_seq_max)} {}
-
-    ~llama_batch_ptr() { llama_batch_free(batch); }
-
-    llama_batch_ptr(const llama_batch_ptr &) = delete;
-    llama_batch_ptr & operator=(const llama_batch_ptr &) = delete;
-    llama_batch_ptr(llama_batch_ptr &&) = default;
-    llama_batch_ptr & operator=(llama_batch_ptr &&) = default;
-
-    llama_batch & get() { return batch; }
-    const llama_batch & get() const { return batch; }
-};
-
 static generation_result generate_tokens(llama_context * ctx, llama_sampler * smpl, int & n_past, int32_t n_predict, llama_seq_id seq_id) {
     generation_result result;
-    llama_batch_ptr batch(1, 0, 1);
+    common_batch batch(ctx);

     for (int i = 0; i < n_predict; i++) {
         std::vector<float> logits;
@@ -104,10 +87,10 @@ static generation_result generate_tokens(llama_context * ctx, llama_sampler * sm
         result.tokens.push_back(next_token);
         result.logits.push_back(std::move(logits));

-        common_batch_clear(batch.get());
-        common_batch_add(batch.get(), next_token, n_past, {seq_id}, true);
+        batch.clear();
+        batch.add(next_token, n_past, seq_id, true);

-        if (llama_decode(ctx, batch.get())) {
+        if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
             LOG_ERR("\n%s: failed to evaluate\n", __func__);
             return {};
         }
@@ -125,7 +108,7 @@ static bool generate_tokens_compare(
         return false;
     }

-    llama_batch_ptr batch(1, 0, 1);
+    common_batch batch(ctx);

     for (int i = 0; i < n_predict; i++) {
         std::vector<float> logits;
@@ -153,10 +136,10 @@ static bool generate_tokens_compare(
             LOG_TRC("%s: sampled token %d differs from expected %d, using expected token\n", __func__, next_token, expected_token);
         }

-        common_batch_clear(batch.get());
-        common_batch_add(batch.get(), expected_token, n_past, {seq_id}, true);
+        batch.clear();
+        batch.add(expected_token, n_past, seq_id, true);

-        if (llama_decode(ctx, batch.get())) {
+        if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
             LOG_ERR("\n%s: failed to evaluate\n", __func__);
             return false;
         }
@@ -222,12 +205,12 @@ static bool test_seq_rm_isolated(

     const size_t n_tokens = tokens.size() < 128 ? tokens.size() : 128;
     for (llama_seq_id seq_id = 0; seq_id < 2; ++seq_id) {
-        llama_batch_ptr batch(n_tokens, 0, 1);
+        common_batch batch(ctx.get());
         for (size_t i = 0; i < n_tokens; ++i) {
-            common_batch_add(batch.get(), tokens[i], i, { seq_id }, i == n_tokens - 1);
+            batch.add(tokens[i], i, seq_id, i == n_tokens - 1);
         }

-        if (llama_decode(ctx.get(), batch.get())) {
+        if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
             LOG_ERR("%s: failed to decode prompt for sequence %d\n", __func__, seq_id);
             return false;
         }
@@ -469,9 +452,9 @@ static bool test_seq_cp_scatter(struct llama_model * model, const struct common_
     const uint32_t flags = on_device ? LLAMA_STATE_SEQ_FLAGS_ON_DEVICE : LLAMA_STATE_SEQ_FLAGS_NONE;

     auto decode_one = [&](llama_token tok, int pos, llama_seq_id seq) {
-        llama_batch_ptr batch(1, 0, 1);
-        common_batch_add(batch.get(), tok, pos, { seq }, true);
-        return llama_decode(ctx.get(), batch.get()) == 0;
+        common_batch batch(ctx.get());
+        batch.add(tok, pos, seq, true);
+        return llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0;
     };

     // seq 0 cells 0,1,4 interleave the seq 1 cells 2,3,5
@@ -554,7 +537,8 @@ static bool test_state_roundtrip(struct llama_model * model, const struct common

     LOGV(LOG_LEVEL_INFO, "\n=== Test 8: state blob round-trip ===\n");

-    if (llama_decode(ctx.get(), llama_batch_get_one(const_cast<llama_token *>(tokens.data()), (int32_t) tokens.size()))) {
+    common_batch batch = common_batch_get_one(ctx.get(), tokens);
+    if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
         LOG_ERR("\n%s: failed to decode prompt\n", __func__);
         return false;
     }
@@ -643,12 +627,12 @@ static bool test_state_restore_failure(struct llama_model * model, const struct
     }

     const auto decode = [&](const llama_tokens & inp, llama_seq_id seq_id, std::vector<float> * logits_out) {
-        llama_batch_ptr batch(inp.size(), 0, 1);
+        common_batch batch(ctx.get());
         for (size_t i = 0; i < inp.size(); ++i) {
-            common_batch_add(batch.get(), inp[i], i, { seq_id }, i == inp.size() - 1);
+            batch.add(inp[i], i, seq_id, i == inp.size() - 1);
         }

-        if (llama_decode(ctx.get(), batch.get())) {
+        if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
             LOG_ERR("%s: failed to decode on sequence %d\n", __func__, seq_id);
             return false;
         }
diff --git a/tests/test-state-restore-fragmented.cpp b/tests/test-state-restore-fragmented.cpp
index 33ce6f276..ea3006949 100644
--- a/tests/test-state-restore-fragmented.cpp
+++ b/tests/test-state-restore-fragmented.cpp
@@ -49,15 +49,15 @@ int main(int argc, char ** argv) {

     // interleave the 3 sequences:
     // 01201230123...
-    llama_batch batch = llama_batch_init(params.n_parallel*tokens.size(), 0, 1);
+    common_batch batch(ctx);
     for (size_t i = 0; i < tokens.size(); i++) {
         for (int s = 0; s < params.n_parallel; ++s) {
-            common_batch_add(batch, tokens[i], i, {s}, false);
+            batch.add(tokens[i], i, s, false);
         }
     }
-    batch.logits[batch.n_tokens - 1] = true;
+    batch.set_output(batch.size() - 1, true);

-    if (llama_decode(ctx, batch)) {
+    if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
         fprintf(stderr, "%s : failed to decode seq 0\n", __func__);
         return 1;
     }
@@ -91,7 +91,6 @@ int main(int argc, char ** argv) {
         fprintf(stderr, "%s : FAILED to restore seq state into fragmented cache (got %zu, expected %zu)\n",
                 __func__, nset, seq_state.size());
         fprintf(stderr, "%s : This is the bug - state restore fails with fragmented KV cache\n", __func__);
-        llama_batch_free(batch);
         return 1;
     }
     fprintf(stderr, "%s : restored state into seq 1, %zu bytes\n", __func__, nset);
@@ -105,13 +104,12 @@ int main(int argc, char ** argv) {
     auto next_token = llama_sampler_sample(smpl, ctx, -1);
     auto next_token_str = common_token_to_piece(ctx, next_token);

-    common_batch_clear(batch);
-    common_batch_add(batch, next_token, (int)tokens.size(), {1}, true);
+    batch.clear();
+    batch.add(next_token, (int)tokens.size(), 1, true);

-    if (llama_decode(ctx, batch)) {
+    if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
         fprintf(stderr, "%s : failed to decode with restored state\n", __func__);
         llama_sampler_free(smpl);
-        llama_batch_free(batch);
         return 1;
     }

@@ -119,7 +117,6 @@ int main(int argc, char ** argv) {
     fprintf(stderr, "%s : SUCCESS - state restore works with fragmented KV cache\n", __func__);

     llama_sampler_free(smpl);
-    llama_batch_free(batch);

     return 0;
 }
diff --git a/tests/test-thread-safety.cpp b/tests/test-thread-safety.cpp
index d0b5946e2..4fbb1a270 100644
--- a/tests/test-thread-safety.cpp
+++ b/tests/test-thread-safety.cpp
@@ -97,7 +97,6 @@ int main(int argc, char ** argv) {
                     return;
                 }

-                llama_batch batch = {};
                 {
                     auto prompt = common_tokenize(ctx.get(), params.prompt, true);
                     if (prompt.empty()) {
@@ -105,8 +104,8 @@ int main(int argc, char ** argv) {
                         failed.store(true);
                         return;
                     }
-                    batch = llama_batch_get_one(prompt.data(), prompt.size());
-                    if (llama_decode(ctx.get(), batch)) {
+                    common_batch batch = common_batch_get_one(ctx.get(), prompt);
+                    if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
                         LOG_ERR("failed to decode prompt\n");
                         failed.store(true);
                         return;
@@ -117,12 +116,7 @@ int main(int argc, char ** argv) {
                 std::string result = params.prompt;

                 for (int i = 0; i < params.n_predict; i++) {
-                    llama_token token;
-                    if (batch.n_tokens > 0) {
-                        token = common_sampler_sample(sampler.get(), ctx.get(), batch.n_tokens - 1);
-                    } else {
-                        token = llama_vocab_bos(vocab);
-                    }
+                    llama_token token = common_sampler_sample(sampler.get(), ctx.get(), -1);

                     result += common_token_to_piece(ctx.get(), token);

@@ -130,9 +124,9 @@ int main(int argc, char ** argv) {
                         break;
                     }

-                    batch = llama_batch_get_one(&token, 1);
+                    common_batch batch = common_batch_get_one(ctx.get(), &token, 1);

-                    int ret = llama_decode(ctx.get(), batch);
+                    int ret = llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get());
                     if (ret == 1 && i > 0) {
                         LOG_INF("Context full, stopping generation.\n");
                         break;
diff --git a/tools/batched-bench/batched-bench.cpp b/tools/batched-bench/batched-bench.cpp
index e2dcd0b2e..c260425d8 100644
--- a/tools/batched-bench/batched-bench.cpp
+++ b/tools/batched-bench/batched-bench.cpp
@@ -76,24 +76,14 @@ int llama_batched_bench(int argc, char ** argv) {

     const int32_t n_kv_max = llama_n_ctx(ctx);

-    llama_batch batch = llama_batch_init(n_kv_max, 0, 1);
+    common_batch batch(ctx);

     // decode in batches of ctx_params.n_batch tokens
-    auto decode_helper = [](llama_context * ctx, llama_batch & batch, int32_t n_batch, bool synchronize) {
-        for (int32_t i = 0; i < batch.n_tokens; i += n_batch) {
-            const int32_t n_tokens = std::min(n_batch, batch.n_tokens - i);
-
-            llama_batch batch_view = {
-                n_tokens,
-                batch.token    + i,
-                nullptr,
-                batch.pos      + i,
-                batch.n_seq_id + i,
-                batch.seq_id   + i,
-                batch.logits   + i,
-            };
-
-            const int ret = llama_decode(ctx, batch_view);
+    auto decode_helper = [](llama_context * ctx, common_batch & batch, int32_t n_batch, bool synchronize) {
+        for (int32_t i = 0; i < batch.size(); i += n_batch) {
+            const int32_t n_tokens = std::min(n_batch, batch.size() - i);
+
+            const int ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get_sub_batch(i, n_tokens));
             if (ret != 0) {
                 LOG_ERR("failed to decode the batch, n_batch = %d, ret = %d\n", n_batch, ret);
                 return false;
@@ -110,7 +100,7 @@ int llama_batched_bench(int argc, char ** argv) {
     // warm up
     {
         for (int i = 0; i < 16; ++i) {
-            common_batch_add(batch, get_token_rand(), i, { 0 }, false);
+            batch.add(get_token_rand(), i, 0, false);
         }

         if (!decode_helper(ctx, batch, ctx_params.n_batch, true)) {
@@ -142,11 +132,11 @@ int llama_batched_bench(int argc, char ** argv) {
                     continue;
                 }

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

                 for (int j = 0; j < (is_pp_shared ? 1 : pl); ++j) {
                     for (int i = 0; i < pp; ++i) {
-                        common_batch_add(batch, get_token_rand(), i, { j }, i == pp - 1);
+                        batch.add(get_token_rand(), i, j, i == pp - 1);
                     }
                 }

@@ -172,8 +162,8 @@ int llama_batched_bench(int argc, char ** argv) {

                     if (!params.kv_unified) {
                         // run one dummy token to apply the memory copy
-                        common_batch_clear(batch);
-                        common_batch_add(batch, get_token_rand(), pp + 0, { 0 }, true);
+                        batch.clear();
+                        batch.add(get_token_rand(), pp + 0, 0, true);
                         if (!decode_helper(ctx, batch, ctx_params.n_batch, true)) {
                             LOG_ERR("%s: llama_decode() failed\n", __func__);
                             llama_free(ctx);
@@ -191,9 +181,9 @@ int llama_batched_bench(int argc, char ** argv) {
                     // 0 0 0 ... 1 1 1 ... 2 2 2 ... 3 3 3 ...
                     for (int j = 0; j < pl; ++j) {
                         for (int i = 0; i < tg; ++i) {
-                            common_batch_clear(batch);
+                            batch.clear();

-                            common_batch_add(batch, get_token_rand(), pp + i, { j }, true);
+                            batch.add(get_token_rand(), pp + i, j, true);

                             if (!decode_helper(ctx, batch, ctx_params.n_batch, true)) {
                                 LOG_ERR("%s: llama_decode() failed\n", __func__);
@@ -207,10 +197,10 @@ int llama_batched_bench(int argc, char ** argv) {
                     // decode pattern:
                     // 0123 0123 0123 ...
                     for (int i = 0; i < tg; ++i) {
-                        common_batch_clear(batch);
+                        batch.clear();

                         for (int j = 0; j < pl; ++j) {
-                            common_batch_add(batch, get_token_rand(), pp + i, { j }, true);
+                            batch.add(get_token_rand(), pp + i, j, true);
                         }

                         if (!decode_helper(ctx, batch, ctx_params.n_batch, true)) {
@@ -251,7 +241,6 @@ int llama_batched_bench(int argc, char ** argv) {
     LOG("\n");
     llama_perf_context_print(ctx);

-    llama_batch_free(batch);

     llama_free(ctx);
     llama_model_free(model);
diff --git a/tools/completion/completion.cpp b/tools/completion/completion.cpp
index 941b7399b..718438ef8 100644
--- a/tools/completion/completion.cpp
+++ b/tools/completion/completion.cpp
@@ -525,10 +525,9 @@ int llama_completion(int argc, char ** argv) {
     }

     if (llama_model_has_encoder(model)) {
-        int enc_input_size = embd_inp.size();
-        llama_token * enc_input_buf = embd_inp.data();
+        common_batch batch = common_batch_get_one(ctx, embd_inp);

-        if (llama_encode(ctx, llama_batch_get_one(enc_input_buf, enc_input_size))) {
+        if (llama_process(ctx, LLAMA_PROCESS_TYPE_ENCODE, batch.get())) {
             LOG_ERR("%s : failed to eval\n", __func__);
             return 1;
         }
diff --git a/tools/cvector-generator/cvector-generator.cpp b/tools/cvector-generator/cvector-generator.cpp
index 558c37e61..af05031f2 100644
--- a/tools/cvector-generator/cvector-generator.cpp
+++ b/tools/cvector-generator/cvector-generator.cpp
@@ -346,7 +346,8 @@ static bool cb_eval(struct ggml_tensor * t, bool ask, void * user_data) {

 static bool get_hidden_layers(llama_context * ctx, std::vector<llama_token> & tokens) {
     llama_memory_clear(llama_get_memory(ctx), true);
-    if (llama_decode(ctx, llama_batch_get_one(tokens.data(), tokens.size()))) {
+    common_batch batch = common_batch_get_one(ctx, tokens);
+    if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
         fprintf(stderr, "%s : failed to eval\n", __func__);
         return false;
     }
diff --git a/tools/imatrix/imatrix.cpp b/tools/imatrix/imatrix.cpp
index f5fee6218..7baa46968 100644
--- a/tools/imatrix/imatrix.cpp
+++ b/tools/imatrix/imatrix.cpp
@@ -845,7 +845,7 @@ static bool compute_imatrix(llama_context * ctx, const common_params & params, c
     GGML_ASSERT(n_batch < n_ctx || n_batch % n_ctx == 0);
     GGML_ASSERT(params.n_ctx == n_seq * n_ctx);

-    llama_batch batch = llama_batch_init(std::min(n_batch, n_ctx*n_seq), 0, 1);
+    common_batch batch(ctx);

     std::vector<float> logits;
     if (params.compute_ppl && num_batches > 1) {
@@ -872,7 +872,7 @@ static bool compute_imatrix(llama_context * ctx, const common_params & params, c
             const int batch_size  = std::min(end - batch_start, n_batch);

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

             for (int seq = 0; seq < n_seq_batch; seq++) {
                 int seq_start = batch_start + seq*n_ctx;
@@ -889,16 +889,15 @@ static bool compute_imatrix(llama_context * ctx, const common_params & params, c
                     //       and also for the perplexity calculation.
                     // TODO: only get outputs when (params.process_output || params.compute_ppl)
                     //       (not possible when this skips FFN computation of the last layer)
-                    common_batch_add(batch, tokens[seq_start + k], j*n_batch + k, { seq }, true);
+                    batch.add(tokens[seq_start + k], j*n_batch + k, seq, true);
                 }

                 // restore the original token in case it was set to BOS
                 tokens[seq_start] = token_org;
             }

-            if (llama_decode(ctx, batch)) {
+            if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
                 LOG_ERR("%s : failed to eval\n", __func__);
-                llama_batch_free(batch);
                 return false;
             }

@@ -960,7 +959,6 @@ static bool compute_imatrix(llama_context * ctx, const common_params & params, c
         }
     }

-    llama_batch_free(batch);

     return true;
 }
diff --git a/tools/llama-bench/llama-bench.cpp b/tools/llama-bench/llama-bench.cpp
index 70e15d044..fd64b4c44 100644
--- a/tools/llama-bench/llama-bench.cpp
+++ b/tools/llama-bench/llama-bench.cpp
@@ -2182,7 +2182,8 @@ static bool test_prompt(llama_context * ctx, int n_prompt, int n_batch, int n_th
         for (int i = 1; i < n_tokens; i++) {
             tokens[i] = std::rand() % n_vocab;
         }
-        int res = llama_decode(ctx, llama_batch_get_one(tokens.data(), n_tokens));
+        common_batch batch = common_batch_get_one(ctx, tokens.data(), n_tokens);
+        int res = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
         if (res != 0) {
             fprintf(stderr, "%s: failed to decode prompt batch, res = %d\n", __func__, res);
             return false;
@@ -2203,8 +2204,13 @@ static bool test_gen(llama_context * ctx, int n_gen, int n_threads) {

     llama_token token = llama_vocab_get_add_bos(vocab) ? llama_vocab_bos(vocab) : std::rand() % n_vocab;

+    common_batch batch(ctx);
+    llama_pos pos = llama_memory_seq_pos_max(llama_get_memory(ctx), 0) + 1;
+
     for (int i = 0; i < n_gen; i++) {
-        int res = llama_decode(ctx, llama_batch_get_one(&token, 1));
+        batch.clear();
+        batch.add(token, pos++, 0, true);
+        int res = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
         if (res != 0) {
             fprintf(stderr, "%s: failed to decode generation batch, res = %d\n", __func__, res);
             return false;
diff --git a/tools/perplexity/perplexity.cpp b/tools/perplexity/perplexity.cpp
index ba41287d8..601361acd 100644
--- a/tools/perplexity/perplexity.cpp
+++ b/tools/perplexity/perplexity.cpp
@@ -366,21 +366,20 @@ static results_perplexity perplexity_v2(llama_context * ctx, const common_params
         // clear the KV cache
         llama_memory_clear(llama_get_memory(ctx), true);

-        llama_batch batch = llama_batch_init(n_batch, 0, 1);
+        common_batch batch(ctx);

         for (int j = 0; j < num_batches; ++j) {
             const int batch_start = start + j * n_batch;
             const int batch_size  = std::min(end - batch_start, n_batch);

-            common_batch_clear(batch);
+            batch.clear();
             for (int i = 0; i < batch_size; i++) {
-                common_batch_add(batch, tokens[batch_start + i], j*n_batch + i, {0}, true);
+                batch.add(tokens[batch_start + i], j*n_batch + i, 0, true);
             }

             //LOG_DBG("    Batch %d: starts at %d, size is %d, n_past is %d\n",j,batch_start,batch_size,j * n_batch);
-            if (llama_decode(ctx, batch)) {
+            if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
                 //LOG_ERR("%s : failed to eval\n", __func__);
-                llama_batch_free(batch);
                 return {tokens, -1, logit_history, prob_history};
             }

@@ -400,7 +399,6 @@ static results_perplexity perplexity_v2(llama_context * ctx, const common_params
             }
         }

-        llama_batch_free(batch);

         const auto t_end = std::chrono::high_resolution_clock::now();

@@ -507,7 +505,7 @@ static results_perplexity perplexity(llama_context * ctx, const common_params &
     GGML_ASSERT(n_batch < n_ctx || n_batch % n_ctx == 0);
     GGML_ASSERT(params.n_ctx == n_seq * n_ctx);

-    llama_batch batch = llama_batch_init(std::min(n_batch, n_ctx*n_seq), 0, 1);
+    common_batch batch(ctx);

     std::vector<float> logits;
     if (num_batches > 1) {
@@ -558,7 +556,7 @@ static results_perplexity perplexity(llama_context * ctx, const common_params &

             int n_outputs = 0;

-            batch.n_tokens = 0;
+            batch.clear();
             for (int seq = 0; seq < n_seq_batch; seq++) {
                 int seq_start = batch_start + seq*n_ctx;

@@ -571,22 +569,17 @@ static results_perplexity perplexity(llama_context * ctx, const common_params &
                 }

                 for (int k = 0; k < batch_size; ++k) {
-                    const int idx = seq*n_ctx + k;
-                    batch.token   [idx]    = tokens[seq_start + k];
-                    batch.pos     [idx]    = j*n_batch + k;
-                    batch.n_seq_id[idx]    = 1;
-                    batch.seq_id  [idx][0] = seq;
-                    batch.logits  [idx]    = batch.pos[idx] >= first ? 1 : 0;
-
-                    n_outputs += batch.logits[idx] != 0;
+                    const llama_pos pos = j*n_batch + k;
+                    const bool need_logits = pos >= first;
+                    batch.add(tokens[seq_start + k], pos, seq, need_logits);
+                    n_outputs += need_logits;
                 }
-                batch.n_tokens += batch_size;

                 // restore the original token in case it was set to BOS
                 tokens[seq_start] = token_org;
             }

-            if (llama_decode(ctx, batch)) {
+            if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
                 LOG_INF("%s : failed to decode\n", __func__);
                 return {tokens, -1, logit_history, prob_history};
             }
@@ -656,35 +649,24 @@ static results_perplexity perplexity(llama_context * ctx, const common_params &
         LOG_ERR("Unexpected negative standard deviation of log(prob)\n");
     }

-    llama_batch_free(batch);

     return {tokens, ppl, logit_history, prob_history};
 }

-static bool decode_helper(llama_context * ctx, llama_batch & batch, std::vector<float> & batch_logits, int n_batch, int n_vocab) {
+static bool decode_helper(llama_context * ctx, common_batch & batch, std::vector<float> & batch_logits, int n_batch, int n_vocab) {
     int prev_outputs = 0;
-    for (int i = 0; i < (int) batch.n_tokens; i += n_batch) {
-        const int n_tokens = std::min<int>(n_batch, batch.n_tokens - i);
-
-        llama_batch batch_view = {
-            n_tokens,
-            batch.token    + i,
-            nullptr,
-            batch.pos      + i,
-            batch.n_seq_id + i,
-            batch.seq_id   + i,
-            batch.logits   + i,
-        };
+    for (int i = 0; i < batch.size(); i += n_batch) {
+        const int n_tokens = std::min<int>(n_batch, batch.size() - i);

-        const int ret = llama_decode(ctx, batch_view);
+        const int ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get_sub_batch(i, n_tokens));
         if (ret != 0) {
             LOG_ERR("failed to decode the batch, n_batch = %d, ret = %d\n", n_batch, ret);
             return false;
         }

         int n_outputs = 0;
-        for (int i = 0; i < n_tokens; ++i) {
-            n_outputs += batch_view.logits[i] != 0;
+        for (int j = i; j < i + n_tokens; ++j) {
+            n_outputs += batch.tokens[j].output;
         }

         memcpy(batch_logits.data() + size_t(prev_outputs)*n_vocab, llama_get_logits(ctx), size_t(n_outputs)*n_vocab*sizeof(float));
@@ -866,7 +848,7 @@ static void hellaswag_score(llama_context * ctx, const common_params & params) {
     const int max_tasks_per_batch = 32;
     const int max_seq = std::min(4*max_tasks_per_batch, (int) llama_n_seq_max(ctx));

-    llama_batch batch = llama_batch_init(n_ctx, 0, 4);
+    common_batch batch(ctx);

     std::vector<float> tok_logits(n_vocab);
     // TODO: this could be made smaller; it's currently the worst-case size
@@ -882,7 +864,7 @@ static void hellaswag_score(llama_context * ctx, const common_params & params) {
         size_t i1 = i0;
         size_t i_logits = 0; // this tells us how many logits were needed before this point in the batch

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

         // batch as much tasks as possible into the available context
         // each task has 4 unique sequence ids - one for each ending
@@ -898,9 +880,9 @@ static void hellaswag_score(llama_context * ctx, const common_params & params) {
             }

             for (size_t i = 0; i < hs_cur.common_prefix; ++i) {
-                common_batch_add(batch, hs_cur.seq_tokens[0][i], i, { s0 + 0, s0 + 1, s0 + 2, s0 + 3 }, false);
+                batch.add(hs_cur.seq_tokens[0][i], i, { s0 + 0, s0 + 1, s0 + 2, s0 + 3 }, false);
             }
-            batch.logits[batch.n_tokens - 1] = true; // we need logits for the last token of the common prefix
+            batch.set_output(batch.size() - 1, true); // we need logits for the last token of the common prefix
             n_logits += 1;

             for (int s = 0; s < 4; ++s) {
@@ -908,7 +890,7 @@ static void hellaswag_score(llama_context * ctx, const common_params & params) {
                 // TODO: don't evaluate the last token of each sequence
                 for (size_t i = hs_cur.common_prefix; i < seq_tokens_size; ++i) {
                     const bool needs_logits = i < seq_tokens_size - 1;
-                    common_batch_add(batch, hs_cur.seq_tokens[s][i], i, { s0 + s }, needs_logits);
+                    batch.add(hs_cur.seq_tokens[s][i], i, s0 + s, needs_logits);
                     n_logits += needs_logits;
                 }
             }
@@ -1009,7 +991,6 @@ static void hellaswag_score(llama_context * ctx, const common_params & params) {
         i0 = i1 - 1;
     }

-    llama_batch_free(batch);

     LOG("\n");
 }
@@ -1164,7 +1145,7 @@ static void winogrande_score(llama_context * ctx, const common_params & params)
     const int max_tasks_per_batch = 128;
     const int max_seq = std::min(2*max_tasks_per_batch, (int) llama_n_seq_max(ctx));

-    llama_batch batch = llama_batch_init(n_ctx, 0, 2);
+    common_batch batch(ctx);

     std::vector<float> tok_logits(n_vocab);
     // TODO: this could be made smaller; it's currently the worst-case size
@@ -1183,7 +1164,7 @@ static void winogrande_score(llama_context * ctx, const common_params & params)
         size_t i1 = i0;
         size_t i_logits = 0;

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

         while (n_cur + (int) data[i1].required_tokens <= n_ctx) {
             int n_logits = 0;
@@ -1193,15 +1174,15 @@ static void winogrande_score(llama_context * ctx, const common_params & params)
             }

             for (size_t i = 0; i < data[i1].common_prefix; ++i) {
-                common_batch_add(batch, data[i1].seq_tokens[0][i], i, { s0 + 0, s0 + 1 }, false);
+                batch.add(data[i1].seq_tokens[0][i], i, { s0 + 0, s0 + 1 }, false);
             }
-            batch.logits[batch.n_tokens - 1] = true;
+            batch.set_output(batch.size() - 1, true);
             n_logits += 1;

             for (int s = 0; s < 2; ++s) {
                 // TODO: end before the last token, no need to predict past the end of the sequences
                 for (size_t i = data[i1].common_prefix; i < data[i1].seq_tokens[s].size(); ++i) {
-                    common_batch_add(batch, data[i1].seq_tokens[s][i], i, { s0 + s }, true);
+                    batch.add(data[i1].seq_tokens[s][i], i, s0 + s, true);
                     n_logits += 1;
                 }
             }
@@ -1518,7 +1499,7 @@ static void multiple_choice_score(llama_context * ctx, const common_params & par
     const int max_tasks_per_batch = 32;
     const int max_seq = std::min(4*max_tasks_per_batch, (int) llama_n_seq_max(ctx));

-    llama_batch batch = llama_batch_init(n_ctx, 0, max_seq);
+    common_batch batch(ctx);

     std::vector<float> tok_logits(n_vocab);
     std::vector<float> batch_logits(size_t(n_ctx)*n_vocab);
@@ -1538,7 +1519,7 @@ static void multiple_choice_score(llama_context * ctx, const common_params & par
         size_t i1 = i0;
         size_t i_logits = 0; // this tells us how many logits were needed before this point in the batch

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

         // batch as much tasks as possible into the available context
         // each task has 4 unique sequence ids - one for each ending
@@ -1568,9 +1549,9 @@ static void multiple_choice_score(llama_context * ctx, const common_params & par

             for (size_t i = 0; i < cur_task.common_prefix; ++i) {
                 //llama_batch_add(batch, cur_task.seq_tokens[0][i], i, { s0 + 0, s0 + 1, s0 + 2, s0 + 3}, false);
-                common_batch_add(batch, cur_task.seq_tokens[0][i], i, batch_indeces, false);
+                batch.add(cur_task.seq_tokens[0][i], i, batch_indeces, false);
             }
-            batch.logits[batch.n_tokens - 1] = true; // we need logits for the last token of the common prefix
+            batch.set_output(batch.size() - 1, true); // we need logits for the last token of the common prefix
             n_logits += 1;

             for (int s = 0; s < int(cur_task.seq_tokens.size()); ++s) {
@@ -1578,7 +1559,7 @@ static void multiple_choice_score(llama_context * ctx, const common_params & par
                 // TODO: don't evaluate the last token of each sequence
                 for (size_t i = cur_task.common_prefix; i < seq_tokens_size; ++i) {
                     const bool needs_logits = i < seq_tokens_size - 1;
-                    common_batch_add(batch, cur_task.seq_tokens[s][i], i, { s0 + s }, needs_logits);
+                    batch.add(cur_task.seq_tokens[s][i], i, s0 + s, needs_logits);
                     n_logits += needs_logits;
                 }
             }
@@ -1677,7 +1658,6 @@ static void multiple_choice_score(llama_context * ctx, const common_params & par
         i0 = i1 - 1;
     }

-    llama_batch_free(batch);

     if (n_done < 100 && (params.multiple_choice_tasks != 0 && params.multiple_choice_tasks < (size_t)n_task)) return;

@@ -1753,7 +1733,7 @@ static void kl_divergence(llama_context * ctx, const common_params & params) {
     const bool add_bos = llama_vocab_get_add_bos(vocab);
     GGML_ASSERT(!llama_vocab_get_add_eos(vocab));

-    llama_batch batch = llama_batch_init(std::min(n_batch, static_cast<int>(n_ctx)*n_seq), 0, 1);
+    common_batch batch(ctx);

     std::vector<uint16_t> log_probs_uint16(size_t(n_ctx - 1 - n_ctx/2) * nv);
     std::vector<float>    kld_values(size_t(n_ctx - 1 - n_ctx/2)*n_chunk);
@@ -1808,7 +1788,7 @@ static void kl_divergence(llama_context * ctx, const common_params & params) {

             int n_outputs = 0;

-            common_batch_clear(batch);
+            batch.clear();
             for (int seq = 0; seq < n_seq_batch; seq++) {
                 int seq_start = batch_start + seq*n_ctx;

@@ -1823,7 +1803,7 @@ static void kl_divergence(llama_context * ctx, const common_params & params) {
                 for (int k = 0; k < batch_size; ++k) {
                     const int pos = j*n_batch + k;
                     const bool need_logits = pos >= first;
-                    common_batch_add(batch, tokens[seq_start + k], pos, { seq }, need_logits);
+                    batch.add(tokens[seq_start + k], pos, seq, need_logits);
                     n_outputs += need_logits;
                 }

@@ -1831,9 +1811,8 @@ static void kl_divergence(llama_context * ctx, const common_params & params) {
                 tokens[seq_start] = token_org;
             }

-            if (llama_decode(ctx, batch)) {
+            if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
                 LOG_ERR("%s : failed to decode\n", __func__);
-                llama_batch_free(batch);
                 return;
             }

@@ -1862,7 +1841,6 @@ static void kl_divergence(llama_context * ctx, const common_params & params) {
         for (int seq = 0; seq < n_seq_batch; seq++) {
             if (in.read((char *)log_probs_uint16.data(), log_probs_uint16.size()*sizeof(uint16_t)).fail()) {
                 LOG_ERR("%s: failed reading log-probs for chunk %d\n", __func__, i + seq);
-                llama_batch_free(batch);
                 return;
             }

@@ -1904,7 +1882,6 @@ static void kl_divergence(llama_context * ctx, const common_params & params) {
         logits.clear();
     }

-    llama_batch_free(batch);
     LOG("\n");

     if (kld.count < 100) return; // we do not wish to do statistics on so few values
diff --git a/tools/results/results.cpp b/tools/results/results.cpp
index f2179ed27..2d6479482 100644
--- a/tools/results/results.cpp
+++ b/tools/results/results.cpp
@@ -32,14 +32,12 @@ static std::vector<float> get_logits(
     const uint32_t n_vocab  = llama_vocab_n_tokens(llama_model_get_vocab(model));
     const uint32_t n_ctx    = llama_n_ctx(lctx);
     const uint32_t n_tokens = tokens.size();
-    llama_batch batch = llama_batch_init(n_ctx, 0, 1);
+    common_batch batch(lctx);
     GGML_ASSERT(n_tokens <= n_ctx);
     for (uint32_t pos = 0; pos < n_tokens; pos++) {
-        common_batch_add(batch, tokens[pos], pos, {0}, true);
+        batch.add(tokens[pos], pos, 0, true);
     }
-    batch.n_tokens = n_tokens;
-    if (llama_decode(lctx, batch)) {
-        llama_batch_free(batch);
+    if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
         throw std::runtime_error("failed to decode batch");
     }

@@ -51,7 +49,6 @@ static std::vector<float> get_logits(
             ret.push_back(logits_ith[j]);
         }
     }
-    llama_batch_free(batch);
     return ret;
 }