From f1ea206218210afb913ae2f5d2c51faed35915da Mon Sep 17 00:00:00 2001 From: Xuan-Son Nguyen Date: Mon, 28 Sep 2026 19:52:45 +0200 Subject: [PATCH] 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 --- common/common.cpp | 188 ++++++--- common/common.h | 49 ++- common/speculative.cpp | 369 +++++++++--------- common/speculative.h | 3 + .../speculative-simple/speculative-simple.cpp | 1 - tools/mtmd/mtmd-cli.cpp | 11 +- tools/mtmd/mtmd-helper-common.h | 152 ++++---- tools/mtmd/mtmd-helper-gen.cpp | 24 +- tools/mtmd/mtmd-helper.cpp | 37 +- tools/mtmd/mtmd-helper.h | 18 +- tools/server/server-context.cpp | 148 +++---- 11 files changed, 544 insertions(+), 456 deletions(-) diff --git a/common/common.cpp b/common/common.cpp index 6099f2ecc4..6f443f1bf0 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= (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 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; } - if (!tokens.empty()) { - llama_batch_ext_set_output_logits(batch.get(), (int32_t) tokens.size() - 1, true); + 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 = 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; + batch.add(tokens[i], pos, 0, output); + pos++; } 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 3294694109..45b15def71 100644 --- a/common/common.h +++ b/common/common.h @@ -8,6 +8,7 @@ #include "ggml.h" #include "llama.h" +#include #include #include #include @@ -879,7 +880,6 @@ void string_process_escapes(std::string & input); std::string string_from(bool value); std::string string_from(const std::vector & values); std::string string_from(const struct llama_context * ctx, const std::vector & 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 & 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 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 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 6fdfa4dc33..82e9e92238 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 zeros; + bool zeros_warned = false; // the substitution is reported once std::vector 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 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 i_batch_beg(n_seq, -1); std::vector 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 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 features_buf; // [n_chunk, n_embd_enc] gathered target features std::vector 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 i_batch_beg(n_seq, -1); std::vector 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 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); + batch.clear(); - 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); - } - - // 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); - set_h(i_batch_beg[seq_id], pending_h[seq_id].data()); + const float * h_row = k == i_batch_beg[seq_id] + ? pending_h[seq_id].data() + : h_tgt + (size_t) (k - 1) * n_embd; + + batch.set_embd(idx, { h_row, 1, (size_t) n_embd }); } 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 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> 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(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 c968750e2d..211fcdabd1 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 863af5a2c7..81aa106f14 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 ba18b3e32b..6fe058fd8c 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 f907346c7b..bc68ed4b9c 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 @@ -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 pos; - std::vector pos_view; // used by mrope - std::vector n_seq_id; - std::vector seq_id_0; - std::vector seq_ids; - std::vector 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 pos; // [n_pos_per_embd, n_tokens], section-major + std::vector pos_view; // sliced positions of the last get_view() + std::vector 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 & 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 1c58d3ae19..5fb7ea9a0d 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(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 bdf8bf6fe4..cd05ade2d7 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 10f2171c0f..7436230f0d 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 611e82a6af..efff549971 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 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 & 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(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(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(); res->id = slot.task->id; res->index = slot.task->index; @@ -2171,8 +2141,8 @@ private: std::vector 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(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(); 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;