server: support inp embd to generate next token (#26313)

* server: support embd for sampled token

* fix ~server_batch()
This commit is contained in:
Xuan-Son Nguyen
2026-07-30 21:40:38 +02:00
committed by GitHub
parent ea63b4d32e
commit b4ca032ae3

View File

@@ -78,31 +78,41 @@ struct server_batch {
};
std::vector<token> tokens;
int32_t n_tokens_alloc = 0;
int32_t n_embd = 0;
// 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.token = nullptr; // sentinel: uninitialized batch
batch.pos = nullptr; // sentinel: uninitialized batch
}
~server_batch() {
if (batch.token != nullptr) {
if (batch.pos != nullptr) {
clear();
llama_batch_free(batch);
}
}
void init(int32_t n_tokens_alloc) {
void init(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;
tokens.reserve(n_tokens_alloc);
}
bool add(int32_t id_slot, llama_token token, llama_pos pos, bool output) {
GGML_ASSERT(batch.token != nullptr);
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;
}
@@ -110,13 +120,30 @@ struct server_batch {
return true;
}
bool add(int32_t id_slot, const std::vector<float> & embd_in, llama_pos pos, bool output) {
GGML_ASSERT(batch.pos != nullptr);
if ((int32_t)tokens.size() >= n_tokens_alloc) {
return false;
}
tokens.push_back({ id_slot, LLAMA_TOKEN_NULL, pos, output });
has_embd = true;
embd.insert(embd.end(), embd_in.begin(), embd_in.end());
return true;
}
void clear() {
tokens.clear();
embd.clear();
common_batch_clear(batch);
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 {
@@ -129,25 +156,33 @@ struct server_batch {
}
void render() {
GGML_ASSERT(batch.token != nullptr);
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.token != nullptr);
GGML_ASSERT(batch.pos != nullptr);
GGML_ASSERT(batch_rendered);
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,
batch.token + off,
nullptr,
token,
embd,
batch.pos + off,
batch.n_seq_id + off,
batch.seq_id + off,
@@ -270,6 +305,10 @@ struct server_slot {
llama_token sampled; // in speculative mode, this is the last accepted token
// for TTS models, this is the embd generated from prev step, decode this to generate next hidden state
// corresponding to one token position (size = n_embd)
std::vector<float> inp_embd;
// stats
size_t n_sent_text = 0; // number of sent text character
@@ -378,7 +417,9 @@ struct server_slot {
bool can_batch_with(server_slot & other_slot) const {
GGML_ASSERT(task);
return task->type == other_slot.task->type && are_lora_equal(lora, other_slot.lora);
return task->type == other_slot.task->type
&& inp_embd.size() == other_slot.inp_embd.size()
&& are_lora_equal(lora, other_slot.lora);
}
bool has_budget(const common_params & global_params) {
@@ -444,7 +485,11 @@ struct server_slot {
// no speculative decoding
i_batch = batch.size();
add_ok &= batch.add(id, sampled, prompt.tokens.pos_next(), true);
if (!inp_embd.empty()) {
add_ok &= batch.add(id, inp_embd, prompt.tokens.pos_next(), true);
} else {
add_ok &= batch.add(id, sampled, prompt.tokens.pos_next(), true);
}
SLT_DBG(*this, "slot decode token, id=%d, n_ctx = %d, n_tokens = %d, truncated = %d\n",
sampled, n_ctx, prompt.n_tokens(), truncated);
@@ -1334,7 +1379,8 @@ private:
// note that n_batch can be > n_ctx (e.g. for non-causal attention models such as BERT where the KV cache is not used)
{
const int32_t n_batch = llama_n_batch(ctx_tgt);
batch.init(std::max(n_batch, params_base.n_parallel));
const int32_t n_embd = llama_model_n_embd_inp(model_tgt);
batch.init(std::max(n_batch, params_base.n_parallel), n_embd);
}
if (params_base.cache_ram_mib != 0) {
@@ -3578,6 +3624,15 @@ private:
n_empty_consecutive = 0;
}
// TODO @ngxson : dft model may have different n_embd than the tgt model, so we check & reject if that's the case
// this case is not currently used by any models, but may need to be supported in the future
if (spec && batch.has_embd) {
if (llama_model_n_embd_inp(model_dft) != llama_model_n_embd_inp(model_tgt)) {
SRV_ERR("%s", "unsupported batch.has_embd + spec case\n");
throw std::runtime_error("unsupported batch.has_embd + spec case");
}
}
const int ret = llama_decode(ctx_tgt, batch_view);
metrics.on_decoded(slots);