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:
Xuan-Son Nguyen
2026-09-28 19:52:45 +02:00
committed by GitHub
parent 6c7a87f7e5
commit f1ea206218
11 changed files with 544 additions and 456 deletions
+58 -90
View File
@@ -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;