};
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;
}
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 {
}
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,
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
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) {
// 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);
// 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) {
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);