mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-27 21:46:57 +02:00
llama: add llama_batch_ext (#24669)
* (wip) add llama_batch_ext * wip * updated design * updated impl * change signature * unused var * demo common_prompt_batch_decode * fix pos * tmp disable test-batch-alloc * fix compat * nits: add const * no more pos_max * add comment about llama_batch_ext_set_embd_state * handle n_embd_out properly * rename api --> embd_token * llama_embd * stub llama_batch_ext_set_embd_state * support both token + embd + state in batch * llama_batch_ext_add_embd * upstream some changes * nits * fix test-batch-alloc * add test for compat
This commit is contained in:
+440
-103
@@ -3,6 +3,9 @@
|
||||
#include "llama-impl.h"
|
||||
#include "llama-vocab.h"
|
||||
#include "llama-memory.h"
|
||||
#include "llama-hparams.h"
|
||||
#include "llama-model.h"
|
||||
#include "llama-context.h"
|
||||
|
||||
#include <cassert>
|
||||
#include <cstring>
|
||||
@@ -23,43 +26,120 @@ llama_batch_allocr::llama_batch_allocr(uint32_t n_pos_per_embd) : n_pos_per_embd
|
||||
}
|
||||
|
||||
bool llama_batch_allocr::init(
|
||||
const llama_batch & batch_inp,
|
||||
const llama_batch_ext & batch_inp,
|
||||
const llama_vocab & vocab,
|
||||
const llama_memory_i * memory,
|
||||
uint32_t n_embd,
|
||||
uint32_t n_seq_max,
|
||||
bool output_all) {
|
||||
clear();
|
||||
|
||||
batch = batch_inp;
|
||||
this->vocab = &vocab;
|
||||
this->n_embd = batch_inp.n_embd > 0 ? batch_inp.n_embd : batch_inp.n_embd_inp;
|
||||
this->n_seq_max = batch_inp.n_seq_max;
|
||||
|
||||
this->vocab = &vocab;
|
||||
const int32_t n_tok = (int32_t) batch_inp.tokens.size();
|
||||
|
||||
GGML_ASSERT(batch.n_tokens > 0);
|
||||
GGML_ASSERT(n_tok > 0);
|
||||
|
||||
//
|
||||
// validate input batch
|
||||
//
|
||||
|
||||
if (n_seq_max > LLAMA_MAX_SEQ) {
|
||||
if ((uint32_t) n_seq_max > LLAMA_MAX_SEQ) {
|
||||
LLAMA_LOG_ERROR("%s: n_seq_max = %d > %d\n", __func__, n_seq_max, LLAMA_MAX_SEQ);
|
||||
return false;
|
||||
}
|
||||
|
||||
if (batch.token) {
|
||||
for (int32_t i = 0; i < batch.n_tokens; ++i) {
|
||||
if (batch.token[i] < 0 || (uint32_t) batch.token[i] >= vocab.n_tokens()) {
|
||||
LLAMA_LOG_ERROR("%s: invalid token[%d] = %d\n", __func__, i, batch.token[i]);
|
||||
const llama_memory_i * mem = batch_inp.mem;
|
||||
|
||||
//
|
||||
// determine the content types of the batch
|
||||
// an entry can carry a token id, a token embedding, or both (e.g. MTP hook batches)
|
||||
// all entries must carry the same combination
|
||||
//
|
||||
|
||||
const bool has_token = batch_inp.tokens[0].id != LLAMA_TOKEN_NULL;
|
||||
const bool has_embd = batch_inp.tokens[0].has_embd;
|
||||
|
||||
for (int32_t i = 1; i < n_tok; ++i) {
|
||||
if ((batch_inp.tokens[i].id != LLAMA_TOKEN_NULL) != has_token ||
|
||||
batch_inp.tokens[i].has_embd != has_embd) {
|
||||
LLAMA_LOG_ERROR("%s: all entries in the batch must have the same content types\n", __func__);
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
if (!has_token && !has_embd) {
|
||||
LLAMA_LOG_ERROR("%s: batch has neither token ids nor embeddings\n", __func__);
|
||||
return false;
|
||||
}
|
||||
|
||||
//
|
||||
// build flat token/embd array
|
||||
//
|
||||
|
||||
if (has_token) {
|
||||
token_vec.resize(n_tok);
|
||||
for (int32_t i = 0; i < n_tok; ++i) {
|
||||
const llama_token id = batch_inp.tokens[i].id;
|
||||
if (id < 0 || id >= batch_inp.n_vocab) {
|
||||
LLAMA_LOG_ERROR("%s: invalid token[%d] = %d\n", __func__, i, id);
|
||||
return false;
|
||||
}
|
||||
token_vec[i] = id;
|
||||
}
|
||||
}
|
||||
|
||||
if (has_embd) {
|
||||
embd_vec = batch_inp.embd;
|
||||
}
|
||||
|
||||
//
|
||||
// build flat pos array
|
||||
// token batch: pos[i] = tokens[i].pos[0]
|
||||
// embedding batch: pos[j*n_tok + i] = tokens[i].pos[j] (section-major)
|
||||
//
|
||||
|
||||
{
|
||||
const int32_t n_pos_total = has_token ? n_tok : n_tok * (int32_t) n_pos_per_embd;
|
||||
pos.resize(n_pos_total);
|
||||
if (has_token) {
|
||||
for (int32_t i = 0; i < n_tok; ++i) {
|
||||
pos[i] = batch_inp.tokens[i].pos[0];
|
||||
}
|
||||
} else {
|
||||
for (int32_t i = 0; i < n_tok; ++i) {
|
||||
for (uint32_t j = 0; j < n_pos_per_embd; ++j) {
|
||||
pos[(int32_t) j * n_tok + i] = batch_inp.tokens[i].pos[j];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (batch.seq_id) {
|
||||
for (int32_t i = 0; i < batch.n_tokens; ++i) {
|
||||
for (int32_t s = 0; s < batch.n_seq_id[i]; ++s) {
|
||||
if (batch.seq_id && (batch.seq_id[i][s] < 0 || batch.seq_id[i][s] >= (llama_seq_id) n_seq_max)) {
|
||||
LLAMA_LOG_ERROR("%s: invalid seq_id[%d][%d] = %d >= %d\n", __func__, i, s, batch.seq_id[i][s], (llama_seq_id) n_seq_max);
|
||||
//
|
||||
// build n_seq_id / seq_id arrays
|
||||
//
|
||||
|
||||
n_seq_id.resize(n_tok);
|
||||
seq_id.resize(n_tok + 1);
|
||||
seq_id[n_tok] = nullptr;
|
||||
|
||||
{
|
||||
size_t total = 0;
|
||||
for (int32_t i = 0; i < n_tok; ++i) {
|
||||
total += batch_inp.tokens[i].seq_ids.size();
|
||||
}
|
||||
seq_id_data.reserve(total);
|
||||
|
||||
for (int32_t i = 0; i < n_tok; ++i) {
|
||||
for (auto sid : batch_inp.tokens[i].seq_ids) {
|
||||
seq_id_data.push_back(sid);
|
||||
}
|
||||
}
|
||||
|
||||
size_t off = 0;
|
||||
for (int32_t i = 0; i < n_tok; ++i) {
|
||||
n_seq_id[i] = (int32_t) batch_inp.tokens[i].seq_ids.size();
|
||||
seq_id[i] = seq_id_data.data() + off;
|
||||
off += n_seq_id[i];
|
||||
|
||||
for (int32_t s = 0; s < n_seq_id[i]; ++s) {
|
||||
if (seq_id[i][s] < 0 || seq_id[i][s] >= (llama_seq_id) n_seq_max) {
|
||||
LLAMA_LOG_ERROR("%s: invalid seq_id[%d][%d] = %d >= %d\n", __func__, i, s, seq_id[i][s], (llama_seq_id) n_seq_max);
|
||||
return false;
|
||||
}
|
||||
}
|
||||
@@ -67,91 +147,43 @@ bool llama_batch_allocr::init(
|
||||
}
|
||||
|
||||
//
|
||||
// auto-generate missing fields
|
||||
// build output/logits array
|
||||
//
|
||||
|
||||
if (!batch.n_seq_id) {
|
||||
n_seq_id.resize(batch.n_tokens);
|
||||
for (int32_t i = 0; i < batch.n_tokens; i++) {
|
||||
n_seq_id[i] = seq_id_0.size();
|
||||
}
|
||||
batch.n_seq_id = n_seq_id.data();
|
||||
}
|
||||
|
||||
if (!batch.seq_id) {
|
||||
seq_id.resize(batch.n_tokens + 1);
|
||||
seq_id[batch.n_tokens] = NULL;
|
||||
for (int32_t i = 0; i < batch.n_tokens; i++) {
|
||||
seq_id[i] = seq_id_0.data();
|
||||
}
|
||||
batch.seq_id = seq_id.data();
|
||||
}
|
||||
|
||||
if (!batch.pos) {
|
||||
pos.resize(batch.n_tokens);
|
||||
|
||||
// initialize the starting position for each sequence based on the positions in the memory
|
||||
llama_pos p0[LLAMA_MAX_SEQ];
|
||||
for (uint32_t s = 0; s < n_seq_max; ++s) {
|
||||
if (!memory) {
|
||||
// if no memory -> start from 0
|
||||
p0[s] = 0;
|
||||
} else {
|
||||
p0[s] = memory->seq_pos_max(s) + 1;
|
||||
}
|
||||
{
|
||||
output.resize(n_tok, 0);
|
||||
for (int32_t i = 0; i < n_tok; ++i) {
|
||||
output[i] = batch_inp.tokens[i].output ? 1 : 0;
|
||||
}
|
||||
|
||||
for (int32_t i = 0; i < batch.n_tokens; i++) {
|
||||
const llama_seq_id seq_id = batch.seq_id[i][0];
|
||||
|
||||
pos[i] = p0[seq_id];
|
||||
|
||||
// update the starting position for all sequences that are assigned to the this token
|
||||
for (int32_t s = 0; s < batch.n_seq_id[i]; ++s) {
|
||||
const llama_seq_id seq_id = batch.seq_id[i][s];
|
||||
|
||||
p0[seq_id] = pos[i] + 1;
|
||||
}
|
||||
}
|
||||
|
||||
batch.pos = pos.data();
|
||||
}
|
||||
|
||||
if (!batch.logits) {
|
||||
if (output_all) {
|
||||
// return the output for all tokens
|
||||
output.resize(batch.n_tokens, true);
|
||||
} else {
|
||||
// return the output only for the last token
|
||||
output.resize(batch.n_tokens, false);
|
||||
output[output.size() - 1] = true;
|
||||
}
|
||||
|
||||
batch.logits = output.data();
|
||||
} else if (output_all) {
|
||||
bool warn = false;
|
||||
|
||||
for (int32_t i = 0; i < batch.n_tokens; ++i) {
|
||||
if (batch.logits[i] == 0) {
|
||||
warn = true;
|
||||
bool warn = false;
|
||||
for (int32_t i = 0; i < n_tok; ++i) {
|
||||
if (!output[i]) { warn = true; break; }
|
||||
}
|
||||
if (warn) {
|
||||
LLAMA_LOG_WARN("%s: embeddings required but some input tokens were not marked as outputs -> overriding\n", __func__);
|
||||
std::fill(output.begin(), output.end(), 1);
|
||||
}
|
||||
}
|
||||
|
||||
if (warn) {
|
||||
LLAMA_LOG_WARN("%s: embeddings required but some input tokens were not marked as outputs -> overriding\n", __func__);
|
||||
|
||||
output.resize(batch.n_tokens, true);
|
||||
batch.logits = output.data();
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
// set up the internal llama_batch to point to our owned arrays
|
||||
//
|
||||
|
||||
batch.n_tokens = n_tok;
|
||||
batch.token = has_token ? token_vec.data() : nullptr;
|
||||
batch.embd = has_embd ? embd_vec.data() : nullptr;
|
||||
batch.pos = pos.data();
|
||||
batch.n_seq_id = n_seq_id.data();
|
||||
batch.seq_id = seq_id.data();
|
||||
batch.logits = output.data();
|
||||
|
||||
//
|
||||
// compute stats
|
||||
//
|
||||
|
||||
this->n_embd = n_embd;
|
||||
this->n_seq_max = n_seq_max;
|
||||
|
||||
// count the outputs in this batch
|
||||
for (int32_t i = 0; i < batch.n_tokens; ++i) {
|
||||
n_outputs += batch.logits[i] != 0;
|
||||
@@ -259,7 +291,7 @@ bool llama_batch_allocr::init(
|
||||
continue;
|
||||
}
|
||||
|
||||
const llama_pos p0 = memory ? memory->seq_pos_max(s) : -1;
|
||||
const llama_pos p0 = mem ? mem->seq_pos_max(s) : -1;
|
||||
|
||||
if (batch.token) {
|
||||
if (p0 >= 0 && p0 >= seq_pos_min(s)) {
|
||||
@@ -292,7 +324,7 @@ bool llama_batch_allocr::init(
|
||||
continue;
|
||||
}
|
||||
|
||||
const llama_pos p0 = memory ? memory->seq_pos_max(s) : -1;
|
||||
const llama_pos p0 = mem ? mem->seq_pos_max(s) : -1;
|
||||
|
||||
if (p0 >= 0) {
|
||||
bool ok = true;
|
||||
@@ -320,12 +352,12 @@ bool llama_batch_allocr::init(
|
||||
}
|
||||
}
|
||||
|
||||
if (memory) {
|
||||
if (mem) {
|
||||
for (uint32_t s0 = 0; s0 < n_seq_max; ++s0) {
|
||||
for (uint32_t s1 = 0; s1 < n_seq_max; ++s1) {
|
||||
if (seq_cpl[s0][s1]) {
|
||||
if (memory->seq_pos_min(s0) != memory->seq_pos_min(s1) ||
|
||||
memory->seq_pos_max(s0) != memory->seq_pos_max(s1)) {
|
||||
if (mem->seq_pos_min(s0) != mem->seq_pos_min(s1) ||
|
||||
mem->seq_pos_max(s0) != mem->seq_pos_max(s1)) {
|
||||
LLAMA_LOG_ERROR("%s: sequence %d is coupled to %d in the input batch, but have divereged\n", __func__, s0, s1);
|
||||
return false;
|
||||
}
|
||||
@@ -725,11 +757,14 @@ void llama_batch_allocr::clear() {
|
||||
|
||||
batch = {};
|
||||
|
||||
pos .clear();
|
||||
n_seq_id .clear();
|
||||
seq_id .clear();
|
||||
seq_id_unq.clear();
|
||||
output .clear();
|
||||
token_vec .clear();
|
||||
embd_vec .clear();
|
||||
seq_id_data .clear();
|
||||
pos .clear();
|
||||
n_seq_id .clear();
|
||||
seq_id .clear();
|
||||
seq_id_unq .clear();
|
||||
output .clear();
|
||||
|
||||
for (auto & cur : seq_pos) {
|
||||
cur.clear();
|
||||
@@ -985,3 +1020,305 @@ void llama_batch_free(struct llama_batch batch) {
|
||||
}
|
||||
if (batch.logits) free(batch.logits);
|
||||
}
|
||||
|
||||
|
||||
// llama_batch_ext
|
||||
|
||||
size_t llama_batch_ext_select_n_embd_inp(llama_context_type ctx_type, llm_arch arch, const llama_hparams & hparams) {
|
||||
if (ctx_type == LLAMA_CONTEXT_TYPE_MTP) {
|
||||
return hparams.n_embd_out();
|
||||
}
|
||||
if (arch == LLM_ARCH_DFLASH) {
|
||||
return hparams.n_embd_inp_enc();
|
||||
}
|
||||
return hparams.n_embd_inp();
|
||||
}
|
||||
|
||||
llama_batch_ext::llama_batch_ext(llama_context * ctx) :
|
||||
n_tokens_max(llama_n_batch(ctx)),
|
||||
n_embd_inp(llama_batch_ext_select_n_embd_inp(ctx->get_cparams().ctx_type, llama_get_model(ctx)->arch, llama_get_model(ctx)->hparams)),
|
||||
n_embd_inp_enc(llama_get_model(ctx)->hparams.n_embd_inp_enc()),
|
||||
n_seq_max(llama_n_seq_max(ctx)),
|
||||
mem(llama_get_memory(ctx)),
|
||||
n_vocab(llama_vocab_n_tokens(llama_model_get_vocab(llama_get_model(ctx)))),
|
||||
n_pos_per_embd(llama_get_model(ctx)->hparams.n_pos_per_embd()) {
|
||||
clear();
|
||||
}
|
||||
|
||||
llama_batch_ext::llama_batch_ext(
|
||||
size_t n_tokens_max,
|
||||
size_t n_embd_inp,
|
||||
size_t n_embd_inp_enc,
|
||||
llama_seq_id n_seq_max,
|
||||
llama_memory_i * mem,
|
||||
llama_token n_vocab,
|
||||
size_t n_pos_per_embd) :
|
||||
n_tokens_max(n_tokens_max),
|
||||
n_embd_inp(n_embd_inp),
|
||||
n_embd_inp_enc(n_embd_inp_enc),
|
||||
n_seq_max(n_seq_max),
|
||||
mem(mem),
|
||||
n_vocab(n_vocab),
|
||||
n_pos_per_embd(n_pos_per_embd) {
|
||||
clear();
|
||||
}
|
||||
|
||||
void llama_batch_ext::clear() {
|
||||
tokens.clear();
|
||||
embd .clear();
|
||||
n_embd = 0;
|
||||
}
|
||||
|
||||
int32_t llama_batch_ext::add_token(llama_seq_id seq_id) {
|
||||
if (tokens.size() >= n_tokens_max) {
|
||||
return -1; // size limit reached
|
||||
}
|
||||
if (seq_id < 0 || seq_id >= n_seq_max) {
|
||||
return -3; // invalid sequence id
|
||||
}
|
||||
|
||||
// position is left undefined; call set_token_pos() before decoding
|
||||
token t;
|
||||
t.seq_ids.insert(seq_id);
|
||||
|
||||
tokens.push_back(t);
|
||||
|
||||
return (int32_t)(tokens.size() - 1);
|
||||
}
|
||||
|
||||
bool llama_batch_ext::add_seq(int32_t idx, llama_seq_id seq_id) {
|
||||
if (idx < 0 || idx >= (int32_t) tokens.size()) {
|
||||
return false;
|
||||
}
|
||||
if (seq_id < 0 || seq_id >= n_seq_max) {
|
||||
return false;
|
||||
}
|
||||
|
||||
token & t = tokens[idx];
|
||||
|
||||
t.seq_ids.insert(seq_id);
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool llama_batch_ext::set_token_id(int32_t idx, llama_token id) {
|
||||
if (idx < 0 || idx >= (int32_t) tokens.size()) {
|
||||
return false;
|
||||
}
|
||||
if (id < 0 || id >= n_vocab) {
|
||||
return false;
|
||||
}
|
||||
tokens[idx].id = id;
|
||||
return true;
|
||||
}
|
||||
|
||||
bool llama_batch_ext::set_token_embd(int32_t idx, llama_embd embd_in) {
|
||||
if (idx < 0 || idx >= (int32_t) tokens.size()) {
|
||||
return false;
|
||||
}
|
||||
if (!embd_in.data) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const size_t n_total = embd_in.n_rows * embd_in.n_embd;
|
||||
if (n_embd == 0) {
|
||||
if (n_total != n_embd_inp && n_total != n_embd_inp_enc) {
|
||||
LLAMA_LOG_ERROR("%s: embedding size mismatch, got %zu rows x %zu = %zu, expected %zu or %zu\n",
|
||||
__func__, embd_in.n_rows, embd_in.n_embd, n_total, n_embd_inp, n_embd_inp_enc);
|
||||
return false;
|
||||
}
|
||||
n_embd = n_total;
|
||||
} else if (n_total != n_embd) {
|
||||
LLAMA_LOG_ERROR("%s: embedding size mismatch, got %zu rows x %zu = %zu, expected %zu\n",
|
||||
__func__, embd_in.n_rows, embd_in.n_embd, n_total, n_embd);
|
||||
return false;
|
||||
}
|
||||
|
||||
token & t = tokens[idx];
|
||||
|
||||
if (t.has_embd) {
|
||||
LLAMA_LOG_ERROR("%s: embedding for token %d is already set\n", __func__, idx);
|
||||
return false;
|
||||
}
|
||||
|
||||
t.has_embd = true;
|
||||
t.embd_off = embd.size();
|
||||
embd.insert(embd.end(), embd_in.data, embd_in.data + n_total);
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool llama_batch_ext::set_token_pos(int32_t idx, const llama_pos * pos_in) {
|
||||
if (idx < 0 || idx >= (int32_t) tokens.size()) {
|
||||
return false;
|
||||
}
|
||||
if (!pos_in) {
|
||||
return false;
|
||||
}
|
||||
|
||||
token & t = tokens[idx];
|
||||
|
||||
size_t n_pos = t.id != LLAMA_TOKEN_NULL ? 1 : n_pos_per_embd;
|
||||
for (size_t i = 0; i < n_pos; ++i) {
|
||||
t.pos[i] = pos_in[i];
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool llama_batch_ext::set_output(int32_t idx, bool output_last) {
|
||||
if (idx < 0 || idx >= (int32_t) tokens.size()) {
|
||||
return false;
|
||||
}
|
||||
tokens[idx].output = output_last;
|
||||
return true;
|
||||
}
|
||||
|
||||
// llama_batch_ext C API
|
||||
|
||||
llama_batch_ext * llama_batch_ext_init(llama_context * ctx) {
|
||||
return new llama_batch_ext(ctx);
|
||||
}
|
||||
|
||||
void llama_batch_ext_free(llama_batch_ext * batch) {
|
||||
delete batch;
|
||||
}
|
||||
|
||||
void llama_batch_ext_clear(llama_batch_ext * batch) {
|
||||
batch->clear();
|
||||
}
|
||||
|
||||
int32_t llama_batch_ext_add(llama_batch_ext * batch, llama_seq_id seq_id) {
|
||||
return batch->add_token(seq_id);
|
||||
}
|
||||
|
||||
int32_t llama_batch_ext_add_token(llama_batch_ext * batch, llama_seq_id seq_id, llama_token id) {
|
||||
int32_t idx = batch->add_token(seq_id);
|
||||
if (idx < 0) {
|
||||
return idx;
|
||||
}
|
||||
if (!batch->set_token_id(idx, id)) {
|
||||
return -2;
|
||||
}
|
||||
return idx;
|
||||
}
|
||||
|
||||
int32_t llama_batch_ext_add_embd(llama_batch_ext * batch, llama_seq_id seq_id, llama_embd embd) {
|
||||
int32_t idx = batch->add_token(seq_id);
|
||||
if (idx < 0) {
|
||||
return idx;
|
||||
}
|
||||
if (!batch->set_token_embd(idx, embd)) {
|
||||
return -2;
|
||||
}
|
||||
return idx;
|
||||
}
|
||||
|
||||
bool llama_batch_ext_add_seq(llama_batch_ext * batch, int32_t idx, llama_seq_id seq_id) {
|
||||
return batch->add_seq(idx, seq_id);
|
||||
}
|
||||
|
||||
bool llama_batch_ext_set_pos(llama_batch_ext * batch, int32_t idx, const llama_pos * pos) {
|
||||
return batch->set_token_pos(idx, pos);
|
||||
}
|
||||
|
||||
bool llama_batch_ext_set_embd_token(llama_batch_ext * batch, int32_t idx, llama_embd embd) {
|
||||
return batch->set_token_embd(idx, embd);
|
||||
}
|
||||
|
||||
bool llama_batch_ext_set_embd_state(llama_batch_ext * batch, int32_t idx, llama_embd embd) {
|
||||
// TODO
|
||||
GGML_UNUSED(batch);
|
||||
GGML_UNUSED(idx);
|
||||
GGML_UNUSED(embd);
|
||||
return false;
|
||||
}
|
||||
|
||||
bool llama_batch_ext_set_output_embd(llama_batch_ext * batch, int32_t idx, bool value) {
|
||||
return batch->set_output(idx, value);
|
||||
}
|
||||
|
||||
bool llama_batch_ext_set_output_logits(llama_batch_ext * batch, int32_t idx, bool value) {
|
||||
return batch->set_output(idx, value);
|
||||
}
|
||||
|
||||
// llama_batch_compat
|
||||
|
||||
void llama_batch_compat::init(llama_batch_ext & dst, const llama_batch & batch_inp, size_t n_embd_row) {
|
||||
llama_batch_ext * batch_ext = &dst;
|
||||
|
||||
if (n_embd_row == 0) {
|
||||
n_embd_row = batch_ext->n_embd_inp;
|
||||
}
|
||||
|
||||
// a batch can carry both, for example the MTP hook batches
|
||||
const bool has_token = batch_inp.token != nullptr;
|
||||
const bool has_embd = batch_inp.embd != nullptr;
|
||||
|
||||
static const llama_seq_id default_seq_id = 0;
|
||||
static const int32_t default_n_seq_id = 1;
|
||||
|
||||
// auto-generates positions locally when batch_inp.pos is null, continuing from memory
|
||||
std::vector<llama_pos> pos_next(batch_ext->n_seq_max);
|
||||
for (llama_seq_id s = 0; s < (llama_seq_id) batch_ext->n_seq_max; ++s) {
|
||||
pos_next[s] = llama_memory_seq_pos_max(batch_ext->mem, s) + 1; // assume next pos
|
||||
}
|
||||
|
||||
for (int32_t i = 0; i < batch_inp.n_tokens; ++i) {
|
||||
const int32_t n_sid = batch_inp.n_seq_id ? batch_inp.n_seq_id[i] : default_n_seq_id;
|
||||
const llama_seq_id * sids = batch_inp.seq_id ? batch_inp.seq_id[i] : &default_seq_id;
|
||||
|
||||
llama_batch_ext::token t;
|
||||
|
||||
// seq_ids
|
||||
for (int32_t s = 0; s < n_sid; ++s) {
|
||||
t.seq_ids.insert(sids[s]);
|
||||
}
|
||||
|
||||
// position(s)
|
||||
if (batch_inp.pos) {
|
||||
if (has_token) {
|
||||
// token batch: one position per token
|
||||
t.pos[0] = batch_inp.pos[i];
|
||||
} else {
|
||||
// embedding batch (M-RoPE): section-major layout pos[j*n_tokens + i]
|
||||
for (uint32_t j = 0; j < batch_ext->n_pos_per_embd; ++j) {
|
||||
t.pos[j] = batch_inp.pos[(int32_t) j * batch_inp.n_tokens + i];
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// auto-generate position from the first seq_id
|
||||
t.pos[0] = pos_next[sids[0]]++;
|
||||
}
|
||||
|
||||
// token id and/or embeddings
|
||||
if (has_token) {
|
||||
t.id = batch_inp.token[i];
|
||||
}
|
||||
|
||||
if (has_embd) {
|
||||
t.has_embd = true;
|
||||
t.embd_off = batch_ext->embd.size();
|
||||
const float * src = batch_inp.embd + (size_t) i * n_embd_row;
|
||||
batch_ext->embd.insert(batch_ext->embd.end(), src, src + n_embd_row);
|
||||
batch_ext->n_embd = n_embd_row;
|
||||
}
|
||||
|
||||
// output flag
|
||||
// if no logits array is given, default to only the last token being an output
|
||||
t.output = batch_inp.logits
|
||||
? (batch_inp.logits[i] != 0)
|
||||
: (i == batch_inp.n_tokens - 1);
|
||||
|
||||
batch_ext->tokens.push_back(t);
|
||||
}
|
||||
}
|
||||
|
||||
llama_batch_compat::llama_batch_compat(llama_context * ctx, const llama_batch & batch_inp, size_t n_embd_row) {
|
||||
batch_ext = new llama_batch_ext(ctx);
|
||||
init(*batch_ext, batch_inp, n_embd_row);
|
||||
}
|
||||
|
||||
llama_batch_compat::~llama_batch_compat() {
|
||||
delete batch_ext;
|
||||
}
|
||||
|
||||
+76
-7
@@ -2,6 +2,7 @@
|
||||
|
||||
#include "llama.h"
|
||||
|
||||
#include "llama-arch.h"
|
||||
#include "llama-cparams.h"
|
||||
|
||||
#include <array>
|
||||
@@ -10,6 +11,7 @@
|
||||
#include <bitset>
|
||||
#include <memory>
|
||||
#include <unordered_map>
|
||||
#include <unordered_set>
|
||||
|
||||
// keep this struct lightweight
|
||||
struct llama_ubatch {
|
||||
@@ -68,19 +70,71 @@ struct llama_ubatch {
|
||||
std::shared_ptr<data_t> data;
|
||||
};
|
||||
|
||||
struct llama_hparams;
|
||||
|
||||
// MTP hook batches carry the target model's hidden state (n_embd_out size).
|
||||
// DFlash batches carry the fused target features at the encoder input width (n_embd_inp_enc size).
|
||||
// Normal batches carry token embeddings (n_embd_inp size).
|
||||
size_t llama_batch_ext_select_n_embd_inp(llama_context_type ctx_type, llm_arch arch, const llama_hparams & hparams);
|
||||
|
||||
struct llama_batch_ext {
|
||||
const size_t n_tokens_max; // max number of tokens that can be stored in the batch
|
||||
const size_t n_embd_inp; // decoder embd row width
|
||||
const size_t n_embd_inp_enc; // encoder embd row width (e.g. eagle3/dflash extracted features)
|
||||
const llama_seq_id n_seq_max; // max number of sequences
|
||||
llama_memory_i * mem; // memory for position inference
|
||||
const llama_token n_vocab; // max token ID that we accept
|
||||
const size_t n_pos_per_embd;
|
||||
|
||||
// actual embd row width of this batch, set by the first set_token_embd()
|
||||
// must be either n_embd_inp or n_embd_inp_enc; encode/decode verify it against the graph input
|
||||
size_t n_embd = 0;
|
||||
|
||||
struct token {
|
||||
llama_token id = LLAMA_TOKEN_NULL;
|
||||
bool has_embd = false; // whether embd_off is set
|
||||
size_t embd_off = 0; // index offset in the embd array
|
||||
bool output = false; // TODO: have dedicated output flags
|
||||
std::unordered_set<llama_seq_id> seq_ids;
|
||||
std::array<llama_pos, GGML_MROPE_SECTIONS> pos = {0, 0, 0, 0};
|
||||
};
|
||||
std::vector<token> tokens;
|
||||
std::vector<float> embd;
|
||||
|
||||
llama_batch_ext(llama_context * ctx);
|
||||
|
||||
// build without a llama_context, used by tests
|
||||
llama_batch_ext(
|
||||
size_t n_tokens_max,
|
||||
size_t n_embd_inp,
|
||||
size_t n_embd_inp_enc,
|
||||
llama_seq_id n_seq_max,
|
||||
llama_memory_i * mem,
|
||||
llama_token n_vocab,
|
||||
size_t n_pos_per_embd);
|
||||
|
||||
void clear();
|
||||
|
||||
// add an entry with an undefined position
|
||||
// the caller must set it explicitly via set_token_pos()
|
||||
int32_t add_token(llama_seq_id seq_id);
|
||||
|
||||
bool add_seq(int32_t idx, llama_seq_id seq_id);
|
||||
bool set_token_id(int32_t idx, llama_token id);
|
||||
bool set_token_embd(int32_t idx, llama_embd embd_in);
|
||||
bool set_token_pos(int32_t idx, const llama_pos * pos_in);
|
||||
bool set_output(int32_t idx, bool output_last);
|
||||
};
|
||||
|
||||
// a helper for sanitizing, fulfilling and splitting a batch
|
||||
class llama_batch_allocr {
|
||||
public:
|
||||
llama_batch_allocr(uint32_t n_pos_per_embd);
|
||||
|
||||
// sanitize and auto-gen missing data in the input batch
|
||||
// memory is optional. if provided will be used to check for sequence continuity and to determine the positions
|
||||
// convert a llama_batch_ext to internal llama_batch and sanitize it
|
||||
bool init(
|
||||
const llama_batch & batch_inp,
|
||||
const llama_batch_ext & batch_inp,
|
||||
const llama_vocab & vocab,
|
||||
const llama_memory_i * memory,
|
||||
uint32_t n_embd,
|
||||
uint32_t n_seq_max,
|
||||
bool output_all);
|
||||
|
||||
const llama_batch & get_batch() const;
|
||||
@@ -137,7 +191,9 @@ private:
|
||||
uint32_t n_seq_max;
|
||||
uint32_t n_outputs;
|
||||
|
||||
std::array<llama_seq_id, 1> seq_id_0 = {{ 0 }}; // default sequence id
|
||||
std::vector<llama_token> token_vec; // owned token IDs built from llama_batch_ext
|
||||
std::vector<float> embd_vec; // owned embeddings built from llama_batch_ext
|
||||
std::vector<llama_seq_id> seq_id_data; // flat storage for seq_id pointers below
|
||||
|
||||
std::vector<llama_pos> pos;
|
||||
std::vector<int32_t> n_seq_id;
|
||||
@@ -172,3 +228,16 @@ private:
|
||||
|
||||
int debug;
|
||||
};
|
||||
|
||||
// RAII translation layer: converts a llama_batch (old API) into a llama_batch_ext
|
||||
struct llama_batch_compat {
|
||||
llama_batch_ext * batch_ext;
|
||||
|
||||
// n_embd_row is the embd row width of batch_inp, 0 = use the decoder width
|
||||
llama_batch_compat(llama_context * ctx, const llama_batch & batch_inp, size_t n_embd_row = 0);
|
||||
~llama_batch_compat();
|
||||
|
||||
// fill an existing llama_batch_ext from a llama_batch (old API)
|
||||
// note: this is called directly by the tests, skipping llama_context creation
|
||||
static void init(llama_batch_ext & batch_ext, const llama_batch & batch_inp, size_t n_embd_row = 0);
|
||||
};
|
||||
|
||||
+52
-32
@@ -1463,24 +1463,25 @@ llm_graph_result * llama_context::process_ubatch(const llama_ubatch & ubatch, ll
|
||||
return res;
|
||||
}
|
||||
|
||||
int llama_context::encode(const llama_batch & batch_inp) {
|
||||
// MTP hook batches carry both token (next-token id) and embd (h_nextn row),
|
||||
// so accept either present rather than requiring exactly one.
|
||||
GGML_ASSERT(batch_inp.token || batch_inp.embd);
|
||||
|
||||
if (batch_inp.n_tokens == 0) {
|
||||
int llama_context::encode(const llama_batch_ext & batch_inp) {
|
||||
if (batch_inp.tokens.empty()) {
|
||||
LLAMA_LOG_ERROR("%s: n_tokens == 0\n", __func__);
|
||||
return -1;
|
||||
}
|
||||
|
||||
const auto & hparams = model.hparams;
|
||||
|
||||
if (batch_inp.n_embd > 0 && batch_inp.n_embd != hparams.n_embd_inp_enc()) {
|
||||
LLAMA_LOG_ERROR("%s: embd row width %zu does not match the encoder input %u\n",
|
||||
__func__, batch_inp.n_embd, hparams.n_embd_inp_enc());
|
||||
return -1;
|
||||
}
|
||||
|
||||
// eagle3/DFlash: features as encoder input, and non-draft paths fall back to model's input dim
|
||||
const int64_t n_embd = hparams.n_embd_inp_enc();
|
||||
const int64_t n_vocab = model.vocab.n_tokens();
|
||||
|
||||
// note: during encode, we always pass the full sequence starting from pos = 0
|
||||
if (!balloc->init(batch_inp, model.vocab, nullptr, n_embd, cparams.kv_unified ? LLAMA_MAX_SEQ : cparams.n_seq_max, true)) {
|
||||
// note: during encode, we always output all tokens and skip position continuity checks (output_all=true)
|
||||
if (!balloc->init(batch_inp, model.vocab, true)) {
|
||||
LLAMA_LOG_ERROR("%s: failed to initialize batch\n", __func__);
|
||||
return -1;
|
||||
}
|
||||
@@ -1701,29 +1702,27 @@ static bool needs_raw_logits(const llama_ubatch & ubatch, const std::map<llama_s
|
||||
return false; // all sequences use backend sampling
|
||||
}
|
||||
|
||||
int llama_context::decode(const llama_batch & batch_inp) {
|
||||
// MTP hook batches carry both token (next-token id) and embd (h_nextn row),
|
||||
// so accept either present rather than requiring exactly one.
|
||||
GGML_ASSERT(batch_inp.token || batch_inp.embd);
|
||||
|
||||
int llama_context::decode(const llama_batch_ext & batch_inp) {
|
||||
if (!memory) {
|
||||
LLAMA_LOG_DEBUG("%s: cannot decode batches with this context (calling encode() instead)\n", __func__);
|
||||
return encode(batch_inp);
|
||||
}
|
||||
|
||||
if (batch_inp.n_tokens == 0) {
|
||||
if (batch_inp.tokens.empty()) {
|
||||
LLAMA_LOG_ERROR("%s: n_tokens == 0\n", __func__);
|
||||
return -1;
|
||||
}
|
||||
|
||||
if (batch_inp.n_embd > 0 && batch_inp.n_embd != batch_inp.n_embd_inp) {
|
||||
LLAMA_LOG_ERROR("%s: embd row width %zu does not match the decoder input %zu\n",
|
||||
__func__, batch_inp.n_embd, batch_inp.n_embd_inp);
|
||||
return -1;
|
||||
}
|
||||
|
||||
const auto & vocab = model.vocab;
|
||||
const auto & hparams = model.hparams;
|
||||
|
||||
const int64_t n_vocab = vocab.n_tokens();
|
||||
const bool mtp_embd = cparams.ctx_type == LLAMA_CONTEXT_TYPE_MTP && batch_inp.embd;
|
||||
// DFlash embd batches carry the fused target features at the encoder input width
|
||||
const bool dflash_embd = model.arch == LLM_ARCH_DFLASH && batch_inp.embd;
|
||||
const int64_t n_embd = mtp_embd ? hparams.n_embd_out() : dflash_embd ? hparams.n_embd_inp_enc() : hparams.n_embd_inp();
|
||||
|
||||
// when computing embeddings, all tokens are output
|
||||
const bool output_all = cparams.embeddings;
|
||||
@@ -1731,20 +1730,17 @@ int llama_context::decode(const llama_batch & batch_inp) {
|
||||
|
||||
const uint32_t n_seq_max = cparams.kv_unified ? LLAMA_MAX_SEQ : cparams.n_seq_max;
|
||||
|
||||
// embedding contexts output every token even when batch.logits is not set
|
||||
if (has_samplers && (output_all || batch_inp.logits)) {
|
||||
// TODO: avoid this workaround in the future
|
||||
// embedding contexts output every token even when no token is explicitly marked as output
|
||||
if (has_samplers) {
|
||||
std::vector<int32_t> seq_output_count(n_seq_max, 0);
|
||||
|
||||
for (int32_t i = 0; i < batch_inp.n_tokens; ++i) {
|
||||
if (!output_all && batch_inp.logits[i] == 0) {
|
||||
for (const auto & tok : batch_inp.tokens) {
|
||||
if (!output_all && !tok.output) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const int ns = batch_inp.n_seq_id ? batch_inp.n_seq_id[i] : 1;
|
||||
|
||||
for (int32_t s = 0; s < ns; ++s) {
|
||||
const llama_seq_id seq_id = batch_inp.seq_id ? batch_inp.seq_id[i][s] : 0;
|
||||
|
||||
for (auto seq_id : tok.seq_ids) {
|
||||
if (seq_id < 0 || (uint32_t) seq_id >= n_seq_max) {
|
||||
continue;
|
||||
}
|
||||
@@ -1762,7 +1758,7 @@ int llama_context::decode(const llama_batch & batch_inp) {
|
||||
}
|
||||
}
|
||||
|
||||
if (!balloc->init(batch_inp, vocab, memory.get(), n_embd, n_seq_max, output_all)) {
|
||||
if (!balloc->init(batch_inp, vocab, output_all)) {
|
||||
LLAMA_LOG_ERROR("%s: failed to initialize batch\n", __func__);
|
||||
return -1;
|
||||
}
|
||||
@@ -3557,9 +3553,13 @@ void llama_context::opt_epoch_iter(
|
||||
batch.logits [pos_batch] = true;
|
||||
}
|
||||
|
||||
if (!balloc->init(batch, model.vocab, nullptr, model.hparams.n_embd_inp(), cparams.kv_unified ? LLAMA_MAX_SEQ : cparams.n_seq_max, true)) {
|
||||
LLAMA_LOG_ERROR("%s: failed to initialize batch\n", __func__);
|
||||
return;
|
||||
// TODO: use llama_batch_ext here
|
||||
{
|
||||
llama_batch_compat compat(this, batch);
|
||||
if (!balloc->init(*compat.batch_ext, model.vocab, true)) {
|
||||
LLAMA_LOG_ERROR("%s: failed to initialize batch\n", __func__);
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
const uint32_t n_tokens_all = balloc->get_n_tokens();
|
||||
@@ -4310,6 +4310,18 @@ size_t llama_state_seq_load_file(llama_context * ctx, const char * filepath, lla
|
||||
}
|
||||
}
|
||||
|
||||
// compat: llama_batch -> llama_batch_ext -> encode/decode
|
||||
|
||||
int llama_context::encode(const llama_batch & batch_inp) {
|
||||
llama_batch_compat compat(this, batch_inp, model.hparams.n_embd_inp_enc());
|
||||
return encode(*compat.batch_ext);
|
||||
}
|
||||
|
||||
int llama_context::decode(const llama_batch & batch_inp) {
|
||||
llama_batch_compat compat(this, batch_inp);
|
||||
return decode(*compat.batch_ext);
|
||||
}
|
||||
|
||||
///
|
||||
|
||||
int32_t llama_encode(
|
||||
@@ -4399,6 +4411,14 @@ void llama_opt_epoch(
|
||||
callback_eval);
|
||||
}
|
||||
|
||||
int32_t llama_process(llama_context * ctx, llama_process_type type, llama_batch_ext * batch) {
|
||||
switch (type) {
|
||||
case LLAMA_PROCESS_TYPE_ENCODE: return ctx->encode(*batch);
|
||||
case LLAMA_PROCESS_TYPE_DECODE: return ctx->decode(*batch);
|
||||
}
|
||||
return -1;
|
||||
}
|
||||
|
||||
//
|
||||
// ext
|
||||
//
|
||||
|
||||
@@ -141,6 +141,10 @@ struct llama_context {
|
||||
llama_memory_context_i * mctx,
|
||||
ggml_status & ret);
|
||||
|
||||
int encode(const llama_batch_ext & batch_inp);
|
||||
int decode(const llama_batch_ext & batch_inp);
|
||||
|
||||
// compat version
|
||||
int encode(const llama_batch & batch_inp);
|
||||
int decode(const llama_batch & batch_inp);
|
||||
|
||||
|
||||
@@ -283,7 +283,8 @@ bool llama_hparams::is_ple(uint32_t il) const {
|
||||
}
|
||||
|
||||
uint32_t llama_hparams::n_pos_per_embd() const {
|
||||
return rope_type == LLAMA_ROPE_TYPE_MROPE || rope_type == LLAMA_ROPE_TYPE_IMROPE ? 4 : 1;
|
||||
return (rope_type == LLAMA_ROPE_TYPE_MROPE || rope_type == LLAMA_ROPE_TYPE_IMROPE)
|
||||
? GGML_MROPE_SECTIONS : 1;
|
||||
}
|
||||
|
||||
bool llama_hparams::is_swa(uint32_t il) const {
|
||||
|
||||
Reference in New Issue
Block a user