mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-29 01:17:36 -05:00
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
This commit is contained in:
+137
-51
@@ -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,29 +2132,138 @@ 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;
|
||||
}
|
||||
|
||||
auto mem = llama_get_memory(ctx);
|
||||
llama_pos pos = mem ? llama_memory_seq_pos_max(mem, 0) + 1 : 0;
|
||||
void common_batch::clear() {
|
||||
tokens.clear();
|
||||
llama_batch_ext_clear(batch.get());
|
||||
}
|
||||
|
||||
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);
|
||||
pos++;
|
||||
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;
|
||||
}
|
||||
|
||||
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;
|
||||
|
||||
+47
-2
@@ -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
|
||||
//
|
||||
|
||||
+183
-186
@@ -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);
|
||||
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<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) {
|
||||
|
||||
@@ -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`
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
};
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
+16
-21
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user