mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-03 03:17:32 -05:00
* llama: preserve original batch order for layer inputs Assisted-by: Codex * tests: cover layer-input order across KV layouts Assisted-by: Codex * tests: exercise layer-input ordering on CUDA devices Assisted-by: Codex * llama: make layer input reordering compatible with tensor split Copy each microbatch tensor from offset zero and restore original row order after synchronization. Extend the layer-input regression to cover tensor split and repeated reads and decodes. Assisted-by: Codex * llama: restore token order for unmasked NextN embeddings Use the original-token mapping for unmasked NextN rows, including when layer-input capture is disabled. Keep masked NextN rows on the logits output mapping and preserve offset-zero tensor copies. Extend the existing regression to cover NextN alone, combined layer capture, and masked outputs with repeated decodes and getters. Validation: all 256 CPU/CUDA/tensor configurations pass. Qwen3.8-27B Q4_K_M MTP completes MT-Bench at concurrency 16 before and after. Assisted-by: Codex * ggml: fix WebGPU reservation and OpenVINO hidden-state capture Reserve WebGPU vector attention scratch across batch sizes and refresh reservations when NextN capture settings change. Preserve requested OpenVINO outputs, dynamic shapes, sequence counts, and current graph bindings. Extend existing WebGPU regression coverage and enable strict allocation checks. Assisted-by: Codex * llama: defer regression test and backend fixes to follow-ups Keep this PR focused on restoring token order for layer inputs and unmasked NextN embeddings. Remove the added regression test, OpenVINO and WebGPU changes, and the separate NextN reservation change. Assisted-by: Codex * llama: keep n_embd declaration in its original position Assisted-by: Codex * llama : pass token count to layer input extraction Assisted-by: Codex * llama : name original batch indices batch_idxs Assisted-by: Codex * llama : name extracted embedding indices embd_batch_idxs Assisted-by: Codex * llama : tag target embedding reordering Assisted-by: Codex * llama : tag extraction and name the index capture flag Assisted-by: Codex
1326 lines
40 KiB
C++
1326 lines
40 KiB
C++
#include "llama-batch.h"
|
|
|
|
#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>
|
|
#include <algorithm>
|
|
#include <sstream>
|
|
|
|
llama_batch_allocr::llama_batch_allocr(uint32_t n_pos_per_embd) : n_pos_per_embd(n_pos_per_embd) {
|
|
const char * LLAMA_BATCH_DEBUG = getenv("LLAMA_BATCH_DEBUG");
|
|
debug = LLAMA_BATCH_DEBUG ? atoi(LLAMA_BATCH_DEBUG) : 0;
|
|
|
|
seq_pos.resize(LLAMA_MAX_SEQ);
|
|
seq_cpl.resize(LLAMA_MAX_SEQ);
|
|
for (auto & cur : seq_cpl) {
|
|
cur.resize(LLAMA_MAX_SEQ);
|
|
}
|
|
|
|
seq_idx.resize(LLAMA_MAX_SEQ, -1);
|
|
}
|
|
|
|
bool llama_batch_allocr::init(
|
|
const llama_batch_ext & batch_inp,
|
|
const llama_vocab & vocab,
|
|
bool output_all) {
|
|
clear();
|
|
|
|
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;
|
|
|
|
const int32_t n_tok = (int32_t) batch_inp.tokens.size();
|
|
|
|
GGML_ASSERT(n_tok > 0);
|
|
|
|
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;
|
|
}
|
|
|
|
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];
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
//
|
|
// 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;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
//
|
|
// build output/logits array
|
|
//
|
|
|
|
{
|
|
output.resize(n_tok, 0);
|
|
for (int32_t i = 0; i < n_tok; ++i) {
|
|
output[i] = batch_inp.tokens[i].output ? 1 : 0;
|
|
}
|
|
|
|
if (output_all) {
|
|
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);
|
|
}
|
|
}
|
|
}
|
|
|
|
//
|
|
// 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
|
|
//
|
|
|
|
// count the outputs in this batch
|
|
for (int32_t i = 0; i < batch.n_tokens; ++i) {
|
|
n_outputs += batch.logits[i] != 0;
|
|
}
|
|
|
|
has_cpl = false;
|
|
|
|
// determine coupled sequences
|
|
// these are pairs of sequences that have at least one token in the input batch that is assigned to both of them
|
|
for (int32_t i = 0; i < batch.n_tokens; ++i) {
|
|
const llama_seq_id s0 = batch.seq_id[i][0];
|
|
|
|
for (int32_t s = 0; s < batch.n_seq_id[i]; ++s) {
|
|
const llama_seq_id s1 = batch.seq_id[i][s];
|
|
|
|
seq_pos[s1].insert(batch.pos[i]);
|
|
|
|
if (s > 0) {
|
|
// mark that sequence s1 is coupled to s0
|
|
seq_cpl[s1][s0] = true;
|
|
|
|
// note: tracking the other way around is not necessary for now
|
|
//seq_cpl[s0][s1] = true;
|
|
|
|
has_cpl = true;
|
|
}
|
|
}
|
|
}
|
|
|
|
// precompute the sequence sets for each token and determine the unique sequence ids that participate in the batch
|
|
{
|
|
seq_set_t seq_set_unq;
|
|
|
|
for (int32_t i = 0; i < batch.n_tokens; ++i) {
|
|
seq_set_t cur;
|
|
for (int32_t s = 0; s < batch.n_seq_id[i]; ++s) {
|
|
const llama_seq_id seq_id = batch.seq_id[i][s];
|
|
|
|
cur .set(seq_id);
|
|
seq_set_unq.set(seq_id);
|
|
}
|
|
|
|
seq_set.push_back(cur);
|
|
seq_set_map[cur].push_back(i);
|
|
}
|
|
|
|
for (uint32_t s = 0; s < n_seq_max; ++s) {
|
|
if (seq_set_unq.test(s)) {
|
|
seq_idx[s] = seq_id_unq.size();
|
|
seq_id_unq.push_back(s);
|
|
}
|
|
}
|
|
}
|
|
|
|
if (debug > 0) {
|
|
LLAMA_LOG_DEBUG("%s: input batch info:\n", __func__);
|
|
|
|
llama_ubatch ubatch {
|
|
/*.b_equal_seqs =*/ false,
|
|
/*.n_tokens =*/ (uint32_t) batch.n_tokens,
|
|
/*.n_seq_tokens =*/ (uint32_t) 1,
|
|
/*.n_seqs =*/ (uint32_t) batch.n_tokens,
|
|
/*.n_seqs_unq =*/ (uint32_t) this->seq_id_unq.size(),
|
|
/*.n_pos =*/ n_pos_per_embd,
|
|
/*.token =*/ batch.token,
|
|
/*.embd =*/ batch.embd,
|
|
/*.pos =*/ batch.pos,
|
|
/*.n_seq_id =*/ batch.n_seq_id,
|
|
/*.seq_id =*/ batch.seq_id,
|
|
/*.seq_id_unq =*/ this->seq_id_unq.data(),
|
|
/*.seq_idx =*/ this->seq_idx.data(),
|
|
/*.output =*/ batch.logits,
|
|
/*.data =*/ {},
|
|
};
|
|
|
|
ubatch_print(ubatch, debug);
|
|
|
|
LLAMA_LOG_DEBUG("%s: seq = [\n", __func__);
|
|
for (int s0 = 0; s0 < (int) seq_pos.size(); ++s0) {
|
|
if (seq_pos[s0].empty()) {
|
|
continue;
|
|
}
|
|
|
|
std::stringstream ss;
|
|
for (int s1 = 0; s1 < (int) seq_cpl[s0].size(); ++s1) {
|
|
if (seq_cpl[s0][s1]) {
|
|
ss << s1 << " ";
|
|
}
|
|
}
|
|
|
|
LLAMA_LOG_DEBUG("%s: %4d: pos = [%4d, %4d], cpl = %s\n",
|
|
__func__, s0, seq_pos_min(s0), seq_pos_max(s0), ss.str().empty() ? "-" : ss.str().c_str());
|
|
}
|
|
LLAMA_LOG_DEBUG("%s: ]\n", __func__);
|
|
}
|
|
|
|
//
|
|
// consistency checks
|
|
//
|
|
|
|
if (n_pos_per_embd > 1) {
|
|
// M-RoPE case: allow position to "jump" forward only (non-continuous positions are allowed)
|
|
for (uint32_t s = 0; s < n_seq_max; ++s) {
|
|
if (seq_pos[s].empty()) {
|
|
continue;
|
|
}
|
|
|
|
const llama_pos p0 = mem ? mem->seq_pos_max(s) : -1;
|
|
|
|
if (batch.token) {
|
|
if (p0 >= 0 && p0 >= seq_pos_min(s)) {
|
|
LLAMA_LOG_ERROR(
|
|
"%s: the tokens of sequence %d in the input batch have inconsistent sequence positions:\n"
|
|
" - the last position stored in the memory module of the context (i.e. the KV cache) for sequence %d is X = %d\n"
|
|
" - the tokens for sequence %d in the input batch have a starting position of Y = %d\n"
|
|
" for M-RoPE, it is required that the position satisfies: X < Y\n",
|
|
__func__, s, s, p0, s, seq_pos_min(s));
|
|
|
|
return false;
|
|
}
|
|
} else {
|
|
// embedding inputs can have overlapping positions
|
|
if (p0 >= 0 && p0 > seq_pos_min(s)) {
|
|
LLAMA_LOG_ERROR(
|
|
"%s: the tokens of sequence %d in the input batch have inconsistent sequence positions:\n"
|
|
" - the last position stored in the memory module of the context (i.e. the KV cache) for sequence %d is X = %d\n"
|
|
" - the tokens for sequence %d in the input batch have a starting position of Y = %d\n"
|
|
" for M-RoPE, it is required that the position satisfies: X <= Y\n",
|
|
__func__, s, s, p0, s, seq_pos_min(s));
|
|
|
|
return false;
|
|
}
|
|
}
|
|
}
|
|
} else {
|
|
for (uint32_t s = 0; s < n_seq_max; ++s) {
|
|
if (seq_pos[s].empty()) {
|
|
continue;
|
|
}
|
|
|
|
const llama_pos p0 = mem ? mem->seq_pos_max(s) : -1;
|
|
|
|
if (p0 >= 0) {
|
|
bool ok = true;
|
|
|
|
if (seq_pos_min(s) != p0 + 1) {
|
|
ok = false;
|
|
}
|
|
|
|
if (!ok) {
|
|
LLAMA_LOG_ERROR(
|
|
"%s: the tokens of sequence %d in the input batch have inconsistent sequence positions:\n"
|
|
" - the last position stored in the memory module of the context (i.e. the KV cache) for sequence %d is X = %d\n"
|
|
" - the tokens for sequence %d in the input batch have a starting position of Y = %d\n"
|
|
" it is required that the sequence positions remain consecutive: Y = X + 1\n",
|
|
__func__, s, s, p0, s, seq_pos_min(s));
|
|
|
|
return false;
|
|
}
|
|
}
|
|
|
|
if (seq_pos_max(s) - seq_pos_min(s) + 1 > (int) seq_pos[s].size()) {
|
|
LLAMA_LOG_ERROR("%s: sequence %d positions are not continuous\n", __func__, s);
|
|
return false;
|
|
}
|
|
}
|
|
}
|
|
|
|
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 (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;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
// disallow partial sequence sub-sets:
|
|
//
|
|
// invalid: x
|
|
// i: 0 1 2 ...
|
|
// ---------------------------------------
|
|
// seq_id[i][0]: 0 0 1
|
|
// seq_id[i][1]: 1 1 2
|
|
// seq_id[i][2]: 2
|
|
//
|
|
// disallow decreasing sequence positions:
|
|
//
|
|
// invalid: x
|
|
// i: 0 1 2 3 4 5 6 ...
|
|
// ---------------------------------------
|
|
// pos[i]: 4 5 0 1 6 2 3
|
|
// seq_id[i][0]: 0 0 1 1 0 1 0
|
|
//
|
|
{
|
|
seq_set_t cur_seq_set[LLAMA_MAX_SEQ];
|
|
for (uint32_t s = 0; s < n_seq_max; ++s) {
|
|
cur_seq_set[s].set();
|
|
}
|
|
|
|
llama_pos cur_seq_pos[LLAMA_MAX_SEQ];
|
|
for (uint32_t s = 0; s < n_seq_max; ++s) {
|
|
cur_seq_pos[s] = -1;
|
|
}
|
|
|
|
for (int32_t i = 0; i < batch.n_tokens; ++i) {
|
|
const llama_pos pos = batch.pos[i];
|
|
|
|
for (int32_t s = 0; s < batch.n_seq_id[i]; ++s) {
|
|
const llama_seq_id seq_id = batch.seq_id[i][s];
|
|
|
|
cur_seq_set[seq_id] &= seq_set[i];
|
|
|
|
if (cur_seq_set[seq_id].none()) {
|
|
LLAMA_LOG_ERROR("%s: sequence %d belongs to incompatible sequence sets (not allowed)\n", __func__, seq_id);
|
|
return false;
|
|
}
|
|
|
|
if (pos < cur_seq_pos[seq_id]) {
|
|
LLAMA_LOG_ERROR("%s: sequence %d positions are decreasing (not allowed)\n", __func__, seq_id);
|
|
return false;
|
|
}
|
|
|
|
cur_seq_pos[seq_id] = pos;
|
|
}
|
|
}
|
|
}
|
|
|
|
split_reset();
|
|
|
|
return true;
|
|
}
|
|
|
|
llama_ubatch llama_batch_allocr::ubatch_reserve(uint32_t n_seq_tokens, uint32_t n_seqs) {
|
|
const uint32_t n_tokens = n_seq_tokens*n_seqs;
|
|
|
|
clear();
|
|
split_reset();
|
|
|
|
const int64_t n_pos_all = (int64_t) n_tokens*n_pos_per_embd;
|
|
|
|
auto udata = std::make_shared<llama_ubatch::data_t>();
|
|
|
|
udata->token .resize(n_tokens);
|
|
udata->embd .clear();
|
|
udata->pos .resize(n_pos_all);
|
|
udata->n_seq_id .resize(n_tokens);
|
|
udata->seq_id .resize(n_tokens);
|
|
udata->seq_id_unq.resize(0);
|
|
udata->seq_idx .resize(LLAMA_MAX_SEQ, -1);
|
|
udata->output .resize(n_tokens);
|
|
|
|
for (uint32_t s = 0; s < n_seqs; ++s) {
|
|
udata->seq_idx[s] = s;
|
|
udata->seq_id_unq.push_back(s);
|
|
}
|
|
|
|
llama_ubatch res {
|
|
/*.b_equal_seqs =*/ true,
|
|
/*.n_tokens =*/ n_tokens,
|
|
/*.n_seq_tokens =*/ n_seq_tokens,
|
|
/*.n_seqs =*/ n_seqs,
|
|
/*.n_seqs_unq =*/ n_seqs,
|
|
/*.n_pos =*/ n_pos_per_embd,
|
|
|
|
/*.token =*/ udata->token.data(),
|
|
/*.embd =*/ nullptr,
|
|
/*.pos =*/ udata->pos.data(),
|
|
/*.n_seq_id =*/ udata->n_seq_id.data(),
|
|
/*.seq_id =*/ udata->seq_id.data(),
|
|
/*.seq_id_unq =*/ udata->seq_id_unq.data(),
|
|
/*.seq_idx =*/ udata->seq_idx.data(),
|
|
/*.output =*/ udata->output.data(),
|
|
/*.data =*/ std::move(udata),
|
|
};
|
|
|
|
return res;
|
|
}
|
|
|
|
const llama_batch & llama_batch_allocr::get_batch() const {
|
|
return batch;
|
|
}
|
|
|
|
uint32_t llama_batch_allocr::get_n_tokens() const {
|
|
return batch.n_tokens;
|
|
}
|
|
|
|
uint32_t llama_batch_allocr::get_n_outputs() const {
|
|
return n_outputs;
|
|
}
|
|
|
|
uint32_t llama_batch_allocr::get_n_used() const {
|
|
return n_used;
|
|
}
|
|
|
|
std::vector<int32_t> & llama_batch_allocr::get_out_ids() {
|
|
return out_ids;
|
|
}
|
|
|
|
llama_pos llama_batch_allocr::seq_pos_min(llama_seq_id seq_id) const {
|
|
return seq_pos[seq_id].empty() ? -1 : *seq_pos[seq_id].begin();
|
|
}
|
|
|
|
llama_pos llama_batch_allocr::seq_pos_max(llama_seq_id seq_id) const {
|
|
return seq_pos[seq_id].empty() ? -1 : *seq_pos[seq_id].rbegin();
|
|
}
|
|
|
|
void llama_batch_allocr::split_reset() {
|
|
out_ids.clear();
|
|
|
|
n_used = 0;
|
|
|
|
used.clear();
|
|
used.resize(get_n_tokens(), false);
|
|
}
|
|
|
|
llama_ubatch llama_batch_allocr::split_simple(uint32_t n_ubatch) {
|
|
// find the first unused token
|
|
uint32_t cur_idx = 0;
|
|
while (cur_idx < used.size() && used[cur_idx]) {
|
|
++cur_idx;
|
|
}
|
|
|
|
// we are done
|
|
if (cur_idx >= used.size()) {
|
|
return {};
|
|
}
|
|
|
|
std::vector<int32_t> idxs;
|
|
|
|
while (true) {
|
|
idxs.push_back(cur_idx);
|
|
|
|
used[cur_idx] = true;
|
|
++n_used;
|
|
|
|
++cur_idx;
|
|
|
|
if (cur_idx >= used.size()) {
|
|
break;
|
|
}
|
|
|
|
if (idxs.size() >= n_ubatch) {
|
|
break;
|
|
}
|
|
}
|
|
|
|
return ubatch_add(idxs, idxs.size(), false);
|
|
}
|
|
|
|
llama_ubatch llama_batch_allocr::split_equal(uint32_t n_ubatch, bool sequential, uint32_t n_keep_tail) {
|
|
if (sequential && has_cpl) {
|
|
LLAMA_LOG_ERROR("%s: sequential split is not supported when there are coupled sequences in the input batch (you may need to use the -kvu flag)\n", __func__);
|
|
|
|
return {};
|
|
}
|
|
|
|
std::vector<seq_set_t> cur_seq_set;
|
|
|
|
llama_seq_id last_seq_id = -1;
|
|
|
|
// determine the non-overlapping sequence sets participating in this ubatch
|
|
for (int32_t i = 0; i < batch.n_tokens; ++i) {
|
|
if (used[i]) {
|
|
continue;
|
|
}
|
|
|
|
bool add = true;
|
|
|
|
for (uint32_t s = 0; s < cur_seq_set.size(); ++s) {
|
|
// no overlap with existing sequence sets:
|
|
if (!(cur_seq_set[s] & seq_set[i]).none()) {
|
|
add = false;
|
|
break;
|
|
}
|
|
}
|
|
|
|
// accept only increasing sequence ids
|
|
if (sequential) {
|
|
add = add && (cur_seq_set.empty() || batch.seq_id[i][0] == last_seq_id + 1);
|
|
}
|
|
|
|
if (add) {
|
|
cur_seq_set.push_back(seq_set[i]);
|
|
|
|
last_seq_id = batch.seq_id[i][0];
|
|
|
|
if (cur_seq_set.size() > n_ubatch) {
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
|
|
uint32_t n_seqs = cur_seq_set.size();
|
|
|
|
// we are done
|
|
if (n_seqs == 0) {
|
|
return {};
|
|
}
|
|
|
|
// the current batch index of each sequence set
|
|
std::vector<int32_t> cur_idx(n_seqs, 0);
|
|
|
|
for (uint32_t s = 0; s < n_seqs; ++s) {
|
|
while (used[seq_set_map[cur_seq_set[s]][cur_idx[s]]]) {
|
|
++cur_idx[s];
|
|
}
|
|
}
|
|
|
|
// the list of batch indices for each sequence set
|
|
// at the end we will concat these to get the final ubatch
|
|
std::vector<idx_vec_t> idxs_per_seq(n_seqs);
|
|
|
|
while (true) {
|
|
// we can only add new n_seq_tokens tokens if all the sequence sets have at least 1 more unused tokens and
|
|
// if we haven't reached n_ubatch
|
|
bool can_expand = true;
|
|
|
|
for (uint32_t s = 0; s < n_seqs; ++s) {
|
|
if (cur_idx[s] >= (int32_t) seq_set_map[cur_seq_set[s]].size()) {
|
|
can_expand = false;
|
|
break;
|
|
}
|
|
}
|
|
|
|
if (!can_expand) {
|
|
break;
|
|
}
|
|
|
|
for (uint32_t s = 0; s < n_seqs; ++s) {
|
|
const int32_t idx = seq_set_map[cur_seq_set[s]][cur_idx[s]];
|
|
|
|
idxs_per_seq[s].push_back(idx);
|
|
|
|
used[idx] = true;
|
|
++n_used;
|
|
|
|
++cur_idx[s];
|
|
}
|
|
|
|
if ((idxs_per_seq[0].size() + 1)*n_seqs > n_ubatch) {
|
|
break;
|
|
}
|
|
}
|
|
|
|
// if n_keep_tail > 0, keep only the seqs that either finish in this ubatch or have at least
|
|
// n_keep_tail tokens remaining for a future ubatch, so that the trailing n_keep_tail tokens
|
|
// of each seq are never split across ubatches
|
|
if (n_keep_tail > 0) {
|
|
GGML_ASSERT(n_ubatch > n_keep_tail);
|
|
|
|
auto n_remaining = [&](uint32_t s) {
|
|
return (uint32_t) (seq_set_map[cur_seq_set[s]].size() - cur_idx[s]);
|
|
};
|
|
|
|
// keep the longest prefix of seqs that satisfy the constraint, to preserve sequential seq ids
|
|
uint32_t n_keep = 0;
|
|
while (n_keep < n_seqs) {
|
|
const uint32_t remaining = n_remaining(n_keep);
|
|
|
|
if (remaining != 0 && remaining < n_keep_tail) {
|
|
break;
|
|
}
|
|
|
|
n_keep++;
|
|
}
|
|
|
|
// all seqs violate the constraint - resolve the first one directly and emit it alone
|
|
if (n_keep == 0) {
|
|
auto & idxs = idxs_per_seq[0];
|
|
|
|
const auto & seq_idxs = seq_set_map[cur_seq_set[0]];
|
|
|
|
if (idxs.size() + n_remaining(0) <= n_ubatch) {
|
|
// extend the seq to completion
|
|
while (n_remaining(0) > 0) {
|
|
const int32_t idx = seq_idxs[cur_idx[0]];
|
|
|
|
idxs.push_back(idx);
|
|
|
|
used[idx] = true;
|
|
++n_used;
|
|
|
|
++cur_idx[0];
|
|
}
|
|
} else {
|
|
// truncate the seq so that at least n_keep_tail tokens remain
|
|
while (n_remaining(0) < n_keep_tail) {
|
|
used[idxs.back()] = false;
|
|
--n_used;
|
|
|
|
idxs.pop_back();
|
|
|
|
--cur_idx[0];
|
|
}
|
|
}
|
|
|
|
n_keep = 1;
|
|
}
|
|
|
|
// return the tokens of the deferred seqs back to the pool
|
|
for (uint32_t s = n_keep; s < n_seqs; ++s) {
|
|
for (const int32_t idx : idxs_per_seq[s]) {
|
|
used[idx] = false;
|
|
--n_used;
|
|
}
|
|
}
|
|
|
|
n_seqs = n_keep;
|
|
}
|
|
|
|
// concat the per-sequence-set lists
|
|
std::vector<int32_t> idxs;
|
|
|
|
for (uint32_t s = 0; s < n_seqs; ++s) {
|
|
idxs.insert(idxs.end(), idxs_per_seq[s].begin(), idxs_per_seq[s].end());
|
|
}
|
|
|
|
return ubatch_add(idxs, n_seqs, true);
|
|
}
|
|
|
|
llama_ubatch llama_batch_allocr::split_seq(uint32_t n_ubatch) {
|
|
// find the first unused token
|
|
uint32_t cur_idx = 0;
|
|
while (cur_idx < used.size() && used[cur_idx]) {
|
|
++cur_idx;
|
|
}
|
|
|
|
// we are done
|
|
if (cur_idx >= used.size()) {
|
|
return {};
|
|
}
|
|
|
|
// this is the starting sequence set
|
|
// we allow adding tokens only if their sequence set is a subset of the current sequence set
|
|
auto cur_seq_set = seq_set[cur_idx];
|
|
|
|
std::vector<int32_t> idxs;
|
|
|
|
while (true) {
|
|
idxs.push_back(cur_idx);
|
|
|
|
used[cur_idx] = true;
|
|
++n_used;
|
|
|
|
if (idxs.size() >= n_ubatch) {
|
|
break;
|
|
}
|
|
|
|
do {
|
|
++cur_idx;
|
|
} while (cur_idx < get_n_tokens() && (used[cur_idx] || ((cur_seq_set & seq_set[cur_idx]) != seq_set[cur_idx])));
|
|
|
|
if (cur_idx == get_n_tokens()) {
|
|
break;
|
|
}
|
|
|
|
cur_seq_set = seq_set[cur_idx];
|
|
}
|
|
|
|
return ubatch_add(idxs, 1, true);
|
|
}
|
|
|
|
void llama_batch_allocr::clear() {
|
|
n_outputs = 0;
|
|
|
|
batch = {};
|
|
|
|
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();
|
|
}
|
|
|
|
for (auto & cur : seq_cpl) {
|
|
std::fill(cur.begin(), cur.end(), false);
|
|
}
|
|
|
|
seq_set.clear();
|
|
|
|
seq_set_map.clear();
|
|
|
|
std::fill(seq_idx.begin(), seq_idx.end(), -1);
|
|
}
|
|
|
|
llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, uint32_t n_seqs, bool equal_seqs) {
|
|
const uint32_t n_tokens = idxs.size();
|
|
|
|
assert(n_tokens%n_seqs == 0);
|
|
|
|
auto udata = std::make_shared<llama_ubatch::data_t>();
|
|
|
|
const int64_t n_embd_all = batch.embd ? (int64_t) n_tokens*n_embd : 0;
|
|
const int64_t n_pos_all = (int64_t) n_tokens*n_pos_per_embd;
|
|
|
|
udata->token .resize(n_tokens);
|
|
udata->embd .resize(n_embd_all);
|
|
udata->pos .resize(n_pos_all);
|
|
udata->n_seq_id .resize(n_tokens);
|
|
udata->seq_id .resize(n_tokens);
|
|
udata->seq_id_unq.resize(0);
|
|
udata->seq_idx .resize(LLAMA_MAX_SEQ, -1);
|
|
udata->output .resize(n_tokens);
|
|
|
|
udata->batch_idxs = idxs;
|
|
udata->seq_id_data.reserve(n_tokens);
|
|
|
|
seq_set_t seq_set_unq;
|
|
|
|
for (size_t i = 0; i < idxs.size(); ++i) {
|
|
if (batch.token) {
|
|
udata->token[i] = batch.token[idxs[i]];
|
|
}
|
|
|
|
if (batch.embd) {
|
|
memcpy(udata->embd.data() + i*n_embd, batch.embd + (int64_t) idxs[i]*n_embd, n_embd*sizeof(float));
|
|
}
|
|
|
|
for (size_t j = 0; j < (size_t)n_pos_per_embd; ++j) {
|
|
// if we are using M-RoPE
|
|
// if the current batch is text, we need to broadcast the same position across all RoPE sections
|
|
// otherwise, the input batch is image embeddings, we copy the positions as-is
|
|
// if we are not using M-RoPE, there is only one position per token (this loop runs only once)
|
|
size_t src_off = batch.token ? 0 : j*batch.n_tokens;
|
|
udata->pos[j*n_tokens + i] = batch.pos[src_off + idxs[i]];
|
|
}
|
|
|
|
udata->n_seq_id[i] = batch.n_seq_id[idxs[i]];
|
|
udata->output[i] = batch.logits[idxs[i]];
|
|
|
|
for (int s = 0; s < udata->n_seq_id[i]; ++s) {
|
|
const llama_seq_id seq_id = batch.seq_id[idxs[i]][s];
|
|
|
|
udata->seq_id_data.push_back(seq_id);
|
|
seq_set_unq.set(seq_id);
|
|
}
|
|
|
|
if (udata->output[i]) {
|
|
out_ids.push_back(idxs[i]);
|
|
}
|
|
}
|
|
|
|
llama_seq_id * seq_id_ptr = udata->seq_id_data.data();
|
|
for (size_t i = 0; i < idxs.size(); ++i) {
|
|
udata->seq_id[i] = seq_id_ptr;
|
|
seq_id_ptr += udata->n_seq_id[i];
|
|
}
|
|
|
|
for (uint32_t s = 0; s < n_seq_max; ++s) {
|
|
if (seq_set_unq.test(s)) {
|
|
udata->seq_idx[s] = udata->seq_id_unq.size();
|
|
udata->seq_id_unq.push_back(s);
|
|
}
|
|
}
|
|
|
|
llama_ubatch res {
|
|
/*.b_equal_seqs =*/ equal_seqs,
|
|
/*.n_tokens =*/ n_tokens,
|
|
/*.n_seq_tokens =*/ n_tokens/n_seqs,
|
|
/*.n_seqs =*/ n_seqs,
|
|
/*.n_seqs_unq =*/ (uint32_t) udata->seq_id_unq.size(),
|
|
/*.n_pos =*/ n_pos_per_embd,
|
|
|
|
/*.token =*/ batch.token ? udata->token.data() : nullptr,
|
|
/*.embd =*/ batch.embd ? udata->embd.data() : nullptr,
|
|
/*.pos =*/ udata->pos.data(),
|
|
/*.n_seq_id =*/ udata->n_seq_id.data(),
|
|
/*.seq_id =*/ udata->seq_id.data(),
|
|
/*.seq_id_unq =*/ udata->seq_id_unq.data(),
|
|
/*.seq_idx =*/ udata->seq_idx.data(),
|
|
/*.output =*/ udata->output.data(),
|
|
/*.data =*/ std::move(udata),
|
|
};
|
|
|
|
if (debug > 0) {
|
|
LLAMA_LOG_DEBUG("%s: added ubatch to split:\n", __func__);
|
|
|
|
ubatch_print(res, debug);
|
|
}
|
|
|
|
return res;
|
|
}
|
|
|
|
void llama_batch_allocr::ubatch_print(const llama_ubatch & ubatch, int debug) {
|
|
if (debug > 0) {
|
|
LLAMA_LOG_DEBUG("%s: equal_seqs = %d\n", __func__, ubatch.equal_seqs());
|
|
LLAMA_LOG_DEBUG("%s: n_tokens = %d\n", __func__, ubatch.n_tokens);
|
|
LLAMA_LOG_DEBUG("%s: n_seq_tokens = %d\n", __func__, ubatch.n_seq_tokens);
|
|
LLAMA_LOG_DEBUG("%s: n_seqs = %d\n", __func__, ubatch.n_seqs);
|
|
LLAMA_LOG_DEBUG("%s: n_seqs_unq = %d\n", __func__, ubatch.n_seqs_unq);
|
|
|
|
std::stringstream ss_seq_id_unq;
|
|
std::stringstream ss_seq_idx;
|
|
|
|
ss_seq_id_unq << "[ ";
|
|
ss_seq_idx << "[";
|
|
|
|
for (uint32_t s = 0; s < ubatch.n_seqs_unq; ++s) {
|
|
ss_seq_id_unq << ubatch.seq_id_unq[s] << " ";
|
|
}
|
|
|
|
for (uint32_t s = 0; s < LLAMA_MAX_SEQ; ++s) {
|
|
if (ubatch.seq_idx[s] >= 0) {
|
|
ss_seq_idx << ubatch.seq_idx[s]%10;
|
|
} else {
|
|
ss_seq_idx << ".";
|
|
}
|
|
}
|
|
|
|
ss_seq_id_unq << "]";
|
|
ss_seq_idx << "]";
|
|
|
|
LLAMA_LOG_DEBUG("%s: token = %p\n", __func__, (void *) ubatch.token);
|
|
LLAMA_LOG_DEBUG("%s: embd = %p\n", __func__, (void *) ubatch.embd);
|
|
LLAMA_LOG_DEBUG("%s: pos = %p\n", __func__, (void *) ubatch.pos);
|
|
LLAMA_LOG_DEBUG("%s: n_seq_id = %p\n", __func__, (void *) ubatch.n_seq_id);
|
|
LLAMA_LOG_DEBUG("%s: seq_id = %p\n", __func__, (void *) ubatch.seq_id);
|
|
LLAMA_LOG_DEBUG("%s: seq_id_unq = %s\n", __func__, ss_seq_id_unq.str().c_str());
|
|
LLAMA_LOG_DEBUG("%s: seq_idx = %s\n", __func__, ss_seq_idx.str().c_str());
|
|
LLAMA_LOG_DEBUG("%s: output = %p\n", __func__, (void *) ubatch.output);
|
|
LLAMA_LOG_DEBUG("%s: n_outputs = %d\n", __func__, n_outputs);
|
|
|
|
if (debug > 0) {
|
|
int seq_id_max = 0;
|
|
for (uint32_t i = 0; i < ubatch.n_tokens; ++i) {
|
|
for (int s = 0; s < ubatch.n_seq_id[i]; ++s) {
|
|
for (int s = 0; s < ubatch.n_seq_id[i]; ++s) {
|
|
seq_id_max = std::max(seq_id_max, ubatch.seq_id[i][s]);
|
|
}
|
|
}
|
|
}
|
|
++seq_id_max;
|
|
|
|
LLAMA_LOG_DEBUG("%s: token = [\n", __func__);
|
|
for (uint32_t i = 0; i < ubatch.n_tokens; ++i) {
|
|
std::vector<int8_t> seq_id(seq_id_max);
|
|
|
|
for (int s = 0; s < ubatch.n_seq_id[i]; ++s) {
|
|
seq_id[ubatch.seq_id[i][s]] = 1;
|
|
}
|
|
|
|
std::stringstream ss;
|
|
for (int s = 0; s < seq_id_max; ++s) {
|
|
if (seq_id[s]) {
|
|
ss << s%10;
|
|
} else {
|
|
ss << ".";
|
|
}
|
|
}
|
|
|
|
if (ubatch.token) {
|
|
LLAMA_LOG_DEBUG("%s: %4d: id = %6d (%16s), pos = %4d, n_seq_id = %2d, seq_id = [%s], output = %d\n",
|
|
__func__, i, ubatch.token[i], vocab->token_to_piece(ubatch.token[i]).c_str(),
|
|
ubatch.pos[i], ubatch.n_seq_id[i], ss.str().c_str(), ubatch.output[i]);
|
|
} else {
|
|
LLAMA_LOG_DEBUG("%s: %4d: [embd], pos = %4d, n_seq_id = %2d, seq_id = [%s], output = %d\n",
|
|
__func__, i, ubatch.pos[i], ubatch.n_seq_id[i], ss.str().c_str(), ubatch.output[i]);
|
|
}
|
|
}
|
|
LLAMA_LOG_DEBUG("%s: ]\n", __func__);
|
|
}
|
|
}
|
|
}
|
|
|
|
//
|
|
// interface implementation
|
|
//
|
|
|
|
struct llama_batch llama_batch_get_one(
|
|
llama_token * tokens,
|
|
int32_t n_tokens) {
|
|
return {
|
|
/*n_tokens =*/ n_tokens,
|
|
/*tokens =*/ tokens,
|
|
/*embd =*/ nullptr,
|
|
/*pos =*/ nullptr,
|
|
/*n_seq_id =*/ nullptr,
|
|
/*seq_id =*/ nullptr,
|
|
/*logits =*/ nullptr,
|
|
};
|
|
}
|
|
|
|
struct llama_batch llama_batch_init(int32_t n_tokens_alloc, int32_t embd, int32_t n_seq_max) {
|
|
llama_batch batch = {
|
|
/*n_tokens =*/ 0,
|
|
/*tokens =*/ nullptr,
|
|
/*embd =*/ nullptr,
|
|
/*pos =*/ nullptr,
|
|
/*n_seq_id =*/ nullptr,
|
|
/*seq_id =*/ nullptr,
|
|
/*logits =*/ nullptr,
|
|
};
|
|
|
|
if (embd) {
|
|
batch.embd = (float *) malloc(sizeof(float) * n_tokens_alloc * embd);
|
|
} else {
|
|
batch.token = (llama_token *) malloc(sizeof(llama_token) * n_tokens_alloc);
|
|
}
|
|
|
|
batch.pos = (llama_pos *) malloc(sizeof(llama_pos) * n_tokens_alloc);
|
|
batch.n_seq_id = (int32_t *) malloc(sizeof(int32_t) * n_tokens_alloc);
|
|
batch.seq_id = (llama_seq_id **) malloc(sizeof(llama_seq_id *) * (n_tokens_alloc + 1));
|
|
for (int i = 0; i < n_tokens_alloc; ++i) {
|
|
batch.seq_id[i] = (llama_seq_id *) malloc(sizeof(llama_seq_id) * n_seq_max);
|
|
}
|
|
batch.seq_id[n_tokens_alloc] = nullptr;
|
|
|
|
batch.logits = (int8_t *) malloc(sizeof(int8_t) * n_tokens_alloc);
|
|
|
|
return batch;
|
|
}
|
|
|
|
void llama_batch_free(struct llama_batch batch) {
|
|
if (batch.token) free(batch.token);
|
|
if (batch.embd) free(batch.embd);
|
|
if (batch.pos) free(batch.pos);
|
|
if (batch.n_seq_id) free(batch.n_seq_id);
|
|
if (batch.seq_id) {
|
|
for (int i = 0; batch.seq_id[i] != nullptr; ++i) {
|
|
free(batch.seq_id[i]);
|
|
}
|
|
free(batch.seq_id);
|
|
}
|
|
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;
|
|
}
|