Commit f1ea20621 for llama.cpp
commit f1ea206218210afb913ae2f5d2c51faed35915da
Author: Xuan-Son Nguyen <son@huggingface.co>
Date: Mon Sep 28 19:52:45 2026 +0200
batch: migrate speculative, mtmd and server to batch_ext (#29385)
* adapt common
* add common_batch
* wip
* wip: spec
* cont
* common_speculative_process
* server_batch to use common_batch
* rm some stale calls
Assisted-by: Claude Fable 5.1
* migrate mtmd
* handle imrope, handle return val of add()/add_embd()
* add spec zeros vector
* add warning on zero fill path
diff --git a/common/common.cpp b/common/common.cpp
index 6099f2ecc..6f443f1bf 100644
--- a/common/common.cpp
+++ b/common/common.cpp
@@ -613,34 +613,6 @@ std::string string_from(const struct llama_context * ctx, const std::vector<llam
return buf.str();
}
-std::string string_from(const struct llama_context * ctx, const struct llama_batch & batch) {
- std::stringstream buf;
-
- buf << "[ ";
-
- bool first = true;
- for (int i = 0; i < batch.n_tokens; ++i) {
- if (!first) {
- buf << ", ";
- } else {
- first = false;
- }
-
- auto detokenized = common_token_to_piece(ctx, batch.token[i]);
-
- buf << "\n" << std::to_string(i)
- << ", token '" << detokenized << "'"
- << ", pos " << std::to_string(batch.pos[i])
- << ", n_seq_id " << std::to_string(batch.n_seq_id[i])
- << ", seq_id " << std::to_string(batch.seq_id[i][0])
- << ", logits " << std::to_string(batch.logits[i]);
- }
-
- buf << " ]";
-
- return buf.str();
-}
-
void string_process_escapes(std::string & input) {
std::size_t input_len = input.length();
std::size_t output_idx = 0;
@@ -1491,7 +1463,8 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
}
if (llama_model_has_encoder(model)) {
- llama_encode(lctx, llama_batch_get_one(tmp.data(), tmp.size()));
+ common_batch batch = common_batch_get_one(lctx, tmp);
+ llama_process(lctx, LLAMA_PROCESS_TYPE_ENCODE, batch.get());
llama_token decoder_start_token_id = llama_model_decoder_start_token(model);
if (decoder_start_token_id == LLAMA_TOKEN_NULL) {
decoder_start_token_id = bos;
@@ -1500,7 +1473,9 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
tmp.push_back(decoder_start_token_id);
}
if (llama_model_has_decoder(model)) {
- llama_decode(lctx, llama_batch_get_one(tmp.data(), std::min(tmp.size(), (size_t) params.n_batch)));
+ tmp.resize(std::min(tmp.size(), (size_t) params.n_batch));
+ common_batch batch = common_batch_get_one(lctx, tmp);
+ llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
}
llama_memory_clear(llama_get_memory(lctx), true);
llama_synchronize(lctx);
@@ -1564,9 +1539,13 @@ common_context_seq_rm_type common_context_can_seq_rm(llama_context * ctx) {
tmp.push_back(0);
tmp.push_back(0);
- int ret = llama_decode(ctx, llama_batch_get_one(tmp.data(), tmp.size()));
+ int ret;
+ {
+ common_batch batch = common_batch_get_one(ctx, tmp);
+ ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+ }
if (ret != 0) {
- COM_ERR("llama_decode() failed: %d\n", ret);
+ COM_ERR("llama_process() failed: %d\n", ret);
res = COMMON_CONTEXT_SEQ_RM_TYPE_NO;
goto done;
}
@@ -2153,31 +2132,140 @@ float lr_opt::get_lr(float epoch) const {
}
bool common_replay_last_token(struct llama_context * ctx, llama_token last_token, int32_t pos) {
- llama_batch batch = llama_batch_get_one(&last_token, 1);
- batch.pos = &pos;
- if (llama_decode(ctx, batch)) {
+ common_batch batch(ctx);
+ batch.add(last_token, pos, 0, true);
+
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("%s: failed to replay last token\n", __func__);
return false;
}
return true;
}
-llama_batch_ext_ptr common_batch_ext_get_one(llama_context * ctx, const llama_tokens & tokens) {
- llama_batch_ext_ptr batch(llama_batch_ext_init(ctx));
+common_batch::common_batch(llama_context * ctx) : batch(llama_batch_ext_init(ctx)) {
+ const auto rope_type = llama_model_rope_type(llama_get_model(ctx));
+ n_pos = rope_type == LLAMA_ROPE_TYPE_MROPE || rope_type == LLAMA_ROPE_TYPE_IMROPE ? GGML_MROPE_SECTIONS : 1;
+}
+
+void common_batch::clear() {
+ tokens.clear();
+ llama_batch_ext_clear(batch.get());
+}
+
+int32_t common_batch::add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output) {
+ const int32_t idx = llama_batch_ext_add_token(batch.get(), seq_id, id);
+ if (idx < 0) {
+ GGML_ABORT("%s: failed to add token %d to the batch (error %d, n_tokens = %d)\n", __func__, id, idx, size());
+ }
+ llama_batch_ext_set_pos(batch.get(), idx, &pos);
+ if (output) {
+ llama_batch_ext_set_output_logits(batch.get(), idx, true);
+ }
+ tokens.push_back({ id, { pos, 0, 0, 0 }, seq_id, output, { nullptr, 0, 0 } });
+ return idx;
+}
+
+bool common_batch::set_output(int32_t idx, bool value) {
+ if (idx < 0 || idx >= (int32_t) tokens.size()) {
+ return false;
+ }
+ tokens[idx].output = value;
+ return llama_batch_ext_set_output_logits(batch.get(), idx, value);
+}
+
+bool common_batch::set_embd(int32_t idx, llama_embd embd) {
+ if (idx < 0 || idx >= (int32_t) tokens.size()) {
+ return false;
+ }
+ if (!llama_batch_ext_set_embd_token(batch.get(), idx, embd)) {
+ return false;
+ }
+ tokens[idx].embd = embd;
+ return true;
+}
+
+int32_t common_batch::add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output) {
+ const int32_t idx = llama_batch_ext_add_embd(batch.get(), seq_id, embd);
+ if (idx < 0) {
+ GGML_ABORT("%s: failed to add embedding to the batch (error %d, n_tokens = %d)\n", __func__, idx, size());
+ }
+ llama_batch_ext_set_pos(batch.get(), idx, pos);
+ if (output) {
+ llama_batch_ext_set_output_logits(batch.get(), idx, true);
+ }
+ token t = { LLAMA_TOKEN_NULL, { 0, 0, 0, 0 }, seq_id, output, embd };
+ for (int32_t j = 0; j < n_pos; ++j) {
+ t.pos[j] = pos[j];
+ }
+ tokens.push_back(t);
+ return idx;
+}
+
+common_batch common_batch_from_llama_batch(llama_context * ctx, const llama_batch & batch) {
+ common_batch res(ctx);
+
+ const bool has_token = batch.token != nullptr;
+ const bool has_embd = batch.embd != nullptr;
+
+ const size_t n_embd = llama_model_n_embd_inp(llama_get_model(ctx));
+
+ // positions continue from the memory when none are given
+ auto * mem = llama_get_memory(ctx);
+ std::vector<llama_pos> pos_next(llama_n_seq_max(ctx));
+ for (llama_seq_id s = 0; s < (llama_seq_id) pos_next.size(); ++s) {
+ pos_next[s] = llama_memory_seq_pos_max(mem, s) + 1;
+ }
+
+ for (int32_t i = 0; i < batch.n_tokens; ++i) {
+ const int32_t n_sid = batch.n_seq_id ? batch.n_seq_id[i] : 1;
+ const llama_seq_id seq_id = batch.seq_id ? batch.seq_id[i][0] : 0;
+
+ llama_pos pos[GGML_MROPE_SECTIONS] = { 0, 0, 0, 0 };
+ if (!batch.pos) {
+ pos[0] = pos_next[seq_id]++;
+ } else if (has_token) {
+ pos[0] = batch.pos[i];
+ } else {
+ // embedding batch: section-major layout pos[j*n_tokens + i]
+ for (int32_t j = 0; j < res.n_pos; ++j) {
+ pos[j] = batch.pos[j * batch.n_tokens + i];
+ }
+ }
+
+ const bool output = batch.logits ? batch.logits[i] != 0 : i == batch.n_tokens - 1;
+
+ const llama_embd embd = { has_embd ? batch.embd + (size_t) i * n_embd : nullptr, 1, n_embd };
+
+ int32_t idx;
+ if (has_token) {
+ idx = res.add(batch.token[i], pos[0], seq_id, output);
+ if (has_embd) {
+ res.set_embd(idx, embd);
+ }
+ } else {
+ idx = res.add_embd(embd, pos, seq_id, output);
+ }
+
+ for (int32_t s = 1; s < n_sid; ++s) {
+ llama_batch_ext_add_seq(res.get(), idx, batch.seq_id[i][s]);
+ }
+ }
+
+ return res;
+}
+
+common_batch common_batch_get_one(llama_context * ctx, const llama_tokens & tokens) {
+ common_batch batch(ctx);
auto mem = llama_get_memory(ctx);
- llama_pos pos = mem ? llama_memory_seq_pos_max(mem, 0) + 1 : 0;
+ llama_pos pos = llama_memory_seq_pos_max(mem, 0) + 1; // -1 + 1 == 0 when the memory is empty
for (size_t i = 0; i < tokens.size(); ++i) {
- const int32_t idx = llama_batch_ext_add_token(batch.get(), 0, tokens[i]);
- llama_batch_ext_set_pos(batch.get(), idx, &pos);
+ const bool output = i == tokens.size() - 1;
+ batch.add(tokens[i], pos, 0, output);
pos++;
}
- if (!tokens.empty()) {
- llama_batch_ext_set_output_logits(batch.get(), (int32_t) tokens.size() - 1, true);
- }
-
return batch;
}
@@ -2205,7 +2293,7 @@ bool common_prompt_batch_decode(
// memory, so we can't just remove the last token from the memory and replay the last token which
// is the reason for this logic.
llama_tokens prefix_tokens(all_tokens.begin() + offset, all_tokens.begin() + offset + n_tokens_before_last);
- llama_batch_ext_ptr batch_prefix = common_batch_ext_get_one(ctx, prefix_tokens);
+ common_batch batch_prefix = common_batch_get_one(ctx, prefix_tokens);
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_prefix.get())) {
COM_ERR("%s", "failed to eval\n");
return false;
@@ -2215,10 +2303,8 @@ bool common_prompt_batch_decode(
llama_state_save_file(ctx, state_path.data(), all_tokens.data(), all_tokens.size());
COM_INF("saved session before last token to %s, n_new = %zu\n", state_path.data(), all_tokens.size());
- llama_token last_token = all_tokens.back();
- llama_batch_ext_ptr batch_last = common_batch_ext_get_one(ctx, { last_token });
- llama_pos pos = n_past;
- llama_batch_ext_set_pos(batch_last.get(), 0, &pos);
+ common_batch batch_last(ctx);
+ batch_last.add(all_tokens.back(), n_past, 0, true);
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_last.get())) {
COM_ERR("%s", "failed to eval last token\n");
@@ -2227,7 +2313,7 @@ bool common_prompt_batch_decode(
n_past++;
} else {
llama_tokens new_tokens(all_tokens.begin() + offset, all_tokens.begin() + offset + n_new);
- llama_batch_ext_ptr batch = common_batch_ext_get_one(ctx, new_tokens);
+ common_batch batch = common_batch_get_one(ctx, new_tokens);
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
COM_ERR("%s", "failed to eval\n");
return false;
diff --git a/common/common.h b/common/common.h
index 329469410..45b15def7 100644
--- a/common/common.h
+++ b/common/common.h
@@ -8,6 +8,7 @@
#include "ggml.h"
#include "llama.h"
+#include <array>
#include <list>
#include <set>
#include <sstream>
@@ -879,7 +880,6 @@ void string_process_escapes(std::string & input);
std::string string_from(bool value);
std::string string_from(const std::vector<int> & values);
std::string string_from(const struct llama_context * ctx, const std::vector<llama_token> & tokens);
-std::string string_from(const struct llama_context * ctx, const struct llama_batch & batch);
bool glob_match(const std::string & pattern, const std::string & str);
@@ -1039,9 +1039,54 @@ void common_batch_add(
const std::vector<llama_seq_id> & seq_ids,
bool logits);
+// wrapper around llama_batch_ext that provide getter functions for downstream code
+struct common_batch {
+ struct token {
+ llama_token id;
+ std::array<llama_pos, GGML_MROPE_SECTIONS> pos; // only pos[0] is used for text tokens
+ llama_seq_id seq_id;
+ bool output;
+ llama_embd embd; // non-owning view of the data passed to add_embd()/set_embd(), data == NULL if none
+ };
+
+ std::vector<token> tokens; // mirror of the entries, tokens[i] describes batch index i
+ llama_batch_ext_ptr batch;
+
+ int32_t n_pos = 1; // positions per embedding entry, GGML_MROPE_SECTIONS for MROPE/IMROPE
+
+ common_batch() = default;
+ common_batch(struct llama_context * ctx);
+
+ llama_batch_ext * get() const { return batch.get(); }
+
+ // content type of the batch, all entries carry the same combination
+ bool has_token() const { return !tokens.empty() && tokens[0].id != LLAMA_TOKEN_NULL; }
+ bool has_embd () const { return !tokens.empty() && tokens[0].embd.data != nullptr; }
+
+ void clear();
+
+ // returns the batch index (>= 0), aborts if the entry cannot be added (batch full, invalid token or seq id)
+ int32_t add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output);
+
+ bool set_output(int32_t idx, bool value);
+
+ // attach a token embedding to the entry at idx, can only be set once per entry
+ bool set_embd(int32_t idx, llama_embd embd);
+
+ // add an embedding-only entry (no token id), aborts like add() on failure
+ // pos points to n_pos positions
+ int32_t add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output);
+
+ int32_t size() const { return (int32_t) tokens.size(); }
+};
+
// create a single-sequence batch from a list of tokens
// last token always have output_logits set to true
-llama_batch_ext_ptr common_batch_ext_get_one(struct llama_context * ctx, const llama_tokens & tokens);
+common_batch common_batch_get_one(struct llama_context * ctx, const llama_tokens & tokens);
+
+// convert a legacy llama_batch, applying its defaults: seq 0, positions continue from memory, last token is output
+// the embd rows are read at the model input width
+common_batch common_batch_from_llama_batch(struct llama_context * ctx, const llama_batch & batch);
// decodes a single batch of tokens for a prompt and manages session tokens
//
diff --git a/common/speculative.cpp b/common/speculative.cpp
index 6fdfa4dc3..82e9e9223 100644
--- a/common/speculative.cpp
+++ b/common/speculative.cpp
@@ -165,7 +165,7 @@ struct common_speculative_impl {
virtual void begin(llama_seq_id seq_id, const llama_tokens & prompt) = 0;
- virtual bool process(const llama_batch & batch) = 0;
+ virtual bool process(const common_batch & batch) = 0;
virtual void draft(common_speculative_draft_params_vec & dparams) = 0;
@@ -179,7 +179,11 @@ struct common_speculative_impl {
struct common_speculative_impl_draft_simple : public common_speculative_impl {
common_params_speculative_draft params;
- llama_batch batch;
+ common_batch batch;
+
+ // zero row at the draft input width, stands in for target embeddings the draft cannot read
+ std::vector<float> zeros;
+ bool zeros_warned = false; // the substitution is reported once
std::vector<common_sampler_ptr> smpls;
@@ -194,6 +198,8 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
throw std::runtime_error("draft-simple requires a draft context");
}
+ zeros.assign(llama_model_n_embd_inp(llama_get_model(ctx_dft)), 0.0f);
+
SPC_TRC("%s", "adding speculative implementation 'draft-simple'\n");
SPC_TRC("- n_max=%d, n_min=%d, p_min=%f\n", this->params.n_max, this->params.n_min, this->params.p_min);
SPC_TRC("- gpu_layers=%d, cache_k=%s, cache_v=%s, ctx_tgt=%s, ctx_dft=%s, devices=[%s]\n",
@@ -204,7 +210,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
ctx_dft ? "yes" : "no",
common_speculative_get_devices_str(this->params.devices).c_str());
- batch = llama_batch_init(llama_n_batch(ctx_dft), 0, 1);
+ batch = common_batch(ctx_dft);
// TODO: optimize or pass from outside?
// {
@@ -249,21 +255,46 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
}
}
- ~common_speculative_impl_draft_simple() override {
- llama_batch_free(batch);
- }
-
void begin(llama_seq_id /*seq_id*/, const llama_tokens & /*prompt*/) override {
// noop
}
- bool process(const llama_batch & batch) override {
+ bool process(const common_batch & batch_in) override {
auto * ctx_dft = params.ctx_dft;
- llama_batch batch_dft = batch;
- batch_dft.logits = nullptr;
+ // copy the entries to a batch owned by the draft context, only the last token is output
+ batch.clear();
+ const int32_t n_tokens = batch_in.size();
+ for (int32_t k = 0; k < n_tokens; ++k) {
+ const auto & t = batch_in.tokens[k];
+ const bool output = k == n_tokens - 1;
+ if (t.id != LLAMA_TOKEN_NULL) {
+ const int32_t idx = batch.add(t.id, t.pos[0], t.seq_id, output);
+ if (t.embd.data) {
+ batch.set_embd(idx, t.embd);
+ }
+ } else {
+ // mtmd input is projected by the target encoder, a draft with a different width cannot read it
+ // it gets zeros instead, keeping its positions contiguous
+ // ref: https://github.com/ggml-org/llama.cpp/pull/29385#discussion_r4124743243
+ const size_t n_embd = t.embd.n_rows * t.embd.n_embd;
+ const bool same_width = n_embd == zeros.size();
+ if (!same_width && !zeros_warned) {
+ SPC_WRN("target embeddings of size %zu do not fit the draft input width %zu, "
+ "the draft receives zero rows for them and drafts after multimodal input will be poor\n",
+ n_embd, zeros.size());
+ zeros_warned = true;
+ }
+ const llama_embd embd = same_width ? t.embd : llama_embd{ zeros.data(), 1, zeros.size() };
+ batch.add_embd(embd, t.pos.data(), t.seq_id, output);
+ }
+ }
- const int ret = llama_decode(ctx_dft, batch_dft);
+ if (batch.size() == 0) {
+ return true;
+ }
+
+ const int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
SPC_ERR("failed to decode draft batch, ret = %d\n", ret);
@@ -277,7 +308,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
void draft(common_speculative_draft_params_vec & dparams) override {
auto & ctx_dft = params.ctx_dft;
- common_batch_clear(batch);
+ batch.clear();
// keep track of which sequences are still drafting
int n_drafting = 0;
@@ -294,12 +325,12 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
drafting[seq_id] = true;
common_sampler_reset(smpls[seq_id].get());
- common_batch_add(batch, dp.id_last, dp.pos0, { seq_id }, true);
+ batch.add(dp.id_last, dp.pos0, seq_id, true);
}
- int ret = llama_decode(ctx_dft, batch);
+ int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
- SPC_ERR("llama_decode returned %d\n", ret);
+ SPC_ERR("llama_process returned %d\n", ret);
return;
}
@@ -308,7 +339,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
while (n_drafting > 0) {
int i_batch = 0;
- common_batch_clear(batch);
+ batch.clear();
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
if (!drafting[seq_id]) {
@@ -353,17 +384,17 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
continue;
}
- common_batch_add(batch, id, dp.pos0 + i + 1, { seq_id }, true);
+ batch.add(id, dp.pos0 + i + 1, seq_id, true);
}
- if (batch.n_tokens == 0) {
+ if (batch.size() == 0) {
break;
}
// evaluate the drafted tokens on the draft model
- ret = llama_decode(ctx_dft, batch);
+ ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
- SPC_ERR("llama_decode[%d] returned %d\n", i, ret);
+ SPC_ERR("llama_process[%d] returned %d\n", i, ret);
break;
}
@@ -423,7 +454,8 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
// encoder+decoder on n_accepted+1 rows).
struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
common_params_speculative_draft params;
- llama_batch batch;
+ common_batch batch; // decoder input, (token, g_embd) pairs
+ common_batch batch_enc; // encoder input, built from the extracted target features
std::vector<common_sampler_ptr> smpls;
@@ -477,11 +509,8 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
n_embd_enc = (int32_t) target_layer_ids_n * n_embd_tgt;
n_layer_tgt = llama_model_n_layer(model_tgt);
- const int32_t n_b = (int32_t) llama_n_batch(ctx_dft);
- batch = llama_batch_init(/*n_tokens=*/ n_b, /*embd=*/ n_embd_dec, /*n_seq_max=*/ 1);
- // llama_batch_init allocates only one of token/embd; eagle3 decoder needs both.
- // TODO: fix, how to call without malloc
- batch.token = (llama_token *) malloc(sizeof(llama_token) * n_b);
+ batch = common_batch(ctx_dft);
+ batch_enc = common_batch(ctx_dft);
smpls.resize(n_seq);
for (auto & s : smpls) {
@@ -543,12 +572,6 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
llama_sampler_free(backend_chains[seq_id]);
}
backend_chains.clear();
-
- if (batch.token != nullptr) {
- free(batch.token);
- batch.token = nullptr;
- }
- llama_batch_free(batch);
}
void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {
@@ -567,16 +590,16 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
}
}
- bool process(const llama_batch & batch_in) override {
- if (batch_in.n_tokens <= 0) {
+ bool process(const common_batch & batch_in) override {
+ if (batch_in.size() <= 0) {
return true;
}
- if (batch_in.token == nullptr || batch_in.embd != nullptr) {
+ if (!batch_in.has_token() || batch_in.has_embd()) {
return true;
}
- const int32_t n_tokens = batch_in.n_tokens;
+ const int32_t n_tokens = batch_in.size();
// i_batch_beg[seq] / i_batch_end[seq]: inclusive batch indices of this seq's
// first/last token in batch_in. Assumes per-seq tokens are contiguous within
@@ -584,8 +607,7 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
std::vector<int32_t> i_batch_beg(n_seq, -1);
std::vector<int32_t> i_batch_end(n_seq, -1);
for (int k = 0; k < n_tokens; ++k) {
- GGML_ASSERT(batch_in.n_seq_id[k] == 1);
- const llama_seq_id seq_id = batch_in.seq_id[k][0];
+ const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {
continue;
}
@@ -619,24 +641,23 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
g_embd_buf.resize((size_t) n_tokens * n_embd_dec);
- // llama_encode() requires the full encoder batch to fit in n_ubatch.
+ // llama_process() requires the full encoder batch to fit in n_ubatch.
// Allow batch > ubatch: eagle3's per-token encoder can be chunked safely.
const int32_t n_ubatch_dft = (int32_t) llama_n_ubatch(ctx_dft);
for (int32_t i = 0; i < n_tokens; i += n_ubatch_dft) {
const int32_t n_chunk = std::min(n_ubatch_dft, n_tokens - i);
- llama_batch enc_batch = {
- /*.n_tokens =*/ n_chunk,
- /*.token =*/ nullptr,
- /*.embd =*/ features_buf.data() + (size_t) i * n_embd_enc,
- /*.pos =*/ nullptr,
- /*.n_seq_id =*/ nullptr,
- /*.seq_id =*/ nullptr,
- /*.logits =*/ nullptr,
- };
- const int32_t rc = llama_encode(ctx_dft, enc_batch);
+ // the per-token encoder does not use positions, generate placeholder ones from the memory state
+ batch_enc.clear();
+ llama_pos pos = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), 0) + 1;
+ for (int32_t j = 0; j < n_chunk; ++j) {
+ batch_enc.add_embd({ features_buf.data() + (size_t) (i + j) * n_embd_enc, 1, (size_t) n_embd_enc }, &pos, 0, true);
+ pos++;
+ }
+
+ const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_ENCODE, batch_enc.get());
if (rc != 0) {
- SPC_ERR("llama_encode(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
+ SPC_ERR("llama_process(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
rc, (int) n_chunk, (int) i);
return false;
}
@@ -664,7 +685,7 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
// deferred boundary, completed by the next process() or draft() call.
// (c) refresh deferred state — stash this ubatch's full g_embd into verify_g,
// update pending_g_last / pending_pos_last to the last row.
- common_batch_clear(batch);
+ batch.clear();
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
const int32_t beg = i_batch_beg[seq_id];
@@ -679,36 +700,34 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
// 2) pending_pos_last + 1 == pos[beg]
// 3) pending_pos_last > dft_pos_max // TODO: is this check needed?
const llama_pos pending_pos = pending_pos_last[seq_id];
- if (pending_pos >= 0 && pending_pos + 1 == batch_in.pos[beg]) {
+ if (pending_pos >= 0 && pending_pos + 1 == batch_in.tokens[beg].pos[0]) {
const llama_pos dft_pos_max = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), seq_id);
if (pending_pos > dft_pos_max) {
- common_batch_add(batch, batch_in.token[beg], pending_pos, { seq_id }, /*logits=*/ false);
- std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,
- pending_g_last[seq_id].data(), row_bytes);
+ const int32_t idx = batch.add(batch_in.tokens[beg].id, pending_pos, seq_id, /*output=*/ false);
+ batch.set_embd(idx, { pending_g_last[seq_id].data(), 1, (size_t) n_embd_dec });
}
}
for (int32_t k = beg; k < end; ++k) {
- common_batch_add(batch, batch_in.token[k + 1], batch_in.pos[k], { seq_id }, /*logits=*/ false);
- std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,
- g_embd + (size_t) k * n_embd_dec, row_bytes);
+ const int32_t idx = batch.add(batch_in.tokens[k + 1].id, batch_in.tokens[k].pos[0], seq_id, /*output=*/ false);
+ batch.set_embd(idx, { g_embd + (size_t) k * n_embd_dec, 1, (size_t) n_embd_dec });
}
// refresh deferred state
const int32_t n_rows = end - beg + 1;
- verify_pos_first[seq_id] = batch_in.pos[beg];
- pending_pos_last[seq_id] = batch_in.pos[end];
+ verify_pos_first[seq_id] = batch_in.tokens[beg].pos[0];
+ pending_pos_last[seq_id] = batch_in.tokens[end].pos[0];
verify_g_rows[seq_id] = n_rows;
verify_g[seq_id].resize((size_t) n_rows * n_embd_dec, 0.0f);
std::memcpy(verify_g[seq_id].data(), g_embd + (size_t) beg * n_embd_dec, row_bytes * n_rows);
std::memcpy(pending_g_last[seq_id].data(), g_embd + (size_t) end * n_embd_dec, row_bytes);
}
- if (batch.n_tokens > 0) {
- const int32_t rc = llama_decode(ctx_dft, batch);
+ if (batch.size() > 0) {
+ const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (rc != 0) {
- SPC_ERR("llama_decode(ctx_dft) failed rc=%d (n_tokens=%d, ubatch_pos[0]=%d)\n",
- rc, (int) batch.n_tokens, (int) batch_in.pos[0]);
+ SPC_ERR("llama_process(ctx_dft) failed rc=%d (n_tokens=%d, ubatch_pos[0]=%d)\n",
+ rc, (int) batch.size(), (int) batch_in.tokens[0].pos[0]);
return false;
}
}
@@ -719,14 +738,12 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
void draft(common_speculative_draft_params_vec & dparams) override {
auto & ctx_dft = params.ctx_dft;
- common_batch_clear(batch);
+ batch.clear();
// keep track of which sequences are still drafting
int n_drafting = 0;
std::vector<bool> drafting(n_seq);
- const size_t row_bytes = (size_t) n_embd_dec * sizeof(float);
-
// Complete the deferred boundary pair (dp.id_last, pending_g_last) at memory
// pos pending_pos_last. dp.id_last is target's freshest sample (= corrected
// token after verify, or first generated token after prefill), matching the
@@ -747,19 +764,17 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
llama_memory_seq_rm(llama_get_memory(ctx_dft), seq_id, pending_pos_last[seq_id], -1);
- common_batch_add(batch, dp.id_last, pending_pos_last[seq_id], { seq_id }, true);
- std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,
- pending_g_last[seq_id].data(),
- row_bytes);
+ const int32_t idx = batch.add(dp.id_last, pending_pos_last[seq_id], seq_id, true);
+ batch.set_embd(idx, { pending_g_last[seq_id].data(), 1, (size_t) n_embd_dec });
}
- if (batch.n_tokens == 0) {
+ if (batch.size() == 0) {
return;
}
- int ret = llama_decode(ctx_dft, batch);
+ int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
- SPC_ERR("llama_decode returned %d\n", ret);
+ SPC_ERR("llama_process returned %d\n", ret);
return;
}
@@ -768,7 +783,7 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
while (n_drafting > 0) {
int i_batch = 0;
- common_batch_clear(batch);
+ batch.clear();
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
if (!drafting[seq_id]) {
@@ -814,17 +829,17 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
continue;
}
- common_batch_add(batch, id, pending_pos_last[seq_id] + (i + 1), { seq_id }, true);
- std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec, prenorm, row_bytes);
+ const int32_t idx = batch.add(id, pending_pos_last[seq_id] + (i + 1), seq_id, true);
+ batch.set_embd(idx, { prenorm, 1, (size_t) n_embd_dec });
}
- if (batch.n_tokens == 0) {
+ if (batch.size() == 0) {
break;
}
- ret = llama_decode(ctx_dft, batch);
+ ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
- SPC_ERR("llama_decode[%d] returned %d\n", i, ret);
+ SPC_ERR("llama_process[%d] returned %d\n", i, ret);
break;
}
@@ -908,8 +923,10 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
struct common_speculative_impl_draft_dflash : public common_speculative_impl {
common_params_speculative_draft params;
- llama_batch batch; // noise tokens
- llama_batch batch_inject; // target features for KV cache injection
+ common_batch batch; // noise tokens
+ common_batch batch_inject; // target features for KV cache injection
+
+ std::vector<float> features_buf; // [n_chunk, n_embd_enc] gathered target features
std::vector<common_sampler_ptr> smpls;
@@ -1005,15 +1022,11 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
}
this->n_max = this->params.n_max;
- batch = llama_batch_init(llama_n_batch(ctx_dft), 0, n_seq);
- batch_inject = llama_batch_init(llama_n_ubatch(ctx_dft), n_embd_enc, n_seq);
+ batch = common_batch(ctx_dft);
+ batch_inject = common_batch(ctx_dft);
- // embd batches on an M-RoPE draft need 4 position rows per token
+ // embd batches on an M-RoPE draft carry 4 position rows per token
is_mrope = llama_model_rope_type(model_dft) == LLAMA_ROPE_TYPE_MROPE;
- if (is_mrope) {
- free(batch_inject.pos);
- batch_inject.pos = (llama_pos *) malloc(sizeof(llama_pos) * 4 * llama_n_batch(ctx_dft));
- }
smpls.resize(n_seq);
for (auto & s : smpls) {
@@ -1062,9 +1075,6 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
llama_sampler_free(backend_chains[seq_id]);
}
backend_chains.clear();
-
- llama_batch_free(batch);
- llama_batch_free(batch_inject);
}
void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {
@@ -1085,8 +1095,8 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
}
}
- bool process(const llama_batch & batch_in) override {
- if (batch_in.n_tokens <= 0) {
+ bool process(const common_batch & batch_in) override {
+ if (batch_in.size() <= 0) {
return true;
}
@@ -1094,20 +1104,19 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
// produce the target-layer features used to seed the draft KV cache, so
// embeddings are injected too, except the pinned ones skipped below.
// TODO: revisit after https://github.com/ggml-org/llama.cpp/pull/24669 is merged
- const bool has_tokens = batch_in.token != nullptr;
- const bool has_embeddings = batch_in.embd != nullptr;
+ const bool has_tokens = batch_in.has_token();
+ const bool has_embeddings = batch_in.has_embd();
if (has_tokens == has_embeddings) {
return true;
}
- const int32_t n_tokens = batch_in.n_tokens;
+ const int32_t n_tokens = batch_in.size();
// per-seq inclusive batch range (assumes each seq's tokens are contiguous in the batch)
std::vector<int32_t> i_batch_beg(n_seq, -1);
std::vector<int32_t> i_batch_end(n_seq, -1);
for (int32_t k = 0; k < n_tokens; ++k) {
- GGML_ASSERT(batch_in.n_seq_id[k] == 1);
- const llama_seq_id seq_id = batch_in.seq_id[k][0];
+ const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {
continue;
}
@@ -1130,7 +1139,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
// an M-RoPE image pins all its rows to one position, so a windowed draft
// cache cannot free cells for it - skip it, the draft can jump over the gap
- const bool pos_pinned = batch_in.pos[i_batch_beg[seq_id]] == batch_in.pos[i_batch_end[seq_id]];
+ const bool pos_pinned = batch_in.tokens[i_batch_beg[seq_id]].pos[0] == batch_in.tokens[i_batch_end[seq_id]].pos[0];
if (has_embeddings && n_rows > 1 && pos_pinned) {
continue;
}
@@ -1140,34 +1149,28 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
// gather target features per extract layer; the fused decode encodes and
// injects them into the K/V cache at the target positions
- batch_inject.n_tokens = n_chunk;
+ features_buf.resize((size_t) n_chunk * n_embd_enc);
for (uint32_t k = 0; k < target_layer_ids_n; ++k) {
const float * layer = llama_get_embeddings_layer_inp(ctx_tgt, (uint32_t) target_layer_ids[k]);
if (!layer) {
GGML_ABORT("DFlash: target layer %d input not extracted.", target_layer_ids[k]);
}
for (int32_t i = 0; i < n_chunk; ++i) {
- float * dst = batch_inject.embd + (size_t) i * n_embd_enc + k * (size_t) n_embd_tgt;
+ float * dst = features_buf.data() + (size_t) i * n_embd_enc + k * (size_t) n_embd_tgt;
const float * src = layer + (size_t) (i_batch_beg[seq_id] + offset + i) * n_embd_tgt;
std::memcpy(dst, src, (size_t) n_embd_tgt * sizeof(float));
}
}
+ batch_inject.clear();
for (int32_t i = 0; i < n_chunk; ++i) {
- const llama_pos p = batch_in.pos[i_batch_beg[seq_id] + offset + i];
- batch_inject.pos[i] = p;
- if (is_mrope) {
- batch_inject.pos[1 * n_chunk + i] = p;
- batch_inject.pos[2 * n_chunk + i] = p;
- batch_inject.pos[3 * n_chunk + i] = 0;
- }
- batch_inject.n_seq_id[i] = 1;
- batch_inject.seq_id[i][0] = seq_id;
- batch_inject.logits[i] = false;
+ const llama_pos p = batch_in.tokens[i_batch_beg[seq_id] + offset + i].pos[0];
+ const llama_pos pos_arr[4] = { p, p, p, 0 };
+ batch_inject.add_embd({ features_buf.data() + (size_t) i * n_embd_enc, 1, (size_t) n_embd_enc }, pos_arr, seq_id, false);
}
- const int32_t rc = llama_decode(ctx_dft, batch_inject);
+ const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch_inject.get());
if (rc != 0) {
- LOG_ERR("%s: llama_decode(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
+ LOG_ERR("%s: llama_process(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
__func__, rc, (int) n_chunk, (int) offset);
return false;
}
@@ -1180,7 +1183,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
void draft(common_speculative_draft_params_vec & dparams) override {
auto & ctx_dft = params.ctx_dft;
- common_batch_clear(batch);
+ batch.clear();
// build one batch holding every drafting sequence's noise block into a single decode)
// record where each block starts and its size
@@ -1200,21 +1203,21 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
const int32_t n_draft = params.n_max;
const int32_t n_block_tokens = n_draft + (is_dspark && sample_from_anchor ? 0 : 1);
- i_block_beg[seq_id] = batch.n_tokens;
+ i_block_beg[seq_id] = batch.size();
n_block [seq_id] = n_block_tokens;
for (int32_t i = 0; i < n_block_tokens; ++i) {
- common_batch_add(batch, i == 0 ? dp.id_last : mask_token_id, n + i, { seq_id }, !is_dflash2);
+ batch.add(i == 0 ? dp.id_last : mask_token_id, n + i, seq_id, !is_dflash2);
}
}
- if (batch.n_tokens == 0) {
+ if (batch.size() == 0) {
return;
}
// decode all sequence's noise block in a single batch
- int ret = llama_decode(ctx_dft, batch);
+ int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
- LOG_WRN("%s: llama_decode returned %d\n", __func__, ret);
+ LOG_WRN("%s: llama_process returned %d\n", __func__, ret);
return;
}
@@ -1328,7 +1331,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
struct common_speculative_impl_draft_mtp : public common_speculative_impl {
common_params_speculative_draft params; // reuses the draft-model params slot (ctx_tgt/ctx_dft)
- llama_batch batch;
+ common_batch batch;
std::vector<common_sampler_ptr> smpls;
@@ -1384,11 +1387,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
ctx_dft ? "yes" : "no",
common_speculative_get_devices_str(this->params.devices).c_str());
- const int32_t n_b = (int32_t) llama_n_batch(ctx_dft);
- batch = llama_batch_init(/*n_tokens=*/ n_b, /*embd=*/ n_embd, /*n_seq_max=*/ 1);
- // llama_batch_init allocates only one of token/embd; MTP needs both.
- // TODO: fix, how to call without malloc
- batch.token = (llama_token *) malloc(sizeof(llama_token) * n_b);
+ batch = common_batch(ctx_dft);
smpls.resize(n_seq);
for (auto & s : smpls) {
@@ -1453,12 +1452,6 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
llama_sampler_free(backend_chains[seq_id]);
}
backend_chains.clear();
-
- if (batch.token != nullptr) {
- free(batch.token);
- batch.token = nullptr;
- }
- llama_batch_free(batch);
}
void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {
@@ -1473,23 +1466,23 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
if (pos_max < N - 1 && !is_mem_shared) {
SPC_WRN("ctx_dft pos_max=%d < N-1=%d - "
"process() hook may not have run on every prefill ubatch "
- "(need_embd / logits=1 on every prompt position?). "
+ "(need_embd / output flag on every prompt position?). "
"Drafts may degrade.\n",
(int) pos_max, N - 1);
}
}
- bool process(const llama_batch & batch_in) override {
- if (batch_in.n_tokens <= 0) {
+ bool process(const common_batch & batch_in) override {
+ if (batch_in.size() <= 0) {
return true;
}
// TODO: how to make it work with vision tokens?
- if (batch_in.token == nullptr || batch_in.embd != nullptr) {
+ if (!batch_in.has_token() || batch_in.has_embd()) {
return true;
}
- const int32_t n_tokens = batch_in.n_tokens;
+ const int32_t n_tokens = batch_in.size();
// remember the first and last batch index for each sequence
std::fill(i_batch_beg.begin(), i_batch_beg.end(), -1);
@@ -1497,9 +1490,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
for (int k = 0; k < n_tokens; ++k) {
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
- GGML_ASSERT(batch_in.n_seq_id[k] == 1);
-
- if (batch_in.seq_id[k][0] == seq_id) {
+ if (batch_in.tokens[k].seq_id == seq_id) {
i_batch_end[seq_id] = k;
if (i_batch_beg[seq_id] < 0) {
i_batch_beg[seq_id] = k;
@@ -1515,33 +1506,26 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
// if kv is shared with target (e.g Gemma4), then we can skip this catch-up decode
if (!is_mem_shared) {
- common_batch_clear(batch);
-
- for (int k = 0; k < n_tokens; ++k) {
- common_batch_add(batch, batch_in.token[k], batch_in.pos[k], { batch_in.seq_id[k][0] }, 0);
- }
+ batch.clear();
- // shift the tgt embeddings to the right by one position
+ // pair each token with the tgt embedding shifted right by one position, and
+ // the first token of each sequence with the pending embedding from a previous run
// assumes that the tokens in the batch are sequential for each sequence
// i.e. we cannot have seq_id like this: [0, 0, 0, 1, 1, 0, 1, 1]
// ^--- this is a problem
// TODO:this is generally true, but would be nice to assert it
- {
- const float * h_tgt = llama_get_embeddings_nextn(ctx_tgt);
- std::memcpy(batch.embd + (size_t) 1 * n_embd, h_tgt, row_bytes * (n_tokens-1));
- }
+ const float * h_tgt = llama_get_embeddings_nextn(ctx_tgt);
- // fill the pending embeddings from a previous run
- auto set_h = [&](int idx, const float * h_row) {
- std::memcpy(batch.embd + (size_t) idx * n_embd, h_row, row_bytes);
- };
+ for (int k = 0; k < n_tokens; ++k) {
+ const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
- for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
- if (i_batch_beg[seq_id] < 0) {
- continue;
- }
+ const int32_t idx = batch.add(batch_in.tokens[k].id, batch_in.tokens[k].pos[0], seq_id, false);
+
+ const float * h_row = k == i_batch_beg[seq_id]
+ ? pending_h[seq_id].data()
+ : h_tgt + (size_t) (k - 1) * n_embd;
- set_h(i_batch_beg[seq_id], pending_h[seq_id].data());
+ batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
}
auto * mem_dft = llama_get_memory(ctx_dft);
@@ -1554,15 +1538,15 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
if (i_batch_beg[seq_id] < 0) {
continue;
}
- llama_memory_seq_rm(mem_dft, seq_id, batch_in.pos[i_batch_beg[seq_id]], -1);
+ llama_memory_seq_rm(mem_dft, seq_id, batch_in.tokens[i_batch_beg[seq_id]].pos[0], -1);
}
llama_set_nextn_layer_offset(ctx_dft, head);
}
- const int32_t rc = llama_decode(ctx_dft, batch);
+ const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (rc != 0) {
- SPC_ERR("llama_decode(ctx_dft) head=%d failed rc=%d (pos=%d)\n",
- head, (int) rc, (int) batch_in.pos[0]);
+ SPC_ERR("llama_process(ctx_dft) head=%d failed rc=%d (pos=%d)\n",
+ head, (int) rc, (int) batch_in.tokens[0].pos[0]);
ok = false;
break;
}
@@ -1600,14 +1584,12 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
void draft(common_speculative_draft_params_vec & dparams) override {
auto & ctx_dft = params.ctx_dft;
- common_batch_clear(batch);
+ batch.clear();
// keep track of which sequences are still drafting
int n_drafting = 0;
std::vector<bool> drafting(n_seq);
- const size_t row_bytes = (size_t) n_embd * sizeof(float);
-
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
auto & dp = dparams[seq_id];
@@ -1619,10 +1601,10 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
drafting[seq_id] = true;
common_sampler_reset(smpls[seq_id].get());
- common_batch_add(batch, dp.id_last, dp.pos0, { seq_id }, true);
- std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, pending_h[seq_id].data(), row_bytes);
+ const int32_t idx = batch.add(dp.id_last, dp.pos0, seq_id, true);
+ batch.set_embd(idx, { pending_h[seq_id].data(), 1, (size_t) n_embd });
- i_last[seq_id] = batch.n_tokens - 1;
+ i_last[seq_id] = idx;
if (chain_heads) {
chain_h[seq_id].assign(pending_h[seq_id].begin(), pending_h[seq_id].end());
@@ -1648,16 +1630,16 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
llama_set_nextn_layer_offset(ctx_dft, i);
}
- int ret = llama_decode(ctx_dft, batch);
+ int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
- SPC_ERR("llama_decode[%d] returned %d\n", i, ret);
+ SPC_ERR("llama_process[%d] returned %d\n", i, ret);
break;
}
// rebuild the batch for the next step: the growing-KV paths re-add only the
// new token (the KV already holds the prefix), while chained heads re-add the
// whole prefix at the next head. dropped sequences are simply not re-added.
- common_batch_clear(batch);
+ batch.clear();
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
if (!drafting[seq_id]) {
@@ -1708,24 +1690,24 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
const int n_rows = (int) result.size() + 1; // id_last + tokens drafted so far
for (int t = 0; t < n_rows; ++t) {
const llama_token tok = (t == 0) ? dp.id_last : result[t - 1];
- common_batch_add(batch, tok, dp.pos0 + t, { seq_id }, t == n_rows - 1);
- std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd,
- chain_h[seq_id].data() + (size_t) t * n_embd, row_bytes);
+ const int32_t idx = batch.add(tok, dp.pos0 + t, seq_id, t == n_rows - 1);
+ batch.set_embd(idx, { chain_h[seq_id].data() + (size_t) t * n_embd, 1, (size_t) n_embd });
+ i_last[seq_id] = idx;
}
} else if (is_mem_shared) {
// note: with shared memory (e.g. Gemma4 assistants) we use the same position for all draft tokens
// ref: https://github.com/huggingface/transformers/blob/effde20942e3f82a1b97449f60b3a48c5ff96145/docs/source/en/model_doc/gemma4_assistant.md?plain=1#L36-L37
- common_batch_add(batch, id, dp.pos0, { seq_id }, true);
- std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, h_row, row_bytes);
+ const int32_t idx = batch.add(id, dp.pos0, seq_id, true);
+ batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
+ i_last[seq_id] = idx;
} else {
- common_batch_add(batch, id, dp.pos0 + i + 1, { seq_id }, true);
- std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, h_row, row_bytes);
+ const int32_t idx = batch.add(id, dp.pos0 + i + 1, seq_id, true);
+ batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
+ i_last[seq_id] = idx;
}
-
- i_last[seq_id] = batch.n_tokens - 1;
}
- if (batch.n_tokens == 0) {
+ if (batch.size() == 0) {
break;
}
@@ -1787,7 +1769,7 @@ struct common_speculative_impl_ngram_simple : public common_speculative_impl {
// noop
}
- bool process(const llama_batch & /*batch*/) override {
+ bool process(const common_batch & /*batch*/) override {
// TODO: implement
return true;
}
@@ -1835,7 +1817,7 @@ struct common_speculative_impl_ngram_map_k : public common_speculative_impl {
common_ngram_map_begin(config[seq_id], prompt);
}
- bool process(const llama_batch & /*batch*/) override {
+ bool process(const common_batch & /*batch*/) override {
// TODO: implement
return true;
}
@@ -1993,7 +1975,7 @@ struct common_speculative_impl_ngram_mod : public common_speculative_impl {
sinfo.n_draft_last = result.size();
}
- bool process(const llama_batch & /*batch*/) override {
+ bool process(const common_batch & /*batch*/) override {
// TODO: implement
return true;
}
@@ -2155,7 +2137,7 @@ struct common_speculative_impl_ngram_cache : public common_speculative_impl {
}
}
- bool process(const llama_batch & /*batch*/) override {
+ bool process(const common_batch & /*batch*/) override {
// TODO: implement
return true;
}
@@ -2181,6 +2163,9 @@ struct common_speculative_impl_ngram_cache : public common_speculative_impl {
struct common_speculative {
common_speculative_draft_params_vec dparams;
+ // the target context, used to convert legacy llama_batch inputs
+ llama_context * ctx_tgt = nullptr;
+
// list of implementations to use and their states
std::vector<std::unique_ptr<common_speculative_impl>> impls;
@@ -2726,6 +2711,7 @@ common_speculative * common_speculative_init(common_params_speculative & params,
common_speculative_ptr result(new common_speculative {
/* .dparams = */ common_speculative_draft_params_vec(n_seq),
+ /* .ctx_tgt = */ params.draft.ctx_tgt,
/* .impls = */ std::move(impls),
/* .impl_last = */ std::vector<common_speculative_impl *>(n_seq, nullptr),
/* .synth_probs = */ {},
@@ -2789,6 +2775,17 @@ void common_speculative_begin(common_speculative * spec, llama_seq_id seq_id, co
}
bool common_speculative_process(common_speculative * spec, const llama_batch & batch) {
+ if (spec == nullptr) {
+ return true;
+ }
+
+ // ngram-only setups have no target context, they do not read the batch anyway
+ const common_batch tmp = spec->ctx_tgt ? common_batch_from_llama_batch(spec->ctx_tgt, batch) : common_batch();
+
+ return common_speculative_process(spec, tmp);
+}
+
+bool common_speculative_process(common_speculative * spec, const common_batch & batch) {
bool result = true;
if (spec == nullptr) {
diff --git a/common/speculative.h b/common/speculative.h
index c968750e2..211fcdabd 100644
--- a/common/speculative.h
+++ b/common/speculative.h
@@ -77,6 +77,9 @@ common_speculative_draft_params & common_speculative_get_draft_params(common_spe
void common_speculative_begin(common_speculative * spec, llama_seq_id seq_id, const llama_tokens & prompt);
// process the batch and update the internal state of the speculative context
+bool common_speculative_process(common_speculative * spec, const common_batch & batch);
+
+// legacy llama_batch input, converted with common_batch_from_llama_batch()
bool common_speculative_process(common_speculative * spec, const llama_batch & batch);
// generate drafts for the sequences specified with `common_speculative_get_draft_params`
diff --git a/examples/speculative-simple/speculative-simple.cpp b/examples/speculative-simple/speculative-simple.cpp
index 863af5a2c..81aa106f1 100644
--- a/examples/speculative-simple/speculative-simple.cpp
+++ b/examples/speculative-simple/speculative-simple.cpp
@@ -228,7 +228,6 @@ int main(int argc, char ** argv) {
common_batch_add(batch_tgt, draft[i], n_past + i, { seq_id }, true);
}
- //LOG_DBG("target batch: %s\n", string_from(ctx_tgt, batch_tgt).c_str());
llama_decode(ctx_tgt, batch_tgt);
}
diff --git a/tools/mtmd/mtmd-cli.cpp b/tools/mtmd/mtmd-cli.cpp
index ba18b3e32..6fe058fd8 100644
--- a/tools/mtmd/mtmd-cli.cpp
+++ b/tools/mtmd/mtmd-cli.cpp
@@ -81,7 +81,7 @@ struct mtmd_cli_context {
llama_context * lctx;
const llama_vocab * vocab;
common_sampler * smpl;
- llama_batch batch;
+ common_batch batch;
int n_batch;
mtmd::bitmaps bitmaps;
@@ -115,7 +115,7 @@ struct mtmd_cli_context {
vocab = llama_model_get_vocab(model);
smpl = common_sampler_init(model, params.sampling);
n_threads = params.cpuparams.n_threads;
- batch = llama_batch_init(1, 0, 1); // batch for next token generation
+ batch = common_batch(lctx); // batch for next token generation
n_batch = params.n_batch;
init_vision_context(params);
@@ -148,7 +148,6 @@ struct mtmd_cli_context {
}
~mtmd_cli_context() {
- llama_batch_free(batch);
common_sampler_free(smpl);
}
@@ -230,9 +229,9 @@ static int generate_response(mtmd_cli_context & ctx, int n_predict) {
}
// eval the token
- common_batch_clear(ctx.batch);
- common_batch_add(ctx.batch, token_id, ctx.n_past++, {0}, true);
- if (llama_decode(ctx.lctx, ctx.batch)) {
+ ctx.batch.clear();
+ ctx.batch.add(token_id, ctx.n_past++, 0, true);
+ if (llama_process(ctx.lctx, LLAMA_PROCESS_TYPE_DECODE, ctx.batch.get())) {
LOG_ERR("failed to decode token\n");
return 1;
}
diff --git a/tools/mtmd/mtmd-helper-common.h b/tools/mtmd/mtmd-helper-common.h
index f907346c7..bc68ed4b9 100644
--- a/tools/mtmd/mtmd-helper-common.h
+++ b/tools/mtmd/mtmd-helper-common.h
@@ -6,6 +6,7 @@
#include "ggml.h"
#include "llama.h"
+#include "llama-cpp.h"
#include "mtmd.h"
#include <cstdarg>
@@ -73,112 +74,99 @@ inline mtmd_helper_logger g_logger;
struct decode_embd_batch {
int n_pos_per_embd;
int n_mmproj_embd;
- std::vector<llama_pos> pos;
- std::vector<llama_pos> pos_view; // used by mrope
- std::vector<int32_t> n_seq_id;
- std::vector<llama_seq_id> seq_id_0;
- std::vector<llama_seq_id *> seq_ids;
- std::vector<int8_t> logits;
- llama_batch batch;
- decode_embd_batch(float * embd, int32_t n_tokens, int n_pos_per_embd, int n_mmproj_embd) : n_pos_per_embd(n_pos_per_embd), n_mmproj_embd(n_mmproj_embd) {
+ int32_t n_tokens;
+ const float * embd; // [n_tokens, n_mmproj_embd], not owned
+ std::vector<llama_pos> pos; // [n_pos_per_embd, n_tokens], section-major
+ std::vector<llama_pos> pos_view; // sliced positions of the last get_view()
+ std::vector<int8_t> logits;
+ llama_seq_id seq_id = 0;
+
+ llama_batch_ext_ptr batch; // rendered sub-batch, see render()
+
+ decode_embd_batch(const float * embd, int32_t n_tokens, int n_pos_per_embd, int n_mmproj_embd)
+ : n_pos_per_embd(n_pos_per_embd), n_mmproj_embd(n_mmproj_embd), n_tokens(n_tokens), embd(embd) {
GGML_ASSERT(n_tokens > 0 && n_pos_per_embd > 0 && n_mmproj_embd > 0);
- pos .resize((size_t) n_tokens * (size_t) n_pos_per_embd);
- n_seq_id.resize(n_tokens);
- seq_ids .resize(n_tokens + 1);
- logits .resize(n_tokens);
- seq_id_0.resize(1);
- seq_ids [n_tokens] = nullptr;
- batch = {
- /*n_tokens =*/ n_tokens,
- /*tokens =*/ nullptr,
- /*embd =*/ embd,
- /*pos =*/ pos.data(),
- /*n_seq_id =*/ n_seq_id.data(),
- /*seq_id =*/ seq_ids.data(),
- /*logits =*/ logits.data(),
- };
+ pos .resize((size_t) n_tokens * (size_t) n_pos_per_embd);
+ logits.resize(n_tokens, 0);
}
void set_position_normal(llama_pos pos_0, llama_seq_id seq_id) {
- seq_id_0[0] = seq_id;
- for (int i = 0; i < batch.n_tokens; i++) {
- batch.pos [i] = pos_0 + i;
- batch.n_seq_id[i] = 1;
- batch.seq_id [i] = seq_id_0.data();
- batch.logits [i] = false;
+ this->seq_id = seq_id;
+ for (int i = 0; i < n_tokens; i++) {
+ pos[i] = pos_0 + i;
}
}
// M-RoPE for image
void set_position_mrope_2d(const std::vector<mtmd_decoder_pos> & rel_pos, llama_seq_id seq_id) {
GGML_ASSERT(n_pos_per_embd == 4);
- GGML_ASSERT(!rel_pos.empty() && (int32_t)rel_pos.size() == batch.n_tokens);
- seq_id_0[0] = seq_id;
- for (int32_t i = 0; i < batch.n_tokens; i++) {
+ GGML_ASSERT(!rel_pos.empty() && (int32_t)rel_pos.size() == n_tokens);
+ this->seq_id = seq_id;
+ for (int32_t i = 0; i < n_tokens; i++) {
const size_t idx = (size_t) i;
- const size_t n_tokens = (size_t) batch.n_tokens;
- pos[idx ] = rel_pos[i].t;
- pos[idx + n_tokens ] = rel_pos[i].y;
- pos[idx + n_tokens * 2 ] = rel_pos[i].x;
- pos[idx + n_tokens * 3 ] = rel_pos[i].z;
- }
- for (int i = 0; i < batch.n_tokens; i++) {
- batch.n_seq_id[i] = 1;
- batch.seq_id [i] = seq_id_0.data();
- batch.logits [i] = false;
+ const size_t n = (size_t) n_tokens;
+ pos[idx ] = rel_pos[i].t;
+ pos[idx + n ] = rel_pos[i].y;
+ pos[idx + n * 2] = rel_pos[i].x;
+ pos[idx + n * 3] = rel_pos[i].z;
}
}
// M-RoPE for audio
void set_position_mrope_1d(llama_pos pos_0, llama_seq_id seq_id) {
GGML_ASSERT(n_pos_per_embd == 4);
- seq_id_0[0] = seq_id;
- for (int i = 0; i < batch.n_tokens; i++) {
+ this->seq_id = seq_id;
+ for (int i = 0; i < n_tokens; i++) {
const size_t idx = (size_t) i;
- const size_t n_tokens = (size_t) batch.n_tokens;
- pos[idx ] = pos_0 + i;
- pos[idx + n_tokens ] = pos_0 + i;
- pos[idx + n_tokens * 2 ] = pos_0 + i;
- pos[idx + n_tokens * 3 ] = pos_0 + i;
- }
- for (int i = 0; i < batch.n_tokens; i++) {
- batch.n_seq_id[i] = 1;
- batch.seq_id [i] = seq_id_0.data();
- batch.logits [i] = false;
+ const size_t n = (size_t) n_tokens;
+ pos[idx ] = pos_0 + i;
+ pos[idx + n ] = pos_0 + i;
+ pos[idx + n * 2] = pos_0 + i;
+ pos[idx + n * 3] = pos_0 + i;
}
}
- llama_batch get_view(int offset, int n_tokens) {
- GGML_ASSERT(offset >= 0 && n_tokens > 0 && offset + n_tokens <= batch.n_tokens);
- llama_pos * pos_ptr;
+ // describe the entries [offset, offset + n) with section-major positions
+ mtmd_helper_embd_batch get_view(int offset, int n) {
+ GGML_ASSERT(offset >= 0 && n > 0 && offset + n <= n_tokens);
pos_view.clear();
- pos_view.reserve((size_t) n_tokens * (size_t) n_pos_per_embd);
- if (n_pos_per_embd > 1) {
- // mrope
- // for example, with layout of src: 1234...1234...1234...1234...
- // offset 2 will give us dst: 34...34...34...34...
- for (int i = 0; i < n_pos_per_embd; i++) {
- // assume n_tokens is less than or equal to batch.n_tokens
- // batch.n_tokens is number of **total** tokens
- // n_tokens is number of viewed token
- size_t src_idx = (size_t) i * (size_t) batch.n_tokens + (size_t) offset;
- pos_view.insert(pos_view.end(),
- pos.data() + src_idx,
- pos.data() + src_idx + n_tokens);
- }
- pos_ptr = pos_view.data();
- } else {
- // normal
- pos_ptr = pos.data() + offset;
+ pos_view.reserve((size_t) n * (size_t) n_pos_per_embd);
+ for (int j = 0; j < n_pos_per_embd; j++) {
+ const size_t src = (size_t) j * (size_t) n_tokens + (size_t) offset;
+ pos_view.insert(pos_view.end(), pos.data() + src, pos.data() + src + n);
}
return {
- /*n_tokens =*/ n_tokens,
- /*tokens =*/ nullptr,
- /*embd =*/ batch.embd + offset * n_mmproj_embd,
- /*pos =*/ pos_ptr,
- /*n_seq_id =*/ batch.n_seq_id + offset,
- /*seq_id =*/ batch.seq_id + offset,
- /*logits =*/ batch.logits + offset,
+ /*n_tokens =*/ n,
+ /*embd =*/ embd + (size_t) offset * n_mmproj_embd,
+ /*n_embd =*/ n_mmproj_embd,
+ /*pos =*/ pos_view.data(),
+ /*n_pos =*/ n_pos_per_embd,
+ /*seq_id =*/ seq_id,
};
}
+
+ // render the entries [offset, offset + n) into a batch owned by this object, ready for llama_process()
+ llama_batch_ext * render(llama_context * lctx, int offset, int n) {
+ GGML_ASSERT(offset >= 0 && n > 0 && offset + n <= n_tokens);
+ if (!batch) {
+ batch.reset(llama_batch_ext_init(lctx));
+ }
+ llama_batch_ext_clear(batch.get());
+ for (int i = offset; i < offset + n; i++) {
+ const llama_embd e = { embd + (size_t) i * n_mmproj_embd, 1, (size_t) n_mmproj_embd };
+ const int32_t idx = llama_batch_ext_add_embd(batch.get(), seq_id, e);
+ GGML_ASSERT(idx >= 0);
+
+ llama_pos p[GGML_MROPE_SECTIONS] = { 0, 0, 0, 0 };
+ for (int j = 0; j < n_pos_per_embd; j++) {
+ p[j] = pos[(size_t) j * (size_t) n_tokens + (size_t) i];
+ }
+ llama_batch_ext_set_pos(batch.get(), idx, p);
+
+ if (logits[i]) {
+ llama_batch_ext_set_output_logits(batch.get(), idx, true);
+ }
+ }
+ return batch.get();
+ }
};
diff --git a/tools/mtmd/mtmd-helper-gen.cpp b/tools/mtmd/mtmd-helper-gen.cpp
index 1c58d3ae1..5fb7ea9a0 100644
--- a/tools/mtmd/mtmd-helper-gen.cpp
+++ b/tools/mtmd/mtmd-helper-gen.cpp
@@ -222,14 +222,13 @@ public:
return 0;
}
const int32_t n_tokens_batch = std::min(n_batch, n_prompt - prompt_pos);
- llama_batch batch_view = prompt_batch->get_view(prompt_pos, n_tokens_batch);
const bool is_last_batch = (prompt_pos + n_tokens_batch) == n_prompt;
if (is_last_batch) {
- batch_view.logits[n_tokens_batch - 1] = 1;
+ prompt_batch->logits[prompt_pos + n_tokens_batch - 1] = 1;
}
- if (llama_decode(lctx, batch_view) != 0) {
+ if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, prompt_batch->render(lctx, prompt_pos, n_tokens_batch)) != 0) {
LOG_ERR("mtmd_helper_gen_audio: prompt decode failed\n");
return -1;
}
@@ -286,10 +285,10 @@ public:
decode_embd_batch batch_embd(fb.data(), 1, n_pos_per_embd, n_embd);
if (mrope) batch_embd.set_position_mrope_1d(pos, seq_id);
else batch_embd.set_position_normal (pos, seq_id);
- batch_embd.batch.logits[0] = 1;
+ batch_embd.logits[0] = 1;
pos++;
- if (llama_decode(lctx, batch_embd.batch) != 0) {
+ if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch_embd.render(lctx, 0, 1)) != 0) {
LOG_ERR("mtmd_helper_gen_audio: decode failed\n");
return 1;
}
@@ -586,13 +585,12 @@ public:
return 0;
}
const int32_t n_tokens_batch = std::min(n_batch, n_prompt - prompt_pos);
- llama_batch batch_view = prompt_batch->get_view(prompt_pos, n_tokens_batch);
if ((prompt_pos + n_tokens_batch) == n_prompt) {
- batch_view.logits[n_tokens_batch - 1] = 1;
+ prompt_batch->logits[prompt_pos + n_tokens_batch - 1] = 1;
}
- if (llama_decode(lctx, batch_view) != 0) {
+ if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, prompt_batch->render(lctx, prompt_pos, n_tokens_batch)) != 0) {
LOG_ERR("mtmd_helper_gen_audio: prompt decode failed\n");
return -1;
}
@@ -646,12 +644,12 @@ public:
}
}
- decode_embd_batch batch_embd(const_cast<float *>(out.embd), 1, 1, n_embd);
+ decode_embd_batch batch_embd(out.embd, 1, 1, n_embd);
batch_embd.set_position_normal(pos, seq_id);
- batch_embd.batch.logits[0] = 1;
+ batch_embd.logits[0] = 1;
pos++;
- if (llama_decode(lctx, batch_embd.batch) != 0) {
+ if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch_embd.render(lctx, 0, 1)) != 0) {
LOG_ERR("mtmd_helper_gen_audio: decode failed\n");
return 1;
}
@@ -842,8 +840,8 @@ private:
GGML_ASSERT(n_rows > 0);
decode_embd_batch batch(prompt_embd_buf.data(), n_rows, 1, n_e);
batch.set_position_normal(pos, seq_id);
- batch.batch.logits[n_rows - 1] = 1;
- if (llama_decode(lctx, batch.batch) != 0) {
+ batch.logits[n_rows - 1] = 1;
+ if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.render(lctx, 0, n_rows)) != 0) {
LOG_ERR("mtmd_helper_gen_audio: chunk prompt decode failed\n");
return 1;
}
diff --git a/tools/mtmd/mtmd-helper.cpp b/tools/mtmd/mtmd-helper.cpp
index bdf8bf6fe..cd05ade2d 100644
--- a/tools/mtmd/mtmd-helper.cpp
+++ b/tools/mtmd/mtmd-helper.cpp
@@ -169,19 +169,19 @@ int32_t mtmd_helper_decode_image_chunk(
while (i_batch < n_img_batches) { // split into batches
int pos_offset = i_batch*n_batch;
int n_tokens_batch = std::min(n_batch, n_tokens - pos_offset);
- llama_batch batch_embd_view = batch_embd.get_view(pos_offset, n_tokens_batch);
LOG_INF("decoding %s batch %d/%d, n_tokens_batch = %d\n", name, i_batch+1, n_img_batches, n_tokens_batch);
int64_t t1 = ggml_time_ms();
- int32_t ret = llama_decode(lctx, batch_embd_view);
+ int32_t ret = llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch_embd.render(lctx, pos_offset, n_tokens_batch));
if (ret != 0) {
LOG_ERR("failed to decode %s\n", name);
return ret;
}
if (callback != nullptr) {
- ret = callback(batch_embd_view, user_data);
+ const mtmd_helper_embd_batch view = batch_embd.get_view(pos_offset, n_tokens_batch);
+ ret = callback(&view, user_data);
if (ret != 0) {
LOG_ERR("post-decode callback failed\n");
return ret;
@@ -209,37 +209,35 @@ int32_t mtmd_helper_eval_chunk_single(mtmd_context * ctx,
llama_pos * new_n_past) {
GGML_ASSERT(n_batch > 0);
int32_t ret;
- llama_batch text_batch = llama_batch_init(n_batch, 0, 1);
auto chunk_type = mtmd_input_chunk_get_type(chunk);
if (chunk_type == MTMD_INPUT_CHUNK_TYPE_TEXT) {
size_t n_tokens;
const auto tokens = mtmd_input_chunk_get_tokens_text(chunk, &n_tokens);
// LOG_INF("decoding text chunk, n_tokens = %zu\n", n_tokens);
+ llama_batch_ext_ptr text_batch(llama_batch_ext_init(lctx));
size_t i = 0;
while (i < n_tokens) { // split into batches
- text_batch.n_tokens = 0; // clear the batch
- for (; i < n_tokens && text_batch.n_tokens < n_batch; i++) {
- int32_t j = text_batch.n_tokens;
- text_batch.token [j] = tokens[i];
- text_batch.pos [j] = n_past++;
- text_batch.n_seq_id[j] = 1;
- text_batch.seq_id [j][0] = seq_id;
- text_batch.logits [j] = false;
-
- text_batch.n_tokens++;
+ llama_batch_ext_clear(text_batch.get());
+ int32_t n_added = 0;
+ int32_t idx = -1;
+ for (; i < n_tokens && n_added < n_batch; i++) {
+ idx = llama_batch_ext_add_token(text_batch.get(), seq_id, tokens[i]);
+ GGML_ASSERT(idx >= 0);
+ llama_pos pos = n_past++;
+ llama_batch_ext_set_pos(text_batch.get(), idx, &pos);
+ n_added++;
}
bool is_last_token = (i == n_tokens);
if (logits_last && is_last_token) {
- text_batch.logits[text_batch.n_tokens - 1] = true;
+ llama_batch_ext_set_output_logits(text_batch.get(), idx, true);
}
- ret = llama_decode(lctx, text_batch);
+ ret = llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, text_batch.get());
if (ret != 0) {
LOG_ERR("failed to decode text\n");
- llama_batch_free(text_batch);
return ret;
}
- *new_n_past += text_batch.n_tokens;
+ *new_n_past += n_added;
}
} else if (chunk_type == MTMD_INPUT_CHUNK_TYPE_IMAGE || chunk_type == MTMD_INPUT_CHUNK_TYPE_AUDIO) {
@@ -251,7 +249,6 @@ int32_t mtmd_helper_eval_chunk_single(mtmd_context * ctx,
ret = mtmd_encode_chunk(ctx, chunk);
if (ret != 0) {
LOG_ERR("failed to encode %s slice\n", name);
- llama_batch_free(text_batch);
return ret;
}
@@ -261,14 +258,12 @@ int32_t mtmd_helper_eval_chunk_single(mtmd_context * ctx,
ret = mtmd_helper_decode_image_chunk(ctx, lctx, chunk, embd, n_past, seq_id, n_batch, new_n_past, nullptr, nullptr);
if (ret != 0) {
LOG_ERR("failed to decode %s\n", name);
- llama_batch_free(text_batch);
return ret;
}
} else {
GGML_ABORT("chunk type not supported");
}
- llama_batch_free(text_batch);
return 0;
}
diff --git a/tools/mtmd/mtmd-helper.h b/tools/mtmd/mtmd-helper.h
index 10f2171c0..7436230f0 100644
--- a/tools/mtmd/mtmd-helper.h
+++ b/tools/mtmd/mtmd-helper.h
@@ -92,9 +92,9 @@ MTMD_API llama_pos mtmd_helper_get_n_pos(const mtmd_input_chunks * chunks);
MTMD_API void mtmd_helper_image_get_decoder_pos(const mtmd_image_tokens * image, llama_pos pos_0, struct mtmd_decoder_pos * out_pos);
// helper function that automatically:
-// 1. run llama_decode() on text chunks
-// 2. run mtmd_encode_chunk() on image chunks, then mtmd_get_output_embd() and then llama_decode()
-// if any of the mtmd_encode_chunk() or llama_decode() calls return non-zero, stop and forward the error
+// 1. decode text chunks
+// 2. run mtmd_encode_chunk() on image chunks, then mtmd_get_output_embd() and then decode the embeddings
+// if any of the mtmd_encode_chunk() or decode calls return non-zero, stop and forward the error
// otherwise, returns 0 on success
// this function is NOT thread-safe
MTMD_API int32_t mtmd_helper_eval_chunks(mtmd_context * ctx,
@@ -117,7 +117,17 @@ MTMD_API int32_t mtmd_helper_eval_chunk_single(mtmd_context * ctx,
bool logits_last,
llama_pos * new_n_past);
-typedef int32_t (*mtmd_helper_post_decode_callback)(struct llama_batch batch, void * user_data);
+// one decoded sub-batch of embeddings, passed to mtmd_helper_post_decode_callback
+struct mtmd_helper_embd_batch {
+ int32_t n_tokens;
+ const float * embd; // [n_tokens, n_embd]
+ int32_t n_embd;
+ const llama_pos * pos; // [n_pos, n_tokens], section-major
+ int32_t n_pos; // 4 for M-RoPE models, 1 otherwise
+ llama_seq_id seq_id;
+};
+
+typedef int32_t (*mtmd_helper_post_decode_callback)(const struct mtmd_helper_embd_batch * batch, void * user_data);
// helper function to decode an image whose embeddings have already been calculated
// this helper will handle batching and pre/post decoding setup (for ex. gemma 3 requires non-causal attention)
diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp
index 611e82a6a..efff54997 100644
--- a/tools/server/server-context.cpp
+++ b/tools/server/server-context.cpp
@@ -109,8 +109,7 @@ enum slot_state {
struct server_slot; // forward declaration
struct server_batch {
- llama_batch batch;
- bool batch_rendered = false;
+ common_batch view; // the rendered sub-batch [off, off + n_tokens), see render()
struct token {
int32_t id_slot;
@@ -126,36 +125,21 @@ struct server_batch {
// track if given slot can be batched with slots already in the batch
server_slot * slot_batched = nullptr;
- // in embd mode, we temporarily swap out the tokens arr and restore it on clear()
bool has_embd = false;
- llama_token * tokens_ptr = nullptr;
std::vector<float> embd;
float alora_scale = -1.0f;
size_t alora_disabled_id = 0;
- server_batch() {
- batch.pos = nullptr; // sentinel: uninitialized batch
- }
-
- ~server_batch() {
- if (batch.pos != nullptr) {
- clear();
- llama_batch_free(batch);
- }
- }
-
- void init(int32_t n_tokens_alloc, int32_t n_embd) {
+ void init(llama_context * ctx, int32_t n_tokens_alloc, int32_t n_embd) {
this->n_tokens_alloc = n_tokens_alloc;
this->n_embd = n_embd;
- batch = llama_batch_init(n_tokens_alloc, 0, 1);
- tokens_ptr = batch.token;
+ view = common_batch(ctx);
tokens.reserve(n_tokens_alloc);
}
bool add(int32_t id_slot, llama_token token, llama_pos pos, bool output, bool is_prompt) {
GGML_ASSERT(!has_embd); // cannot mix tokens + embd in same batch
- GGML_ASSERT(batch.pos != nullptr);
if ((int32_t)tokens.size() >= n_tokens_alloc) {
return false;
}
@@ -164,7 +148,6 @@ struct server_batch {
}
bool add(int32_t id_slot, const std::vector<float> & embd_in, llama_pos pos, bool output, bool is_prompt) {
- GGML_ASSERT(batch.pos != nullptr);
if ((int32_t)tokens.size() >= n_tokens_alloc) {
return false;
}
@@ -177,16 +160,11 @@ struct server_batch {
void clear() {
tokens.clear();
embd.clear();
- common_batch_clear(batch);
+ view.clear();
slot_batched = nullptr;
alora_scale = -1.0f;
alora_disabled_id = 0;
- batch_rendered = false;
has_embd = false;
- if (batch.token == nullptr) {
- batch.token = tokens_ptr;
- batch.embd = nullptr;
- }
}
int32_t size() const {
@@ -198,41 +176,22 @@ struct server_batch {
tokens[idx].output = output;
}
- void render() {
- GGML_ASSERT(!batch_rendered);
- GGML_ASSERT(batch.pos != nullptr);
- common_batch_clear(batch);
- for (int32_t i = 0; i < size(); i++) {
- const auto & t = tokens[i];
- common_batch_add(batch, t.token, t.pos, { t.id_slot }, t.output);
- }
- if (has_embd) {
- batch.token = nullptr; // will be restored on clear()
- batch.embd = embd.data();
- }
- batch_rendered = true;
- }
-
- llama_batch get_view(int32_t off, int32_t n_tokens) const {
- GGML_ASSERT(batch.pos != nullptr);
- GGML_ASSERT(batch_rendered);
+ // render the sub-batch [off, off + n_tokens) into view, index i in view is index off + i here
+ void render(int32_t off, int32_t n_tokens) {
GGML_ASSERT(off >= 0 && off < size());
GGML_ASSERT(n_tokens > 0 && off + n_tokens <= size());
- auto * token = batch.token ? batch.token + off : nullptr;
- auto * embd = batch.embd ? batch.embd + off * n_embd : nullptr;
-
- llama_batch view = {
- n_tokens,
- token,
- embd,
- batch.pos + off,
- batch.n_seq_id + off,
- batch.seq_id + off,
- batch.logits + off,
- };
-
- return view;
+ view.clear();
+ for (int32_t i = off; i < off + n_tokens; i++) {
+ const auto & t = tokens[i];
+ if (has_embd) {
+ // text embeddings broadcast the same position across the M-RoPE sections
+ const llama_pos pos[GGML_MROPE_SECTIONS] = { t.pos, t.pos, t.pos, 0 };
+ view.add_embd({ embd.data() + (size_t) i * n_embd, 1, (size_t) n_embd }, pos, t.id_slot, t.output);
+ } else {
+ view.add(t.token, t.pos, t.id_slot, t.output);
+ }
+ }
}
};
@@ -761,13 +720,24 @@ static int process_mtmd_chunk(const server_slot & slot, mtmd::batch_ptr & mbatch
if (mbatch) {
float * embd = mtmd_batch_get_output_embd(mbatch.get(), chunk.get());
if (embd) {
- void * cb_data = slot.spec;
- static auto cb = [](llama_batch batch, void * user_data) {
- common_speculative * spec = static_cast<common_speculative *>(user_data);
- if (!common_speculative_process(spec, batch)) {
- return 1;
+ struct cb_data_t {
+ common_speculative * spec;
+ llama_context * ctx;
+ } cb_data = { slot.spec, slot.ctx_tgt };
+
+ static auto cb = [](const mtmd_helper_embd_batch * b, void * user_data) {
+ const auto * data = static_cast<cb_data_t *>(user_data);
+
+ common_batch batch(data->ctx);
+ for (int32_t i = 0; i < b->n_tokens; ++i) {
+ llama_pos pos[GGML_MROPE_SECTIONS] = { 0, 0, 0, 0 };
+ for (int32_t j = 0; j < b->n_pos; ++j) {
+ pos[j] = b->pos[j * b->n_tokens + i];
+ }
+ batch.add_embd({ b->embd + (size_t) i * b->n_embd, 1, (size_t) b->n_embd }, pos, b->seq_id, false);
}
- return 0;
+
+ return common_speculative_process(data->spec, batch) ? 0 : 1;
};
llama_pos new_n_past; // unused for now
@@ -781,7 +751,7 @@ static int process_mtmd_chunk(const server_slot & slot, mtmd::batch_ptr & mbatch
llama_n_batch(slot.ctx_tgt),
&new_n_past,
cb,
- cb_data
+ &cb_data
);
if (res != 0) {
SLT_ERR(slot, "failed to decode mtmd chunk, idx = %zu, res = %d\n", idx, res);
@@ -1356,7 +1326,7 @@ private:
{
const int32_t n_batch = llama_n_batch(ctx_tgt);
const int32_t n_embd = llama_model_n_embd_inp(model_tgt);
- batch.init(std::max(n_batch, params_base.n_parallel), n_embd);
+ batch.init(ctx_tgt, std::max(n_batch, params_base.n_parallel), n_embd);
}
if (params_base.cache_ram_mib != 0) {
@@ -2160,7 +2130,7 @@ private:
queue_results.send(std::move(res));
}
- void send_embedding(const server_slot & slot, const llama_batch & batch) {
+ void send_embedding(const server_slot & slot, const common_batch & batch) {
auto res = std::make_unique<server_task_result_embd>();
res->id = slot.task->id;
res->index = slot.task->index;
@@ -2171,8 +2141,8 @@ private:
std::vector<float> embd_res(n_embd_out, 0.0f);
- for (int i = 0; i < batch.n_tokens; ++i) {
- if (!batch.logits[i] || batch.seq_id[i][0] != slot.id) {
+ for (int i = 0; i < batch.size(); ++i) {
+ if (!batch.tokens[i].output || batch.tokens[i].seq_id != slot.id) {
continue;
}
@@ -2180,11 +2150,11 @@ private:
if (llama_pooling_type(slot.ctx_tgt) == LLAMA_POOLING_TYPE_NONE) {
embd = llama_get_embeddings_ith(slot.ctx_tgt, i);
} else {
- embd = llama_get_embeddings_seq(slot.ctx_tgt, batch.seq_id[i][0]);
+ embd = llama_get_embeddings_seq(slot.ctx_tgt, batch.tokens[i].seq_id);
}
if (embd == nullptr) {
- SLT_ERR(slot, "failed to get embeddings, token = %d, seq_id = %d\n", batch.token[i], batch.seq_id[i][0]);
+ SLT_ERR(slot, "failed to get embeddings, token = %d, seq_id = %d\n", batch.tokens[i].id, batch.tokens[i].seq_id);
res->embedding.push_back(std::vector<float>(n_embd_out, 0.0f));
continue;
@@ -2205,24 +2175,24 @@ private:
queue_results.send(std::move(res));
}
- void send_rerank(const server_slot & slot, const llama_batch & batch) {
+ void send_rerank(const server_slot & slot, const common_batch & batch) {
auto res = std::make_unique<server_task_result_rerank>();
res->id = slot.task->id;
res->index = slot.task->index;
res->n_tokens = slot.task->n_tokens();
- for (int i = 0; i < batch.n_tokens; ++i) {
- if (!batch.logits[i] || batch.seq_id[i][0] != slot.id) {
+ for (int i = 0; i < batch.size(); ++i) {
+ if (!batch.tokens[i].output || batch.tokens[i].seq_id != slot.id) {
continue;
}
- const float * embd = llama_get_embeddings_seq(ctx_tgt, batch.seq_id[i][0]);
+ const float * embd = llama_get_embeddings_seq(ctx_tgt, batch.tokens[i].seq_id);
if (embd == NULL) {
embd = llama_get_embeddings_ith(ctx_tgt, i);
}
if (embd == NULL) {
- SLT_ERR(slot, "failed to get embeddings, token = %d, seq_id = %d\n", batch.token[i], batch.seq_id[i][0]);
+ SLT_ERR(slot, "failed to get embeddings, token = %d, seq_id = %d\n", batch.tokens[i].id, batch.tokens[i].seq_id);
res->score = -1e6;
continue;
@@ -2845,7 +2815,6 @@ private:
try {
scoped_timer t(t_pre_decode, n_pre_decode);
pre_decode();
- batch.render();
} catch (const std::exception & e) {
SRV_ERR("pre_decode() failed: %s\n", e.what());
abort_all_slots("pre_decode() failed: " + std::string(e.what()));
@@ -2875,7 +2844,6 @@ private:
llama_set_embeddings(ctx_tgt, slot_batched->need_embd());
}
- llama_batch batch_view;
int32_t off_next = 0;
int32_t n_batch = llama_n_batch(ctx_tgt);
for (int32_t off = 0; off < batch.size(); off = off_next) {
@@ -2884,8 +2852,8 @@ private:
scoped_timer t(t_decode, n_decode);
// TODO @ngxson : maybe handle n_batch == 1 here instead of inside decode()
- batch_view = batch.get_view(off, n_tokens);
- bool ok = decode(n_batch, off, batch_view);
+ batch.render(off, n_tokens);
+ bool ok = decode(n_batch, off);
#ifdef DEBUG_TIMINGS
llama_synchronize(ctx_tgt);
#endif
@@ -2908,7 +2876,7 @@ private:
try {
scoped_timer t(t_post_decode, n_post_decode);
- post_decode(n_tokens, off, batch_view);
+ post_decode(n_tokens, off);
} catch (const std::exception & e) {
SRV_ERR("post_decode() failed: %s\n", e.what());
abort_all_slots("post_decode() failed: " + std::string(e.what()));
@@ -3655,7 +3623,7 @@ private:
// returns true = success ; false = retry with smaller batch size
// throw std::runtime_error on fatal error
- bool decode(int32_t & n_batch, int32_t off, llama_batch & batch_view) {
+ bool decode(int32_t & n_batch, int32_t off) {
SRV_DBG("n_batch (effective) = %d, off = %d\n", n_batch, off);
metrics_pre_decode();
@@ -3682,7 +3650,7 @@ private:
}
bool has_output = false;
- for (int i = off; i < off + batch_view.n_tokens; ++i) {
+ for (int i = off; i < off + batch.view.size(); ++i) {
has_output |= batch.tokens[i].output;
}
@@ -3690,7 +3658,7 @@ private:
// note: the sync is done here too, so that the wait is also covered by the yield
int ret = 0;
queue_tasks.yield_to_queue([&]() {
- ret = llama_decode(ctx_tgt, batch_view);
+ ret = llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, batch.view.get());
if (ret == 0 && has_output) {
llama_synchronize(ctx_tgt);
}
@@ -3746,7 +3714,7 @@ private:
return false; // retry with the updated n_batch
} else {
// success, apply batch metrics
- metrics_post_decode(off, batch_view.n_tokens, has_output);
+ metrics_post_decode(off, batch.view.size(), has_output);
}
// TODO: avoid restoring the draft context and re-evaluating the drafted tokens when not needed [TAG_SPEC_AVOID_DRAFT_REEVAL]
@@ -3755,7 +3723,7 @@ private:
if (spec) {
bool ok = true;
queue_tasks.yield_to_queue([&]() {
- ok = common_speculative_process(spec.get(), batch_view);
+ ok = common_speculative_process(spec.get(), batch.view);
});
if (!ok) {
@@ -3792,8 +3760,8 @@ private:
return true;
}
- void post_decode(int32_t n_batch_tokens, int32_t off, llama_batch & batch_view) {
- // for checking if a given batch index is inside batch_view
+ void post_decode(int32_t n_batch_tokens, int32_t off) {
+ // for checking if a given batch index is inside the current sub-batch
auto is_inside_view = [&](int32_t idx) {
return idx >= off && idx < off + n_batch_tokens;
};
@@ -3829,14 +3797,14 @@ private:
if (slot.state == SLOT_STATE_DONE_PROMPT) {
if (slot.task->type == SERVER_TASK_TYPE_EMBEDDING) {
// prompt evaluated for embedding
- send_embedding(slot, batch_view);
+ send_embedding(slot, batch.view);
slot.release();
slot.i_batch = -1;
return;
}
if (slot.task->type == SERVER_TASK_TYPE_RERANK) {
- send_rerank(slot, batch_view);
+ send_rerank(slot, batch.view);
slot.release();
slot.i_batch = -1;
return;