Commit 60e9cf7a7 for llama.cpp
commit 60e9cf7a7be81e4a1be7f15de6e8470d0e5b8b29
Author: Xuan-Son Nguyen <son@huggingface.co>
Date: Wed Sep 30 18:08:43 2026 +0200
batch: migrate the rest of examples to llama_batch_ext (#29601)
* migrate the rest
* test-thread-safety
* rm common_batch_staged
diff --git a/common/common.cpp b/common/common.cpp
index 598a97d10..401de1dc2 100644
--- a/common/common.cpp
+++ b/common/common.cpp
@@ -1738,33 +1738,6 @@ void common_threadpools::init(llama_context * ctx, const common_params & params)
llama_attach_threadpool(ctx, threadpool, threadpool_batch);
}
-//
-// Batch utils
-//
-
-void common_batch_clear(struct llama_batch & batch) {
- batch.n_tokens = 0;
-}
-
-void common_batch_add(
- struct llama_batch & batch,
- llama_token id,
- llama_pos pos,
- const std::vector<llama_seq_id> & seq_ids,
- bool logits) {
- GGML_ASSERT(batch.seq_id[batch.n_tokens] && "llama_batch size exceeded");
-
- batch.token [batch.n_tokens] = id;
- batch.pos [batch.n_tokens] = pos;
- batch.n_seq_id[batch.n_tokens] = seq_ids.size();
- for (size_t i = 0; i < seq_ids.size(); ++i) {
- batch.seq_id[batch.n_tokens][i] = seq_ids[i];
- }
- batch.logits [batch.n_tokens] = logits;
-
- batch.n_tokens++;
-}
-
//
// Vocab utils
//
@@ -2118,35 +2091,41 @@ common_batch::common_batch(llama_context * ctx) : batch(llama_batch_ext_init(ctx
void common_batch::clear() {
tokens.clear();
- llama_batch_ext_clear(batch.get());
}
int32_t common_batch::add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output) {
- const int32_t idx = llama_batch_ext_add_token(batch.get(), seq_id, id);
- if (idx < 0) {
- GGML_ABORT("%s: failed to add token %d to the batch (error %d, n_tokens = %d)\n", __func__, id, idx, size());
- }
- llama_batch_ext_set_pos(batch.get(), idx, &pos);
- if (output) {
- llama_batch_ext_set_output_logits(batch.get(), idx, true);
+ tokens.push_back({ id, { pos, 0, 0, 0 }, seq_id, output, { nullptr, 0, 0 }, {} });
+ return size() - 1;
+}
+
+int32_t common_batch::add(llama_token id, llama_pos pos, const std::vector<llama_seq_id> & seq_ids, bool output) {
+ GGML_ASSERT(!seq_ids.empty());
+
+ const int32_t idx = add(id, pos, seq_ids[0], output);
+ for (size_t s = 1; s < seq_ids.size(); ++s) {
+ add_seq(idx, seq_ids[s]);
}
- tokens.push_back({ id, { pos, 0, 0, 0 }, seq_id, output, { nullptr, 0, 0 } });
return idx;
}
+bool common_batch::add_seq(int32_t idx, llama_seq_id seq_id) {
+ if (idx < 0 || idx >= size()) {
+ return false;
+ }
+ tokens[idx].seq_ids_extra.push_back(seq_id);
+ return true;
+}
+
bool common_batch::set_output(int32_t idx, bool value) {
- if (idx < 0 || idx >= (int32_t) tokens.size()) {
+ if (idx < 0 || idx >= size()) {
return false;
}
tokens[idx].output = value;
- return llama_batch_ext_set_output_logits(batch.get(), idx, value);
+ return true;
}
bool common_batch::set_embd(int32_t idx, llama_embd embd) {
- if (idx < 0 || idx >= (int32_t) tokens.size()) {
- return false;
- }
- if (!llama_batch_ext_set_embd_token(batch.get(), idx, embd)) {
+ if (idx < 0 || idx >= size() || tokens[idx].embd.data != nullptr) {
return false;
}
tokens[idx].embd = embd;
@@ -2154,83 +2133,64 @@ bool common_batch::set_embd(int32_t idx, llama_embd embd) {
}
int32_t common_batch::add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output) {
- const int32_t idx = llama_batch_ext_add_embd(batch.get(), seq_id, embd);
- if (idx < 0) {
- GGML_ABORT("%s: failed to add embedding to the batch (error %d, n_tokens = %d)\n", __func__, idx, size());
- }
- llama_batch_ext_set_pos(batch.get(), idx, pos);
- if (output) {
- llama_batch_ext_set_output_logits(batch.get(), idx, true);
- }
- token t = { LLAMA_TOKEN_NULL, { 0, 0, 0, 0 }, seq_id, output, embd };
+ token t = { LLAMA_TOKEN_NULL, { 0, 0, 0, 0 }, seq_id, output, embd, {} };
for (int32_t j = 0; j < n_pos; ++j) {
t.pos[j] = pos[j];
}
tokens.push_back(t);
- return idx;
+ return size() - 1;
}
-common_batch common_batch_from_llama_batch(llama_context * ctx, const llama_batch & batch) {
- common_batch res(ctx);
-
- const bool has_token = batch.token != nullptr;
- const bool has_embd = batch.embd != nullptr;
+llama_batch_ext * common_batch::get_sub_batch(int32_t off, int32_t n) {
+ GGML_ASSERT(batch && "common_batch was not initialized with a context");
+ GGML_ASSERT(off >= 0 && n >= 0 && off + n <= size());
- const size_t n_embd = llama_model_n_embd_inp(llama_get_model(ctx));
+ llama_batch_ext * res = batch.get();
+ llama_batch_ext_clear(res);
- // positions continue from the memory when none are given
- auto * mem = llama_get_memory(ctx);
- std::vector<llama_pos> pos_next(llama_n_seq_max(ctx));
- for (llama_seq_id s = 0; s < (llama_seq_id) pos_next.size(); ++s) {
- pos_next[s] = llama_memory_seq_pos_max(mem, s) + 1;
- }
-
- for (int32_t i = 0; i < batch.n_tokens; ++i) {
- const int32_t n_sid = batch.n_seq_id ? batch.n_seq_id[i] : 1;
- const llama_seq_id seq_id = batch.seq_id ? batch.seq_id[i][0] : 0;
+ for (int32_t i = off; i < off + n; ++i) {
+ const token & t = tokens[i];
- llama_pos pos[GGML_MROPE_SECTIONS] = { 0, 0, 0, 0 };
- if (!batch.pos) {
- pos[0] = pos_next[seq_id]++;
- } else if (has_token) {
- pos[0] = batch.pos[i];
+ int32_t idx;
+ if (t.id != LLAMA_TOKEN_NULL) {
+ idx = llama_batch_ext_add_token(res, t.seq_id, t.id);
+ if (idx < 0) {
+ GGML_ABORT("%s: failed to add token %d at index %d (error %d, n = %d)\n", __func__, t.id, i, idx, n);
+ }
+ llama_batch_ext_set_pos(res, idx, t.pos.data());
+ if (t.embd.data && !llama_batch_ext_set_embd_token(res, idx, t.embd)) {
+ GGML_ABORT("%s: failed to set the embedding of token %d at index %d\n", __func__, t.id, i);
+ }
} else {
- // embedding batch: section-major layout pos[j*n_tokens + i]
- for (int32_t j = 0; j < res.n_pos; ++j) {
- pos[j] = batch.pos[j * batch.n_tokens + i];
+ idx = llama_batch_ext_add_embd(res, t.seq_id, t.embd);
+ if (idx < 0) {
+ GGML_ABORT("%s: failed to add embedding at index %d (error %d, n = %d)\n", __func__, i, idx, n);
}
+ llama_batch_ext_set_pos(res, idx, t.pos.data());
}
+ GGML_ASSERT(idx == i - off);
- const bool output = batch.logits ? batch.logits[i] != 0 : i == batch.n_tokens - 1;
-
- const llama_embd embd = { has_embd ? batch.embd + (size_t) i * n_embd : nullptr, 1, n_embd };
-
- int32_t idx;
- if (has_token) {
- idx = res.add(batch.token[i], pos[0], seq_id, output);
- if (has_embd) {
- res.set_embd(idx, embd);
+ for (const llama_seq_id seq_id : t.seq_ids_extra) {
+ if (!llama_batch_ext_add_seq(res, idx, seq_id)) {
+ GGML_ABORT("%s: failed to add seq %d to the entry at index %d\n", __func__, seq_id, i);
}
- } else {
- idx = res.add_embd(embd, pos, seq_id, output);
}
-
- for (int32_t s = 1; s < n_sid; ++s) {
- llama_batch_ext_add_seq(res.get(), idx, batch.seq_id[i][s]);
+ if (t.output) {
+ llama_batch_ext_set_output_logits(res, idx, true);
}
}
return res;
}
-common_batch common_batch_get_one(llama_context * ctx, const llama_tokens & tokens) {
+common_batch common_batch_get_one(llama_context * ctx, const llama_token * tokens, int32_t n_tokens) {
common_batch batch(ctx);
auto mem = llama_get_memory(ctx);
llama_pos pos = llama_memory_seq_pos_max(mem, 0) + 1; // -1 + 1 == 0 when the memory is empty
- for (size_t i = 0; i < tokens.size(); ++i) {
- const bool output = i == tokens.size() - 1;
+ for (int32_t i = 0; i < n_tokens; ++i) {
+ const bool output = i == n_tokens - 1;
batch.add(tokens[i], pos, 0, output);
pos++;
}
@@ -2238,6 +2198,10 @@ common_batch common_batch_get_one(llama_context * ctx, const llama_tokens & toke
return batch;
}
+common_batch common_batch_get_one(llama_context * ctx, const llama_tokens & tokens) {
+ return common_batch_get_one(ctx, tokens.data(), (int32_t) tokens.size());
+}
+
bool common_prompt_batch_decode(
struct llama_context * ctx,
const llama_tokens & all_tokens,
diff --git a/common/common.h b/common/common.h
index e95eb2fd0..dcc5ec1ae 100644
--- a/common/common.h
+++ b/common/common.h
@@ -1031,23 +1031,16 @@ struct common_memory {
// Batch utils
//
-void common_batch_clear(struct llama_batch & batch);
-
-void common_batch_add(
- struct llama_batch & batch,
- llama_token id,
- llama_pos pos,
- const std::vector<llama_seq_id> & seq_ids,
- bool logits);
-
// wrapper around llama_batch_ext that provide getter functions for downstream code
+// entries can exceed n_batch, use get_sub_batch() to decode them in chunks
struct common_batch {
struct token {
llama_token id;
std::array<llama_pos, GGML_MROPE_SECTIONS> pos; // only pos[0] is used for text tokens
- llama_seq_id seq_id;
+ llama_seq_id seq_id; // the first sequence id, see add_seq()
bool output;
llama_embd embd; // non-owning view of the data passed to add_embd()/set_embd(), data == NULL if none
+ std::vector<llama_seq_id> seq_ids_extra; // see add_seq()
};
std::vector<token> tokens; // mirror of the entries, tokens[i] describes batch index i
@@ -1058,7 +1051,10 @@ struct common_batch {
common_batch() = default;
common_batch(struct llama_context * ctx);
- llama_batch_ext * get() const { return batch.get(); }
+ llama_batch_ext * get() { return get_sub_batch(0, size()); }
+
+ // render entries [off, off + n) into batch, the result is overwritten by the next call
+ llama_batch_ext * get_sub_batch(int32_t off, int32_t n);
// content type of the batch, all entries carry the same combination
bool has_token() const { return !tokens.empty() && tokens[0].id != LLAMA_TOKEN_NULL; }
@@ -1066,15 +1062,21 @@ struct common_batch {
void clear();
- // returns the batch index (>= 0), aborts if the entry cannot be added (batch full, invalid token or seq id)
+ // returns the batch index
int32_t add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output);
+ // same, with the entry shared by all seq_ids (must not be empty)
+ int32_t add(llama_token id, llama_pos pos, const std::vector<llama_seq_id> & seq_ids, bool output);
+
+ // add the entry at idx to another sequence, tokens[idx].seq_id keeps the first one
+ bool add_seq(int32_t idx, llama_seq_id seq_id);
+
bool set_output(int32_t idx, bool value);
// attach a token embedding to the entry at idx, can only be set once per entry
bool set_embd(int32_t idx, llama_embd embd);
- // add an embedding-only entry (no token id), aborts like add() on failure
+ // add an embedding-only entry (no token id)
// pos points to n_pos positions
int32_t add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output);
@@ -1082,13 +1084,10 @@ struct common_batch {
};
// create a single-sequence batch from a list of tokens
-// last token always have output_logits set to true
+// positions continue from the memory, last token always have output_logits set to true
+common_batch common_batch_get_one(struct llama_context * ctx, const llama_token * tokens, int32_t n_tokens);
common_batch common_batch_get_one(struct llama_context * ctx, const llama_tokens & tokens);
-// convert a legacy llama_batch, applying its defaults: seq 0, positions continue from memory, last token is output
-// the embd rows are read at the model input width
-common_batch common_batch_from_llama_batch(struct llama_context * ctx, const llama_batch & batch);
-
// decodes a single batch of tokens for a prompt and manages session tokens
//
// Note: We save state before the last token so that we can replay it to ensure
diff --git a/common/speculative.cpp b/common/speculative.cpp
index 82e9e9223..b1244d9a6 100644
--- a/common/speculative.cpp
+++ b/common/speculative.cpp
@@ -2163,9 +2163,6 @@ struct common_speculative_impl_ngram_cache : public common_speculative_impl {
struct common_speculative {
common_speculative_draft_params_vec dparams;
- // the target context, used to convert legacy llama_batch inputs
- llama_context * ctx_tgt = nullptr;
-
// list of implementations to use and their states
std::vector<std::unique_ptr<common_speculative_impl>> impls;
@@ -2711,7 +2708,6 @@ common_speculative * common_speculative_init(common_params_speculative & params,
common_speculative_ptr result(new common_speculative {
/* .dparams = */ common_speculative_draft_params_vec(n_seq),
- /* .ctx_tgt = */ params.draft.ctx_tgt,
/* .impls = */ std::move(impls),
/* .impl_last = */ std::vector<common_speculative_impl *>(n_seq, nullptr),
/* .synth_probs = */ {},
@@ -2774,17 +2770,6 @@ void common_speculative_begin(common_speculative * spec, llama_seq_id seq_id, co
}
}
-bool common_speculative_process(common_speculative * spec, const llama_batch & batch) {
- if (spec == nullptr) {
- return true;
- }
-
- // ngram-only setups have no target context, they do not read the batch anyway
- const common_batch tmp = spec->ctx_tgt ? common_batch_from_llama_batch(spec->ctx_tgt, batch) : common_batch();
-
- return common_speculative_process(spec, tmp);
-}
-
bool common_speculative_process(common_speculative * spec, const common_batch & batch) {
bool result = true;
diff --git a/common/speculative.h b/common/speculative.h
index 211fcdabd..d46b21eb7 100644
--- a/common/speculative.h
+++ b/common/speculative.h
@@ -79,9 +79,6 @@ void common_speculative_begin(common_speculative * spec, llama_seq_id seq_id, co
// process the batch and update the internal state of the speculative context
bool common_speculative_process(common_speculative * spec, const common_batch & batch);
-// legacy llama_batch input, converted with common_batch_from_llama_batch()
-bool common_speculative_process(common_speculative * spec, const llama_batch & batch);
-
// generate drafts for the sequences specified with `common_speculative_get_draft_params`
void common_speculative_draft(common_speculative * spec);
diff --git a/examples/batched/batched.cpp b/examples/batched/batched.cpp
index 830e45f5a..9da844c7e 100644
--- a/examples/batched/batched.cpp
+++ b/examples/batched/batched.cpp
@@ -117,7 +117,7 @@ int main(int argc, char ** argv) {
// create a llama_batch
// we use this object to submit token data for decoding
- llama_batch batch = llama_batch_init(std::max(tokens_list.size(), (size_t) n_parallel), 0, n_parallel);
+ common_batch batch(ctx);
std::vector<llama_seq_id> seq_ids(n_parallel, 0);
for (int32_t i = 0; i < n_parallel; ++i) {
@@ -126,12 +126,12 @@ int main(int argc, char ** argv) {
// evaluate the initial prompt
for (size_t i = 0; i < tokens_list.size(); ++i) {
- common_batch_add(batch, tokens_list[i], i, seq_ids, false);
+ batch.add(tokens_list[i], i, seq_ids, false);
}
- GGML_ASSERT(batch.n_tokens == (int) tokens_list.size());
+ GGML_ASSERT(batch.size() == (int) tokens_list.size());
if (llama_model_has_encoder(model)) {
- if (llama_encode(ctx, batch)) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_ENCODE, batch.get())) {
LOG_ERR("%s : failed to eval\n", __func__);
return 1;
}
@@ -141,14 +141,14 @@ int main(int argc, char ** argv) {
decoder_start_token_id = llama_vocab_bos(vocab);
}
- common_batch_clear(batch);
- common_batch_add(batch, decoder_start_token_id, 0, seq_ids, false);
+ batch.clear();
+ batch.add(decoder_start_token_id, 0, seq_ids, false);
}
// llama_decode will output logits only for the last token of the prompt
- batch.logits[batch.n_tokens - 1] = true;
+ batch.set_output(batch.size() - 1, true);
- if (llama_decode(ctx, batch) != 0) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
LOG_ERR("%s: llama_decode() failed\n", __func__);
return 1;
}
@@ -170,16 +170,16 @@ int main(int argc, char ** argv) {
// remember the batch index of the last token for each parallel sequence
// we need this to determine which logits to sample from
- std::vector<int32_t> i_batch(n_parallel, batch.n_tokens - 1);
+ std::vector<int32_t> i_batch(n_parallel, batch.size() - 1);
- int n_cur = batch.n_tokens;
+ int n_cur = batch.size();
int n_decode = 0;
const auto t_main_start = ggml_time_us();
while (n_cur <= n_predict) {
// prepare the next batch
- common_batch_clear(batch);
+ batch.clear();
// sample the next token for each parallel sequence / stream
for (int32_t i = 0; i < n_parallel; ++i) {
@@ -208,23 +208,23 @@ int main(int argc, char ** argv) {
streams[i] += common_token_to_piece(ctx, new_token_id);
- i_batch[i] = batch.n_tokens;
+ i_batch[i] = batch.size();
// push this new token for next evaluation
- common_batch_add(batch, new_token_id, n_cur, { i }, true);
+ batch.add(new_token_id, n_cur, i, true);
n_decode += 1;
}
// all streams are finished
- if (batch.n_tokens == 0) {
+ if (batch.size() == 0) {
break;
}
n_cur += 1;
// evaluate the current batch with the transformer model
- if (llama_decode(ctx, batch)) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("%s : failed to eval, return code %d\n", __func__, 1);
return 1;
}
@@ -249,7 +249,6 @@ int main(int argc, char ** argv) {
fprintf(stderr, "\n");
- llama_batch_free(batch);
for (auto & sampler_config : sampler_configs) {
llama_sampler_free(sampler_config.sampler);
diff --git a/examples/debug/debug.cpp b/examples/debug/debug.cpp
index 761e7a2db..18a264b63 100644
--- a/examples/debug/debug.cpp
+++ b/examples/debug/debug.cpp
@@ -194,7 +194,8 @@ static bool run(llama_context * ctx, const common_params & params) {
return false;
}
- if (llama_decode(ctx, llama_batch_get_one(tokens.data(), tokens.size()))) {
+ common_batch batch = common_batch_get_one(ctx, tokens);
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("%s : failed to eval\n", __func__);
return false;
}
diff --git a/examples/diffusion/diffusion.cpp b/examples/diffusion/diffusion.cpp
index 97d6b6944..0ff57d953 100644
--- a/examples/diffusion/diffusion.cpp
+++ b/examples/diffusion/diffusion.cpp
@@ -1,5 +1,7 @@
#include "diffusion.h"
+#include "common.h"
+
#include "log.h"
#include <algorithm>
@@ -144,8 +146,7 @@ void diffusion_generate(llama_context * ctx,
struct llama_sampler * dist_sampler = llama_sampler_init_dist(params.seed);
- llama_batch batch = llama_batch_init(params.max_length, 0, 1);
- batch.n_tokens = params.max_length;
+ common_batch batch(ctx);
// Pre-allocate buffers for CFG if needed
int32_t logits_size = n_vocab * params.max_length;
@@ -202,18 +203,15 @@ void diffusion_generate(llama_context * ctx,
}
// Setup batch
+ batch.clear();
for (int32_t i = 0; i < params.max_length; i++) {
- batch.token[i] = output_tokens[i];
- batch.pos[i] = i;
- batch.n_seq_id[i] = 1;
- batch.seq_id[i][0] = 0;
- batch.logits[i] = 1;
+ batch.add(output_tokens[i], i, 0, true);
}
float * logits = nullptr;
if (params.cfg_scale > 0.0f) {
- int ret = llama_decode(ctx, batch);
+ int ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
LOG_ERR("Failed to generate conditional");
break;
@@ -227,10 +225,11 @@ void diffusion_generate(llama_context * ctx,
un_x_buffer[i] = params.mask_token_id;
}
+ batch.clear();
for (int32_t i = 0; i < params.max_length; i++) {
- batch.token[i] = un_x_buffer[i];
+ batch.add(un_x_buffer[i], i, 0, true);
}
- ret = llama_decode(ctx, batch);
+ ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
LOG_ERR("Failed to generate unconditional");
break;
@@ -244,7 +243,7 @@ void diffusion_generate(llama_context * ctx,
}
logits = cond_logits_buffer.data();
} else {
- int ret = llama_decode(ctx, batch);
+ int ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret != 0) {
LOG_ERR("%s: failed to decode at step %d, ret = %d\n", __func__, global_step, ret);
break;
@@ -400,7 +399,6 @@ void diffusion_generate(llama_context * ctx,
total_time / 1000.0 / params.steps,
total_sampling_time / 1000.0 / params.steps);
- llama_batch_free(batch);
llama_sampler_free(sampler);
llama_sampler_free(dist_sampler);
diff --git a/examples/embedding/embedding.cpp b/examples/embedding/embedding.cpp
index f6a20ef9d..a59f04cae 100644
--- a/examples/embedding/embedding.cpp
+++ b/examples/embedding/embedding.cpp
@@ -27,27 +27,27 @@ static std::vector<std::string> split_lines(const std::string & s, const std::st
return lines;
}
-static void batch_add_seq(llama_batch & batch, const std::vector<int32_t> & tokens, llama_seq_id seq_id) {
+static void batch_add_seq(common_batch & batch, const std::vector<int32_t> & tokens, llama_seq_id seq_id) {
size_t n_tokens = tokens.size();
for (size_t i = 0; i < n_tokens; i++) {
- common_batch_add(batch, tokens[i], i, { seq_id }, true);
+ batch.add(tokens[i], i, seq_id, true);
}
}
-static void batch_decode(llama_context * ctx, llama_batch & batch, float * output, int n_seq, int n_embd_out, int embd_norm) {
+static void batch_decode(llama_context * ctx, common_batch & batch, float * output, int n_seq, int n_embd_out, int embd_norm) {
const enum llama_pooling_type pooling_type = llama_pooling_type(ctx);
// clear previous kv_cache values (irrelevant for embeddings)
llama_memory_clear(llama_get_memory(ctx), true);
// run model
- LOG_INF("%s: n_tokens = %d, n_seq = %d\n", __func__, batch.n_tokens, n_seq);
- if (llama_decode(ctx, batch) < 0) {
+ LOG_INF("%s: n_tokens = %d, n_seq = %d\n", __func__, batch.size(), n_seq);
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) < 0) {
LOG_ERR("%s : failed to process\n", __func__);
}
- for (int i = 0; i < batch.n_tokens; i++) {
- if (!batch.logits[i]) {
+ for (int i = 0; i < batch.size(); i++) {
+ if (!batch.tokens[i].output) {
continue;
}
@@ -61,8 +61,8 @@ static void batch_decode(llama_context * ctx, llama_batch & batch, float * outpu
GGML_ASSERT(embd != NULL && "failed to get token embeddings");
} else {
// try to get sequence embeddings - supported only when pooling_type is not NONE
- embd = llama_get_embeddings_seq(ctx, batch.seq_id[i][0]);
- embd_pos = batch.seq_id[i][0];
+ embd = llama_get_embeddings_seq(ctx, batch.tokens[i].seq_id);
+ embd_pos = batch.tokens[i].seq_id;
GGML_ASSERT(embd != NULL && "failed to get sequence embeddings");
}
@@ -242,7 +242,7 @@ int main(int argc, char ** argv) {
// initialize batch
const int n_prompts = prompts.size();
- struct llama_batch batch = llama_batch_init(n_batch, 0, 1);
+ common_batch batch(ctx);
// count number of embeddings
int n_embd_count = 0;
@@ -269,12 +269,12 @@ int main(int argc, char ** argv) {
const uint64_t n_toks = inp.size();
// encode if at capacity
- if (batch.n_tokens + n_toks > n_batch || s >= n_seq_max) {
+ if (batch.size() + n_toks > n_batch || s >= n_seq_max) {
float * out = emb + e * n_embd_out;
batch_decode(ctx, batch, out, s, n_embd_out, params.embd_normalize);
- e += pooling_type == LLAMA_POOLING_TYPE_NONE ? batch.n_tokens : s;
+ e += pooling_type == LLAMA_POOLING_TYPE_NONE ? batch.size() : s;
s = 0;
- common_batch_clear(batch);
+ batch.clear();
}
// add to batch
@@ -407,7 +407,6 @@ int main(int argc, char ** argv) {
llama_perf_context_print(ctx);
// clean up
- llama_batch_free(batch);
llama_backend_free();
return 0;
diff --git a/examples/eval-callback/eval-callback.cpp b/examples/eval-callback/eval-callback.cpp
index 4ce8d600b..703ce130b 100644
--- a/examples/eval-callback/eval-callback.cpp
+++ b/examples/eval-callback/eval-callback.cpp
@@ -26,7 +26,8 @@ static bool run(llama_context * ctx, const common_params & params) {
LOG_INF(" %d\n", tokens[i]);
}
- if (llama_decode(ctx, llama_batch_get_one(tokens.data(), tokens.size()))) {
+ common_batch batch = common_batch_get_one(ctx, tokens);
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("%s : failed to eval\n", __func__);
return false;
}
diff --git a/examples/idle/idle.cpp b/examples/idle/idle.cpp
index 409fd25c1..ddbda7993 100644
--- a/examples/idle/idle.cpp
+++ b/examples/idle/idle.cpp
@@ -57,12 +57,13 @@ int main(int argc, char ** argv) {
return 1;
}
- llama_batch batch = llama_batch_get_one(prompt_tokens.data(), prompt_tokens.size());
-
const int n_iters = 3;
// warm-up
- llama_decode(ctx, batch);
+ {
+ common_batch batch = common_batch_get_one(ctx, prompt_tokens);
+ llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+ }
llama_memory_clear(llama_get_memory(ctx), true);
llama_synchronize(ctx);
@@ -71,13 +72,16 @@ int main(int argc, char ** argv) {
double t_sum2_us = 0.0;
for (int i = 0; i < n_iters; i++) {
+ // positions continue from the memory
+ common_batch batch = common_batch_get_one(ctx, prompt_tokens);
+
// this pause is important - it simulates "idle GPU"
std::this_thread::sleep_for(std::chrono::milliseconds(t_pause_ms));
const int64_t t_start_us = llama_time_us();
// this should take constant time
- llama_decode(ctx, batch);
+ llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
llama_synchronize(ctx);
const int64_t t_end_us = llama_time_us();
diff --git a/examples/llama.android/lib/src/main/cpp/ai_chat.cpp b/examples/llama.android/lib/src/main/cpp/ai_chat.cpp
index 03ab96cfd..ca6f7172d 100644
--- a/examples/llama.android/lib/src/main/cpp/ai_chat.cpp
+++ b/examples/llama.android/lib/src/main/cpp/ai_chat.cpp
@@ -35,7 +35,7 @@ constexpr float DEFAULT_SAMPLER_TEMP = 0.3f;
static llama_model * g_model;
static llama_context * g_context;
-static llama_batch g_batch;
+static common_batch g_batch;
static common_chat_templates_ptr g_chat_templates;
static common_sampler * g_sampler;
@@ -116,7 +116,7 @@ Java_com_arm_aichat_internal_InferenceEngineImpl_prepare(JNIEnv * /*env*/, jobje
auto *context = init_context(g_model);
if (!context) { return 1; }
g_context = context;
- g_batch = llama_batch_init(BATCH_SIZE, 0, 1);
+ g_batch = common_batch(context);
g_chat_templates = common_chat_templates_init(g_model, "");
g_sampler = new_sampler(DEFAULT_SAMPLER_TEMP);
return 0;
@@ -164,18 +164,18 @@ Java_com_arm_aichat_internal_InferenceEngineImpl_benchModel(JNIEnv *env, jobject
for (nri = 0; nri < nr; nri++) {
LOGi("Benchmark prompt processing (pp = %d)", pp);
- common_batch_clear(g_batch);
+ common_batch batch(context);
const int n_tokens = pp;
for (i = 0; i < n_tokens; i++) {
- common_batch_add(g_batch, 0, i, {0}, false);
+ batch.add(0, i, 0, false);
}
- g_batch.logits[g_batch.n_tokens - 1] = true;
+ batch.set_output(batch.size() - 1, true);
llama_memory_clear(llama_get_memory(context), false);
const auto t_pp_start = ggml_time_us();
- if (llama_decode(context, g_batch) != 0) {
+ if (llama_process(context, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
LOGe("llama_decode() failed during prompt processing");
}
const auto t_pp_end = ggml_time_us();
@@ -187,12 +187,12 @@ Java_com_arm_aichat_internal_InferenceEngineImpl_benchModel(JNIEnv *env, jobject
llama_memory_clear(llama_get_memory(context), false);
const auto t_tg_start = ggml_time_us();
for (i = 0; i < tg; i++) {
- common_batch_clear(g_batch);
+ batch.clear();
for (j = 0; j < pl; j++) {
- common_batch_add(g_batch, 0, i, {j}, true);
+ batch.add(0, i, j, true);
}
- if (llama_decode(context, g_batch) != 0) {
+ if (llama_process(context, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
LOGe("llama_decode() failed during text generation");
}
}
@@ -315,7 +315,7 @@ static void reset_short_term_states() {
static int decode_tokens_in_batches(
llama_context *context,
- llama_batch &batch,
+ common_batch &batch,
const llama_tokens &tokens,
const llama_pos start_pos,
const bool compute_last_logit = false) {
@@ -323,7 +323,7 @@ static int decode_tokens_in_batches(
LOGd("%s: Decode %d tokens starting at position %d", __func__, (int) tokens.size(), start_pos);
for (int i = 0; i < (int) tokens.size(); i += BATCH_SIZE) {
const int cur_batch_size = std::min((int) tokens.size() - i, BATCH_SIZE);
- common_batch_clear(batch);
+ batch.clear();
LOGv("%s: Preparing a batch size of %d starting at: %d", __func__, cur_batch_size, i);
// Shift context if current batch cannot fit into the context
@@ -337,11 +337,11 @@ static int decode_tokens_in_batches(
const llama_token token_id = tokens[i + j];
const llama_pos position = start_pos + i + j;
const bool want_logit = compute_last_logit && (i + j == tokens.size() - 1);
- common_batch_add(batch, token_id, position, {0}, want_logit);
+ batch.add(token_id, position, 0, want_logit);
}
// Decode this batch
- const int decode_result = llama_decode(context, batch);
+ const int decode_result = llama_process(context, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (decode_result) {
LOGe("%s: llama_decode failed w/ %d", __func__, decode_result);
return 1;
@@ -506,9 +506,9 @@ Java_com_arm_aichat_internal_InferenceEngineImpl_generateNextToken(
common_sampler_accept(g_sampler, new_token_id, true);
// Populate the batch with new token, then decode
- common_batch_clear(g_batch);
- common_batch_add(g_batch, new_token_id, current_position, {0}, true);
- if (llama_decode(g_context, g_batch) != 0) {
+ g_batch.clear();
+ g_batch.add(new_token_id, current_position, 0, true);
+ if (llama_process(g_context, LLAMA_PROCESS_TYPE_DECODE, g_batch.get()) != 0) {
LOGe("%s: llama_decode() failed for generated token", __func__);
return nullptr;
}
@@ -553,7 +553,7 @@ Java_com_arm_aichat_internal_InferenceEngineImpl_unload(JNIEnv * /*unused*/, job
// Free up resources
common_sampler_free(g_sampler);
g_chat_templates.reset();
- llama_batch_free(g_batch);
+ g_batch = common_batch();
llama_free(g_context);
llama_model_free(g_model);
}
diff --git a/examples/lookahead/lookahead.cpp b/examples/lookahead/lookahead.cpp
index b7f5c6de8..62772814f 100644
--- a/examples/lookahead/lookahead.cpp
+++ b/examples/lookahead/lookahead.cpp
@@ -101,8 +101,13 @@ int main(int argc, char ** argv) {
const auto t_enc_start = ggml_time_us();
// eval the prompt
- llama_decode(ctx, llama_batch_get_one( inp.data(), n_input - 1));
- llama_decode(ctx, llama_batch_get_one(&inp.back(), 1));
+ {
+ common_batch batch = common_batch_get_one(ctx, inp.data(), n_input - 1);
+ llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+
+ batch = common_batch_get_one(ctx, &inp.back(), 1);
+ llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+ }
for (int s = 1; s < W + G + 1; ++s) {
llama_memory_seq_cp(mem, 0, s, -1, -1);
@@ -124,7 +129,7 @@ int main(int argc, char ** argv) {
// seq_id == 0 : the current input token
// seq_id [1, W] : tokens from the past N - 1 Jacobi iterations
// seq_id [W + 1, W + G] : verification n-grams
- llama_batch batch = llama_batch_init(llama_n_ctx(ctx), 0, W + G + 1);
+ common_batch batch(ctx);
// target model sampling context
struct common_sampler * smpl = common_sampler_init(model, params.sampling);
@@ -204,10 +209,10 @@ int main(int argc, char ** argv) {
// V V V V V V
// id
{
- common_batch_clear(batch);
+ batch.clear();
// current token - first token of the first level
- common_batch_add(batch, id, n_past, seq_id_all, true);
+ batch.add(id, n_past, seq_id_all, true);
// verification n-grams - queue this before the lookahead tokens for less KV cache fragmentation
{
@@ -230,9 +235,9 @@ int main(int argc, char ** argv) {
const llama_token t = ngrams_observed.tokens[idx + j];
ngrams_cur[g].tokens [j + 1] = t;
- ngrams_cur[g].i_batch[j + 1] = batch.n_tokens;
+ ngrams_cur[g].i_batch[j + 1] = batch.size();
- common_batch_add(batch, t, n_past + j + 1, { W + 1 + g }, true);
+ batch.add(t, n_past + j + 1, W + 1 + g, true);
}
}
}
@@ -244,18 +249,18 @@ int main(int argc, char ** argv) {
seq_id_look[j] = i + j + 1;
}
- common_batch_add(batch, tokens_j[0][i], n_past + i, seq_id_look, false);
+ batch.add(tokens_j[0][i], n_past + i, seq_id_look, false);
}
// fill the rest of the levels
for (int j = 1; j < N - 1; j++) {
for (int i = 0; i < W; i++) {
- common_batch_add(batch, tokens_j[j][i], n_past + j + i, { i + 1 }, j == N - 2);
+ batch.add(tokens_j[j][i], n_past + j + i, i + 1, j == N - 2);
}
}
}
- if (llama_decode(ctx, batch) != 0) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
LOG_ERR("\n\n%s: llama_decode failed - increase KV cache size\n", __func__);
return 1;
}
@@ -473,7 +478,6 @@ int main(int argc, char ** argv) {
common_sampler_free(smpl);
- llama_batch_free(batch);
llama_backend_free();
diff --git a/examples/lookup/lookup.cpp b/examples/lookup/lookup.cpp
index 662105865..004062c57 100644
--- a/examples/lookup/lookup.cpp
+++ b/examples/lookup/lookup.cpp
@@ -98,8 +98,13 @@ int main(int argc, char ** argv){
const auto t_enc_start = ggml_time_us();
- llama_decode(ctx, llama_batch_get_one( inp.data(), n_input - 1));
- llama_decode(ctx, llama_batch_get_one(&inp.back(), 1));
+ {
+ common_batch batch = common_batch_get_one(ctx, inp.data(), n_input - 1);
+ llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+
+ batch = common_batch_get_one(ctx, &inp.back(), 1);
+ llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+ }
const auto t_enc_end = ggml_time_us();
@@ -115,7 +120,7 @@ int main(int argc, char ** argv){
std::vector<llama_token> draft;
- llama_batch batch_tgt = llama_batch_init(llama_n_ctx(ctx), 0, 1);
+ common_batch batch_tgt(ctx);
const auto t_dec_start = ggml_time_us();
@@ -192,8 +197,8 @@ int main(int argc, char ** argv){
// clean the cache of draft tokens that weren't accepted
llama_memory_seq_rm(llama_get_memory(ctx), 0, n_past, -1);
- common_batch_clear(batch_tgt);
- common_batch_add(batch_tgt, draft[0], n_past, { 0 }, true);
+ batch_tgt.clear();
+ batch_tgt.add(draft[0], n_past, 0, true);
// Draft already contains a single token sampled from the model:
GGML_ASSERT(draft.size() == 1);
@@ -203,13 +208,13 @@ int main(int argc, char ** argv){
common_ngram_cache_draft(inp, draft, n_draft, LLAMA_NGRAM_MIN, LLAMA_NGRAM_MAX, ngram_cache_context, ngram_cache_dynamic, ngram_cache_static);
for (size_t i = 1; i < draft.size(); ++i) {
- common_batch_add(batch_tgt, draft[i], n_past + i, { 0 }, true);
+ batch_tgt.add(draft[i], n_past + i, 0, true);
}
t_draft_us += ggml_time_us() - t_start_draft_us;
n_drafted += draft.size() - 1;
- llama_decode(ctx, batch_tgt);
+ llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_tgt.get());
++n_past;
draft.erase(draft.begin());
@@ -241,7 +246,6 @@ int main(int argc, char ** argv){
common_sampler_free(smpl);
- llama_batch_free(batch_tgt);
llama_backend_free();
diff --git a/examples/parallel/parallel.cpp b/examples/parallel/parallel.cpp
index 4b74540f0..de877ca36 100644
--- a/examples/parallel/parallel.cpp
+++ b/examples/parallel/parallel.cpp
@@ -224,8 +224,6 @@ int main(int argc, char ** argv) {
LOG_INF("\n\n");
- const int n_ctx = llama_n_ctx(ctx);
-
if (sseed >= 0) {
LOG_INF("%s: initializing all samplers with the same RNG seed: %d (use a negative seed to have different seeds)\n", __func__, sseed);
} else {
@@ -252,7 +250,7 @@ int main(int argc, char ** argv) {
// the max batch size is as large as the context to handle cases where we get very long input prompt from multiple
// users. regardless of the size, the main loop will chunk the batch into a maximum of params.n_batch tokens at a time
- llama_batch batch = llama_batch_init(n_ctx, 0, 1);
+ common_batch batch(ctx);
int32_t n_total_prompt = 0;
int32_t n_total_gen = 0;
@@ -268,10 +266,10 @@ int main(int argc, char ** argv) {
LOG_INF("%s: Evaluating the system prompt ...\n", __func__);
for (int32_t i = 0; i < n_tokens_system; ++i) {
- common_batch_add(batch, tokens_system[i], i, { 0 }, false);
+ batch.add(tokens_system[i], i, 0, false);
}
- if (llama_decode(ctx, batch) != 0) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
LOG_ERR("%s: llama_decode() failed\n", __func__);
return 1;
}
@@ -287,7 +285,7 @@ int main(int argc, char ** argv) {
LOG_INF("Processing requests ...\n\n");
while (true) {
- common_batch_clear(batch);
+ batch.clear();
// decode any currently ongoing sequences
for (auto & client : clients) {
@@ -295,14 +293,14 @@ int main(int argc, char ** argv) {
continue;
}
- client.i_batch = batch.n_tokens;
+ client.i_batch = batch.size();
- common_batch_add(batch, client.sampled, client.n_past++, { client.id + 1 }, true);
+ batch.add(client.sampled, client.n_past++, client.id + 1, true);
client.n_decoded += 1;
}
- if (batch.n_tokens == 0) {
+ if (batch.size() == 0) {
// all sequences have ended - clear the entire KV cache
for (int i = 1; i <= n_clients; ++i) {
llama_memory_seq_rm(mem, i, -1, -1);
@@ -314,7 +312,7 @@ int main(int argc, char ** argv) {
}
// insert new sequences for decoding
- if (cont_batching || batch.n_tokens == 0) {
+ if (cont_batching || batch.size() == 0) {
for (auto & client : clients) {
if (client.seq_id == -1 && g_seq_id < n_seq) {
client.seq_id = g_seq_id;
@@ -350,17 +348,17 @@ int main(int argc, char ** argv) {
tokens_prompt = common_tokenize(ctx, client.prompt, false);
for (size_t i = 0; i < tokens_prompt.size(); ++i) {
- common_batch_add(batch, tokens_prompt[i], client.n_past++, { client.id + 1 }, false);
+ batch.add(tokens_prompt[i], client.n_past++, client.id + 1, false);
}
// extract the logits only for the last token
- if (batch.n_tokens > 0) {
- batch.logits[batch.n_tokens - 1] = true;
+ if (batch.size() > 0) {
+ batch.set_output(batch.size() - 1, true);
}
client.n_prompt = tokens_prompt.size();
client.n_decoded = 0;
- client.i_batch = batch.n_tokens - 1;
+ client.i_batch = batch.size() - 1;
LOG_INF("\033[31mClient %3d, seq %4d, junk = %4d, prompt = %d, started decoding ...\033[0m\n", client.id, client.seq_id, n_junk_cur, client.n_prompt);
@@ -374,7 +372,7 @@ int main(int argc, char ** argv) {
}
}
- if (batch.n_tokens == 0) {
+ if (batch.size() == 0) {
break;
}
@@ -383,27 +381,17 @@ int main(int argc, char ** argv) {
int32_t i_next = 0;
- for (int32_t i = 0; i < batch.n_tokens; i = i_next) {
+ for (int32_t i = 0; i < batch.size(); i = i_next) {
// experiment: process in powers of 2
- //if (i + n_batch > (int32_t) batch.n_tokens && n_batch > 32) {
+ //if (i + n_batch > (int32_t) batch.size() && n_batch > 32) {
// n_batch /= 2;
// i -= n_batch;
// continue;
//}
- const int32_t n_tokens = std::min(n_batch, batch.n_tokens - i);
-
- llama_batch batch_view = {
- n_tokens,
- batch.token + i,
- nullptr,
- batch.pos + i,
- batch.n_seq_id + i,
- batch.seq_id + i,
- batch.logits + i,
- };
+ const int32_t n_tokens = std::min(n_batch, batch.size() - i);
- const int ret = llama_decode(ctx, batch_view);
+ const int ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get_sub_batch(i, n_tokens));
if (ret != 0) {
if (n_batch == 1 || ret < 0) {
// if you get here, it means the KV cache is full - try increasing it via the context size
@@ -511,7 +499,6 @@ int main(int argc, char ** argv) {
// TODO: print sampling/grammar timings for all clients
llama_perf_context_print(ctx);
- llama_batch_free(batch);
llama_backend_free();
diff --git a/examples/passkey/passkey.cpp b/examples/passkey/passkey.cpp
index 8440a2bf7..9ac8a0170 100644
--- a/examples/passkey/passkey.cpp
+++ b/examples/passkey/passkey.cpp
@@ -125,7 +125,7 @@ int main(int argc, char ** argv) {
LOG_INF("prompt tokens: %d\n", n_tokens_all);
//LOG_INF("prompt: %s\n", params.prompt.c_str());
- llama_batch batch = llama_batch_init(params.n_batch, 0, 1);
+ common_batch batch(ctx);
int n_past = 0;
@@ -144,17 +144,17 @@ int main(int argc, char ** argv) {
n_past = llama_memory_seq_pos_max(mem, 0) + 1;
}
- common_batch_clear(batch);
+ batch.clear();
for (int j = 0; j < n_batch && i + j < n_tokens_all; j++) {
- common_batch_add(batch, tokens_list[i + j], n_past++, { 0 }, false);
+ batch.add(tokens_list[i + j], n_past++, 0, false);
}
if (i + n_batch >= n_tokens_all) {
- batch.logits[batch.n_tokens - 1] = true;
+ batch.set_output(batch.size() - 1, true);
}
- if (llama_decode(ctx, batch) != 0) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
LOG_INF("%s: llama_decode() failed\n", __func__);
return 1;
}
@@ -176,17 +176,17 @@ int main(int argc, char ** argv) {
n_past = llama_memory_seq_pos_max(mem, 0) + 1;
- common_batch_clear(batch);
+ batch.clear();
for (int j = 0; j < n_batch && i + j < n_tokens_all; j++) {
- common_batch_add(batch, tokens_list[i + j], n_past++, { 0 }, false);
+ batch.add(tokens_list[i + j], n_past++, 0, false);
}
if (i + n_batch >= n_tokens_all) {
- batch.logits[batch.n_tokens - 1] = true;
+ batch.set_output(batch.size() - 1, true);
}
- if (llama_decode(ctx, batch) != 0) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
LOG_ERR("%s: llama_decode() failed\n", __func__);
return 1;
}
@@ -223,7 +223,7 @@ int main(int argc, char ** argv) {
while (n_cur <= n_len) {
// sample the next token
{
- const llama_token new_token_id = llama_sampler_sample(smpl, ctx, batch.n_tokens - 1);
+ const llama_token new_token_id = llama_sampler_sample(smpl, ctx, batch.size() - 1);
// is it an end of generation?
if (llama_vocab_is_eog(vocab, new_token_id) || n_cur == n_len) {
@@ -237,16 +237,16 @@ int main(int argc, char ** argv) {
n_decode += 1;
// prepare the next batch
- common_batch_clear(batch);
+ batch.clear();
// push this new token for next evaluation
- common_batch_add(batch, new_token_id, n_past++, { 0 }, true);
+ batch.add(new_token_id, n_past++, 0, true);
}
n_cur += 1;
// evaluate the current batch with the transformer model
- if (llama_decode(ctx, batch)) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("%s : failed to eval, return code %d\n", __func__, 1);
return 1;
}
@@ -266,7 +266,6 @@ int main(int argc, char ** argv) {
llama_sampler_free(smpl);
- llama_batch_free(batch);
llama_free(ctx);
llama_model_free(model);
diff --git a/examples/retrieval/retrieval.cpp b/examples/retrieval/retrieval.cpp
index 7d93ab117..8793e5751 100644
--- a/examples/retrieval/retrieval.cpp
+++ b/examples/retrieval/retrieval.cpp
@@ -75,30 +75,30 @@ static std::vector<chunk> chunk_file(const std::string & filename, int chunk_siz
return chunks;
}
-static void batch_add_seq(llama_batch & batch, const std::vector<int32_t> & tokens, llama_seq_id seq_id) {
+static void batch_add_seq(common_batch & batch, const std::vector<int32_t> & tokens, llama_seq_id seq_id) {
size_t n_tokens = tokens.size();
for (size_t i = 0; i < n_tokens; i++) {
- common_batch_add(batch, tokens[i], i, { seq_id }, true);
+ batch.add(tokens[i], i, seq_id, true);
}
}
-static void batch_process(llama_context * ctx, llama_batch & batch, float * output, int n_seq, int n_embd) {
+static void batch_process(llama_context * ctx, common_batch & batch, float * output, int n_seq, int n_embd) {
// clear previous kv_cache values (irrelevant for embeddings)
llama_memory_clear(llama_get_memory(ctx), false);
// run model
- LOG_INF("%s: n_tokens = %d, n_seq = %d\n", __func__, batch.n_tokens, n_seq);
- if (llama_decode(ctx, batch) < 0) {
+ LOG_INF("%s: n_tokens = %d, n_seq = %d\n", __func__, batch.size(), n_seq);
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) < 0) {
LOG_ERR("%s : failed to process\n", __func__);
}
- for (int i = 0; i < batch.n_tokens; i++) {
- if (!batch.logits[i]) {
+ for (int i = 0; i < batch.size(); i++) {
+ if (!batch.tokens[i].output) {
continue;
}
// try to get sequence embeddings - supported only when pooling_type is not NONE
- const float * embd = llama_get_embeddings_seq(ctx, batch.seq_id[i][0]);
+ const float * embd = llama_get_embeddings_seq(ctx, batch.tokens[i].seq_id);
if (embd == NULL) {
embd = llama_get_embeddings_ith(ctx, i);
if (embd == NULL) {
@@ -107,7 +107,7 @@ static void batch_process(llama_context * ctx, llama_batch & batch, float * outp
}
}
- float * out = output + batch.seq_id[i][0] * n_embd;
+ float * out = output + batch.tokens[i].seq_id * n_embd;
common_embd_normalize(embd, out, n_embd, 2);
}
}
@@ -217,7 +217,7 @@ int main(int argc, char ** argv) {
// initialize batch
const int n_chunks = chunks.size();
- struct llama_batch batch = llama_batch_init(n_batch, 0, 1);
+ common_batch batch(ctx);
// allocate output
const int n_embd_out = llama_model_n_embd_out(model);
@@ -234,10 +234,10 @@ int main(int argc, char ** argv) {
const uint64_t n_toks = inp.size();
// encode if at capacity
- if (batch.n_tokens + n_toks > n_batch || s >= llama_n_seq_max(ctx)) {
+ if (batch.size() + n_toks > n_batch || s >= llama_n_seq_max(ctx)) {
float * out = emb + p * n_embd_out;
batch_process(ctx, batch, out, s, n_embd_out);
- common_batch_clear(batch);
+ batch.clear();
p += s;
s = 0;
}
@@ -258,7 +258,7 @@ int main(int argc, char ** argv) {
chunks[i].tokens.clear();
}
- struct llama_batch query_batch = llama_batch_init(n_batch, 0, 1);
+ common_batch query_batch(ctx);
// start loop, receive query and return top k similar chunks based on cosine similarity
std::string query;
@@ -272,7 +272,7 @@ int main(int argc, char ** argv) {
std::vector<float> query_emb(n_embd_out, 0);
batch_process(ctx, query_batch, query_emb.data(), 1, n_embd_out);
- common_batch_clear(query_batch);
+ query_batch.clear();
// compute cosine similarities
{
@@ -302,6 +302,5 @@ int main(int argc, char ** argv) {
llama_perf_context_print(ctx);
// clean up
- llama_batch_free(query_batch);
llama_backend_free();
}
diff --git a/examples/simple-chat/simple-chat.cpp b/examples/simple-chat/simple-chat.cpp
index 30a0966e0..0cad9652c 100644
--- a/examples/simple-chat/simple-chat.cpp
+++ b/examples/simple-chat/simple-chat.cpp
@@ -6,6 +6,17 @@
#include <string>
#include <vector>
+// fill the batch with tokens at consecutive positions starting from pos_0, output logits only for the last one
+static void batch_set_tokens(llama_batch_ext * batch, const llama_token * tokens, int32_t n_tokens, llama_pos pos_0) {
+ llama_batch_ext_clear(batch);
+ for (int32_t i = 0; i < n_tokens; ++i) {
+ const int32_t idx = llama_batch_ext_add_token(batch, 0, tokens[i]);
+ const llama_pos pos = pos_0 + i;
+ llama_batch_ext_set_pos(batch, idx, &pos);
+ }
+ llama_batch_ext_set_output_logits(batch, n_tokens - 1, true);
+}
+
static void print_usage(int, char ** argv) {
printf("\nexample usage:\n");
printf("\n %s -m model.gguf [-c context_size] [-ngl n_gpu_layers]\n", argv[0]);
@@ -96,6 +107,8 @@ int main(int argc, char ** argv) {
llama_sampler_chain_add(smpl, llama_sampler_init_temp(0.8f));
llama_sampler_chain_add(smpl, llama_sampler_init_dist(LLAMA_DEFAULT_SEED));
+ llama_batch_ext * batch = llama_batch_ext_init(ctx);
+
// helper function to evaluate a prompt and generate a response
auto generate = [&](const std::string & prompt) {
std::string response;
@@ -109,20 +122,25 @@ int main(int argc, char ** argv) {
GGML_ABORT("failed to tokenize the prompt\n");
}
- // prepare a batch for the prompt
- llama_batch batch = llama_batch_get_one(prompt_tokens.data(), prompt_tokens.size());
+ // the tokens to evaluate next: the prompt, then the sampled token
+ const llama_token * tokens = prompt_tokens.data();
+ int n_tokens = prompt_tokens.size();
+
llama_token new_token_id;
while (true) {
// check if we have enough space in the context to evaluate this batch
int n_ctx = llama_n_ctx(ctx);
int n_ctx_used = llama_memory_seq_pos_max(llama_get_memory(ctx), 0) + 1;
- if (n_ctx_used + batch.n_tokens > n_ctx) {
+ if (n_ctx_used + n_tokens > n_ctx) {
printf("\033[0m\n");
fprintf(stderr, "context size exceeded\n");
exit(0);
}
- int ret = llama_decode(ctx, batch);
+ // positions continue from the memory
+ batch_set_tokens(batch, tokens, n_tokens, n_ctx_used);
+
+ int ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch);
if (ret != 0) {
GGML_ABORT("failed to decode, ret = %d\n", ret);
}
@@ -147,7 +165,8 @@ int main(int argc, char ** argv) {
response += piece;
// prepare the next batch with the sampled token
- batch = llama_batch_get_one(&new_token_id, 1);
+ tokens = &new_token_id;
+ n_tokens = 1;
}
return response;
@@ -201,6 +220,7 @@ int main(int argc, char ** argv) {
for (auto & msg : messages) {
free(const_cast<char *>(msg.content));
}
+ llama_batch_ext_free(batch);
llama_sampler_free(smpl);
llama_free(ctx);
llama_model_free(model);
diff --git a/examples/simple/simple.cpp b/examples/simple/simple.cpp
index 982a4d860..ebb1969fb 100644
--- a/examples/simple/simple.cpp
+++ b/examples/simple/simple.cpp
@@ -5,6 +5,17 @@
#include <string>
#include <vector>
+// fill the batch with tokens at consecutive positions starting from pos_0, output logits only for the last one
+static void batch_set_tokens(llama_batch_ext * batch, const llama_token * tokens, int32_t n_tokens, llama_pos pos_0) {
+ llama_batch_ext_clear(batch);
+ for (int32_t i = 0; i < n_tokens; ++i) {
+ const int32_t idx = llama_batch_ext_add_token(batch, 0, tokens[i]);
+ const llama_pos pos = pos_0 + i;
+ llama_batch_ext_set_pos(batch, idx, &pos);
+ }
+ llama_batch_ext_set_output_logits(batch, n_tokens - 1, true);
+}
+
static void print_usage(int, char ** argv) {
printf("\nexample usage:\n");
printf("\n %s -m model.gguf [-n n_predict] [-ngl n_gpu_layers] [prompt]\n", argv[0]);
@@ -144,10 +155,13 @@ int main(int argc, char ** argv) {
// prepare a batch for the prompt
- llama_batch batch = llama_batch_get_one(prompt_tokens.data(), prompt_tokens.size());
+ llama_batch_ext * batch = llama_batch_ext_init(ctx);
+ int n_tokens = n_prompt; // number of tokens in the current batch
+
+ batch_set_tokens(batch, prompt_tokens.data(), n_prompt, 0);
if (llama_model_has_encoder(model)) {
- if (llama_encode(ctx, batch)) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_ENCODE, batch)) {
fprintf(stderr, "%s : failed to eval\n", __func__);
return 1;
}
@@ -157,7 +171,8 @@ int main(int argc, char ** argv) {
decoder_start_token_id = llama_vocab_bos(vocab);
}
- batch = llama_batch_get_one(&decoder_start_token_id, 1);
+ batch_set_tokens(batch, &decoder_start_token_id, 1, 0);
+ n_tokens = 1;
}
// main loop
@@ -166,14 +181,14 @@ int main(int argc, char ** argv) {
int n_decode = 0;
llama_token new_token_id;
- for (int n_pos = 0; n_pos + batch.n_tokens < n_prompt + n_predict; ) {
+ for (int n_pos = 0; n_pos + n_tokens < n_prompt + n_predict; ) {
// evaluate the current batch with the transformer model
- if (llama_decode(ctx, batch)) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch)) {
fprintf(stderr, "%s : failed to eval, return code %d\n", __func__, 1);
return 1;
}
- n_pos += batch.n_tokens;
+ n_pos += n_tokens;
// sample the next token
{
@@ -195,7 +210,8 @@ int main(int argc, char ** argv) {
fflush(stdout);
// prepare the next batch with the sampled token
- batch = llama_batch_get_one(&new_token_id, 1);
+ batch_set_tokens(batch, &new_token_id, 1, n_pos);
+ n_tokens = 1;
n_decode += 1;
}
@@ -213,6 +229,7 @@ int main(int argc, char ** argv) {
llama_perf_context_print(ctx);
fprintf(stderr, "\n");
+ llama_batch_ext_free(batch);
llama_sampler_free(smpl);
llama_free(ctx);
llama_model_free(model);
diff --git a/examples/speculative-simple/speculative-simple.cpp b/examples/speculative-simple/speculative-simple.cpp
index 81aa106f1..08a1f2a88 100644
--- a/examples/speculative-simple/speculative-simple.cpp
+++ b/examples/speculative-simple/speculative-simple.cpp
@@ -125,12 +125,12 @@ int main(int argc, char ** argv) {
// eval the prompt on the target and feed it to the speculative implementation(s)
{
- llama_batch batch_prompt = llama_batch_init(inp.size(), 0, 1);
+ common_batch batch_prompt(ctx_tgt);
for (size_t i = 0; i < inp.size() - 1; ++i) {
- common_batch_add(batch_prompt, inp[i], i, { seq_id }, false);
+ batch_prompt.add(inp[i], i, seq_id, false);
}
- llama_decode(ctx_tgt, batch_prompt);
+ llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, batch_prompt.get());
if (!common_speculative_process(spec, batch_prompt)) {
LOG_ERR("%s", "failed to process speculative prompt\n");
@@ -149,7 +149,7 @@ int main(int argc, char ** argv) {
common_speculative_begin(spec, seq_id, prompt_tgt);
- llama_batch batch_tgt = llama_batch_init(llama_n_batch(ctx_tgt), 0, 1);
+ common_batch batch_tgt(ctx_tgt);
llama_tokens draft;
@@ -219,17 +219,17 @@ int main(int argc, char ** argv) {
}
// always have a token to evaluate from before - id_last
- common_batch_clear(batch_tgt);
- common_batch_add (batch_tgt, id_last, n_past++, { seq_id }, true);
+ batch_tgt.clear();
+ batch_tgt.add(id_last, n_past++, seq_id, true);
// evaluate the target model on [id_last, draft0, draft1, ..., draftN-1]
{
for (size_t i = 0; i < draft.size(); ++i) {
- common_batch_add(batch_tgt, draft[i], n_past + i, { seq_id }, true);
+ batch_tgt.add(draft[i], n_past + i, seq_id, true);
}
- llama_decode(ctx_tgt, batch_tgt);
+ llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, batch_tgt.get());
}
// feed the batch to the speculative implementation(s) - this drives the draft model, MTP, Eagle3, etc.
@@ -364,7 +364,6 @@ int main(int argc, char ** argv) {
LOG_INF("target:\n\n");
common_perf_print(ctx_tgt, smpl.get());
- llama_batch_free(batch_tgt);
common_speculative_free(spec);
diff --git a/examples/speculative/speculative.cpp b/examples/speculative/speculative.cpp
index 17071aa05..1bb47594b 100644
--- a/examples/speculative/speculative.cpp
+++ b/examples/speculative/speculative.cpp
@@ -190,9 +190,16 @@ int main(int argc, char ** argv) {
const auto t_enc_start = ggml_time_us();
// eval the prompt with both models
- llama_decode(ctx_tgt, llama_batch_get_one( inp.data(), n_input - 1));
- llama_decode(ctx_tgt, llama_batch_get_one(&inp.back(), 1));
- llama_decode(ctx_dft, llama_batch_get_one( inp.data(), n_input));
+ {
+ common_batch batch = common_batch_get_one(ctx_tgt, inp.data(), n_input - 1);
+ llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+
+ batch = common_batch_get_one(ctx_tgt, &inp.back(), 1);
+ llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+
+ batch = common_batch_get_one(ctx_dft, inp.data(), n_input);
+ llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+ }
const auto t_enc_end = ggml_time_us();
@@ -223,8 +230,8 @@ int main(int argc, char ** argv) {
drafts[s].smpl = common_sampler_init(model_dft, params.sampling);
}
- llama_batch batch_dft = llama_batch_init(llama_n_batch(ctx_dft), 0, 1);
- llama_batch batch_tgt = llama_batch_init(llama_n_batch(ctx_tgt), 0, n_seq_dft);
+ common_batch batch_dft(ctx_dft);
+ common_batch batch_tgt(ctx_tgt);
const auto t_dec_start = ggml_time_us();
@@ -465,12 +472,12 @@ int main(int argc, char ** argv) {
drafts[0].dists.push_back(std::vector<llama_token_data>());
drafts[0].i_batch_tgt.push_back(0);
- common_batch_clear(batch_dft);
- common_batch_add (batch_dft, token_id, n_past_dft, { 0 }, true);
+ batch_dft.clear();
+ batch_dft.add(token_id, n_past_dft, 0, true);
llama_memory_seq_rm(mem_dft, 0, n_past_dft, -1);
// LOG_DBG("dft batch: %s\n", LOG_BATCH_TOSTR_PRETTY(ctx_dft, batch_dft).c_str());
- llama_decode(ctx_dft, batch_dft);
+ llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch_dft.get());
++n_past_dft;
}
@@ -495,12 +502,12 @@ int main(int argc, char ** argv) {
drafts[0].drafting = true;
drafts[0].i_batch_dft = 0;
- common_batch_clear(batch_tgt);
- common_batch_add (batch_tgt, drafts[0].tokens[0], n_past_tgt, { 0 }, true);
+ batch_tgt.clear();
+ batch_tgt.add(drafts[0].tokens[0], n_past_tgt, 0, true);
// sample n_draft tokens from the draft model using tree-based sampling
for (int i = 0; i < n_draft; ++i) {
- batch_dft.n_tokens = 0;
+ batch_dft.clear();
for (int s = 0; s < n_seq_dft; ++s) {
drafts[s].skip = false;
@@ -531,14 +538,8 @@ int main(int argc, char ** argv) {
llama_memory_seq_cp(mem_dft, s, n_seq_cur, -1, -1);
// all previous tokens from this branch are now also part of the new branch
- for (int t = 0; t < batch_tgt.n_tokens; ++t) {
- for (int p = 0; p < batch_tgt.n_seq_id[t]; ++p) {
- if (batch_tgt.seq_id[t][p] == s) {
- batch_tgt.seq_id[t][batch_tgt.n_seq_id[t]] = n_seq_cur;
- batch_tgt.n_seq_id[t]++;
- break;
- }
- }
+ for (int t : drafts[s].i_batch_tgt) {
+ batch_tgt.add_seq(t, n_seq_cur);
}
// copy the draft state
@@ -577,32 +578,32 @@ int main(int argc, char ** argv) {
drafts[s].dists.push_back({cur_p->data, cur_p->data + cur_p->size});
// add unique drafted tokens to the target batch
- drafts[s].i_batch_tgt.push_back(batch_tgt.n_tokens);
+ drafts[s].i_batch_tgt.push_back(batch_tgt.size());
- common_batch_add(batch_tgt, id, n_past_tgt + i + 1, { s }, true);
+ batch_tgt.add(id, n_past_tgt + i + 1, s, true);
// add the token to the batch for batched decoding with the draft model
- drafts[s].i_batch_dft = batch_dft.n_tokens;
+ drafts[s].i_batch_dft = batch_dft.size();
- common_batch_add(batch_dft, id, n_past_cur, { s }, true);
+ batch_dft.add(id, n_past_cur, s, true);
- if (batch_tgt.n_tokens > n_draft) {
+ if (batch_tgt.size() > n_draft) {
drafts[s].drafting = false;
}
}
}
// no sequence is drafting anymore
- if (batch_dft.n_tokens == 0) {
+ if (batch_dft.size() == 0) {
break;
}
// evaluate the drafted tokens on the draft model
- llama_decode(ctx_dft, batch_dft);
+ llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch_dft.get());
++n_past_cur;
++n_drafted;
- if (batch_tgt.n_tokens > n_draft) {
+ if (batch_tgt.size() > n_draft) {
break;
}
}
@@ -615,7 +616,7 @@ int main(int argc, char ** argv) {
}
// LOG_DBG("target batch: %s\n", LOG_BATCH_TOSTR_PRETTY(ctx_tgt, batch_tgt).c_str());
- llama_decode(ctx_tgt, batch_tgt);
+ llama_process(ctx_tgt, LLAMA_PROCESS_TYPE_DECODE, batch_tgt.get());
++n_past_tgt;
}
@@ -658,7 +659,6 @@ int main(int argc, char ** argv) {
common_sampler_free(drafts[s].smpl);
}
- llama_batch_free(batch_dft);
llama_backend_free();
diff --git a/tests/test-backend-sampler.cpp b/tests/test-backend-sampler.cpp
index c23e7248d..56736ac46 100644
--- a/tests/test-backend-sampler.cpp
+++ b/tests/test-backend-sampler.cpp
@@ -129,7 +129,7 @@ struct test_context {
GGML_ASSERT(ctx);
last_batch_info.clear();
- llama_batch batch = llama_batch_init(512, 0, prompts.size());
+ common_batch batch(ctx.get());
for (const auto & [seq_id, prompt] : prompts) {
std::vector<llama_token> tokens;
@@ -141,7 +141,6 @@ struct test_context {
false, false);
if (n_tokens < 0) {
fprintf(stderr, "Warning: tokenization failed for seq_id %d\n", seq_id);
- llama_batch_free(batch);
return false;
}
@@ -155,7 +154,7 @@ struct test_context {
int32_t start_pos = seq_positions[seq_id];
for (size_t i = 0; i < tokens.size(); i++) {
- common_batch_add(batch, tokens[i], start_pos + i, { seq_id }, i == tokens.size() - 1);
+ batch.add(tokens[i], start_pos + i, seq_id, i == tokens.size() - 1);
}
seq_positions[seq_id] = start_pos + tokens.size();
@@ -163,31 +162,18 @@ struct test_context {
printf("Batch contents:\n");
- printf("n_tokens: %d\n", batch.n_tokens);
- for (int i = 0; i < batch.n_tokens; i++) {
- printf("token[%d]: tok=%-5d, pos=%d, n_seq_id=%d, seq_ids=[", i, batch.token[i], batch.pos[i], batch.n_seq_id[i]);
-
- for (int j = 0; j < batch.n_seq_id[i]; j++) {
- printf("%d%s", batch.seq_id[i][j], j < batch.n_seq_id[i]-1 ? ", " : "");
- }
- printf("], logits=%d\n", batch.logits[i]);
+ printf("n_tokens: %d\n", batch.size());
+ for (int i = 0; i < batch.size(); i++) {
+ const auto & t = batch.tokens[i];
+ printf("token[%d]: tok=%-5d, pos=%d, seq_id=%d, logits=%d\n", i, t.id, t.pos[0], t.seq_id, t.output);
}
- if (llama_decode(ctx.get(), batch) != 0) {
+ if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
fprintf(stderr, "Warning: llama_decode failed\n");
- llama_batch_free(batch);
return false;
}
- // Build mapping from seq id to batch token idx
- for (int i = 0; i < batch.n_tokens; i++) {
- if (batch.logits[i]) {
- llama_seq_id seq_id = batch.seq_id[i][0];
- last_batch_info[seq_id] = i;
- }
- }
-
- llama_batch_free(batch);
+ update_batch_info(batch);
return true;
}
@@ -200,11 +186,12 @@ struct test_context {
return it->second;
}
- void update_batch_info(const llama_batch & batch) {
+ // build mapping from seq id to batch token idx
+ void update_batch_info(const common_batch & batch) {
last_batch_info.clear();
- for (int i = 0; i < batch.n_tokens; i++) {
- if (batch.logits[i]) {
- llama_seq_id cur_seq = batch.seq_id[i][0];
+ for (int i = 0; i < batch.size(); i++) {
+ if (batch.tokens[i].output) {
+ llama_seq_id cur_seq = batch.tokens[i].seq_id;
last_batch_info[cur_seq] = i;
}
}
@@ -213,20 +200,18 @@ struct test_context {
bool decode_token(llama_token token, llama_seq_id seq_id = 0) {
GGML_ASSERT(ctx);
- llama_batch batch = llama_batch_init(1, 0, 1);
+ common_batch batch(ctx.get());
int32_t pos = seq_positions[seq_id];
- common_batch_add(batch, token, pos, { seq_id }, true);
+ batch.add(token, pos, seq_id, true);
- if (llama_decode(ctx.get(), batch) != 0) {
+ if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
fprintf(stderr, "Warning: llama_decode failed for token %d in seq %d\n", token, seq_id);
- llama_batch_free(batch);
return false;
}
update_batch_info(batch);
seq_positions[seq_id]++;
- llama_batch_free(batch);
return true;
}
@@ -234,16 +219,15 @@ struct test_context {
bool decode_tokens(const std::map<llama_seq_id, llama_token> & seq_tokens) {
GGML_ASSERT(ctx);
- llama_batch batch = llama_batch_init(seq_tokens.size(), 0, seq_tokens.size());
+ common_batch batch(ctx.get());
for (const auto & [seq_id, token] : seq_tokens) {
int32_t pos = seq_positions[seq_id];
- common_batch_add(batch, token, pos, { seq_id }, true);
+ batch.add(token, pos, seq_id, true);
}
- if (llama_decode(ctx.get(), batch) != 0) {
+ if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
fprintf(stderr, "Warning: llama_decode failed for batch tokens\n");
- llama_batch_free(batch);
return false;
}
@@ -253,8 +237,6 @@ struct test_context {
update_batch_info(batch);
- llama_batch_free(batch);
-
return true;
}
@@ -1607,18 +1589,16 @@ static void test_backend_multi_output_limit(const test_params & params) {
std::vector<llama_sampler_seq_config> configs = {{ seq_id, chain.get() }};
test_context test_ctx(params, configs, 1, 3, 0, 2);
- llama_batch batch = llama_batch_init(3, 0, 1);
+ common_batch batch(test_ctx.ctx.get());
for (int i = 0; i < 3; ++i) {
- common_batch_add(batch, llama_vocab_bos(test_ctx.vocab), i, { seq_id }, true);
+ batch.add(llama_vocab_bos(test_ctx.vocab), i, seq_id, true);
}
printf(">>> test_backend_multi_output_limit expected error start:\n");
- const int ret = llama_decode(test_ctx.ctx.get(), batch);
+ const int ret = llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get());
GGML_ASSERT(ret != 0 && "llama_decode should reject outputs above the per-sequence limit");
printf("<<< test_backend_multi_output_limit expected error end.\n");
- llama_batch_free(batch);
-
printf("backend multi-output limit test PASSED\n");
}
@@ -1649,14 +1629,22 @@ static void test_backend_multi_sequence_multi_output_dist(const test_params & pa
{ llama_vocab_eos(vocab), llama_vocab_bos(vocab) },
};
- llama_batch batch = llama_batch_init(4, 0, 1);
- for (int pos = 0; pos < 2; ++pos) {
- common_batch_add(batch, seq_tokens[0][pos], pos, { 0 }, true);
- common_batch_add(batch, seq_tokens[1][pos], pos, { 1 }, true);
- }
+ // a batch belongs to one context, so it is built per context
+ auto make_batch = [&](llama_context * ctx) {
+ common_batch batch(ctx);
+ for (int pos = 0; pos < 2; ++pos) {
+ batch.add(seq_tokens[0][pos], pos, 0, true);
+ batch.add(seq_tokens[1][pos], pos, 1, true);
+ }
+ return batch;
+ };
- GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
- GGML_ASSERT(llama_decode(reference_ctx.ctx.get(), batch) == 0);
+ common_batch batch = make_batch(test_ctx.ctx.get());
+ GGML_ASSERT(llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0);
+ {
+ common_batch batch_ref = make_batch(reference_ctx.ctx.get());
+ GGML_ASSERT(llama_process(reference_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch_ref.get()) == 0);
+ }
std::mt19937 reference_rngs[] = {
std::mt19937(seeds[0]),
@@ -1664,8 +1652,8 @@ static void test_backend_multi_sequence_multi_output_dist(const test_params & pa
};
std::uniform_real_distribution<double> reference_dist(0.0, 1.0);
- for (int i = 0; i < batch.n_tokens; ++i) {
- const llama_seq_id seq_id = batch.seq_id[i][0];
+ for (int i = 0; i < batch.size(); ++i) {
+ const llama_seq_id seq_id = batch.tokens[i].seq_id;
GGML_ASSERT(seq_id == 0 || seq_id == 1);
llama_sampler * chain = seq_id == 0 ? chain_0.get() : chain_1.get();
@@ -1706,8 +1694,6 @@ static void test_backend_multi_sequence_multi_output_dist(const test_params & pa
GGML_ASSERT(rnd <= cumsum_sampled + 1e-4f);
}
- llama_batch_free(batch);
-
printf("backend multi-sequence multi-output dist test PASSED\n");
}
@@ -1750,33 +1736,28 @@ static void test_backend_multi_output_dist_transaction(const test_params & param
int32_t pos = 0;
auto decode = [&]() {
- llama_batch batch = llama_batch_init(3, 0, 1);
+ common_batch batch(test_ctx.ctx.get());
for (int32_t i = 0; i < 3; ++i) {
- common_batch_add(batch, llama_vocab_bos(vocab), pos++, { seq_id }, true);
+ batch.add(llama_vocab_bos(vocab), pos++, seq_id, true);
}
- GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
- return batch;
+ GGML_ASSERT(llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0);
};
- llama_batch batch = decode();
+ decode();
verify_random(0, randoms[0], false);
- llama_batch_free(batch);
- batch = decode();
+ decode();
verify_random(0, randoms[0]);
verify_random(1, randoms[1]);
- llama_batch_free(batch);
- batch = decode();
+ decode();
llama_sampler_ptr saved(llama_sampler_clone(chain.get()));
verify_random(0, randoms[2]);
- llama_batch_free(batch);
llama_sampler_copy(saved.get(), chain.get());
- batch = decode();
+ decode();
verify_random(0, randoms[2]);
- llama_batch_free(batch);
printf("backend multi-output dist transaction test PASSED\n");
}
@@ -1817,19 +1798,23 @@ static void test_backend_multi_output_sampling_chain(const test_params & params)
llama_sampler_ptr reference_temp(llama_sampler_init_temp(temp));
std::vector<llama_token_data> reference_data(n_vocab);
- auto make_batch = [&](int32_t pos) {
- llama_batch batch = llama_batch_init(2, 0, 1);
+ // a batch belongs to one context, so it is built per context
+ auto make_batch = [&](llama_context * ctx, int32_t pos) {
+ common_batch batch(ctx);
for (int i = 0; i < 2; ++i) {
- common_batch_add(batch, llama_vocab_bos(vocab), pos + i, { seq_id }, true);
+ batch.add(llama_vocab_bos(vocab), pos + i, seq_id, true);
}
return batch;
};
- llama_batch batch = make_batch(0);
- GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
- GGML_ASSERT(llama_decode(reference_ctx.ctx.get(), batch) == 0);
+ common_batch batch = make_batch(test_ctx.ctx.get(), 0);
+ GGML_ASSERT(llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0);
+ {
+ common_batch batch_ref = make_batch(reference_ctx.ctx.get(), 0);
+ GGML_ASSERT(llama_process(reference_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch_ref.get()) == 0);
+ }
- for (int i = 0; i < batch.n_tokens; ++i) {
+ for (int i = 0; i < batch.size(); ++i) {
const llama_token backend_token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), i);
const float * sampled_logits = llama_get_sampled_logits_ith(test_ctx.ctx.get(), i);
const float * sampled_probs = llama_get_sampled_probs_ith(test_ctx.ctx.get(), i);
@@ -1922,11 +1907,8 @@ static void test_backend_multi_output_sampling_chain(const test_params & params)
GGML_ASSERT(std::fabs(prob_sum - 1.0f) <= 1e-3f);
}
- llama_batch_free(batch);
-
- batch = make_batch(2);
- GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
- llama_batch_free(batch);
+ batch = make_batch(test_ctx.ctx.get(), 2);
+ GGML_ASSERT(llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0);
printf("backend multi-output sampling chain test PASSED\n");
}
@@ -1950,17 +1932,15 @@ static void test_backend_multi_output_cpu_suffix(const test_params & params) {
std::vector<llama_sampler_seq_config> configs = {{ seq_id, chain.get() }};
test_context test_ctx(params, configs, 1, 1, 0, 4);
- llama_batch batch = llama_batch_init(1, 0, 1);
- common_batch_add(batch, llama_vocab_bos(vocab), 0, { seq_id }, true);
- GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
+ common_batch batch(test_ctx.ctx.get());
+ batch.add(llama_vocab_bos(vocab), 0, seq_id, true);
+ GGML_ASSERT(llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0);
GGML_ASSERT(sampler_ctx->backend_initialized);
GGML_ASSERT(sampler_ctx->backend_outputs_max_per_seq == 1);
GGML_ASSERT(sampler_ctx->backend_apply_count > 0);
GGML_ASSERT(sampler_ctx->apply_count == 0);
GGML_ASSERT(llama_get_sampled_token_ith(test_ctx.ctx.get(), 0) != LLAMA_TOKEN_NULL);
-
- llama_batch_free(batch);
}
{
@@ -1969,25 +1949,23 @@ static void test_backend_multi_output_cpu_suffix(const test_params & params) {
std::vector<llama_sampler_seq_config> configs = {{ seq_id, chain.get() }};
test_context test_ctx(params, configs, 1, 2, 0, 0);
- llama_batch batch = llama_batch_init(2, 0, 1);
+ common_batch batch(test_ctx.ctx.get());
for (int i = 0; i < 2; ++i) {
- common_batch_add(batch, llama_vocab_bos(vocab), i, { seq_id }, true);
+ batch.add(llama_vocab_bos(vocab), i, seq_id, true);
}
- GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
+ GGML_ASSERT(llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0);
GGML_ASSERT(!sampler_ctx->backend_initialized);
GGML_ASSERT(sampler_ctx->backend_outputs_max_per_seq == 2);
GGML_ASSERT(sampler_ctx->backend_apply_count == 0);
- for (int i = 0; i < batch.n_tokens; ++i) {
+ for (int i = 0; i < batch.size(); ++i) {
GGML_ASSERT(llama_get_sampled_token_ith(test_ctx.ctx.get(), i) == LLAMA_TOKEN_NULL);
GGML_ASSERT(llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), i) == (uint32_t) k);
GGML_ASSERT(llama_get_sampled_candidates_count_ith(test_ctx.ctx.get(), i) == (uint32_t) k);
const llama_token token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), i);
GGML_ASSERT(token >= 0 && token < llama_vocab_n_tokens(vocab));
}
- GGML_ASSERT(sampler_ctx->apply_count == batch.n_tokens);
-
- llama_batch_free(batch);
+ GGML_ASSERT(sampler_ctx->apply_count == batch.size());
}
printf("backend multi-output CPU suffix test PASSED\n");
diff --git a/tests/test-fusion.cpp b/tests/test-fusion.cpp
index 65d444ecf..55e4f9df4 100644
--- a/tests/test-fusion.cpp
+++ b/tests/test-fusion.cpp
@@ -145,13 +145,11 @@ static llama_context_ptr create_ctx(llama_model * model, int n_ubatch) {
// decode all tokens in one batch; returns the logits of every token
static std::vector<float> decode_prefill(llama_model * model, llama_context * lctx, const std::vector<llama_token> & tokens) {
const uint32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model));
- llama_batch batch = llama_batch_init(tokens.size(), 0, 1);
+ common_batch batch(lctx);
for (size_t i = 0; i < tokens.size(); i++) {
- common_batch_add(batch, tokens[i], i, { 0 }, true);
+ batch.add(tokens[i], i, 0, true);
}
- batch.n_tokens = tokens.size();
- if (llama_decode(lctx, batch)) {
- llama_batch_free(batch);
+ if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
throw std::runtime_error("prefill decode failed");
}
@@ -163,20 +161,18 @@ static std::vector<float> decode_prefill(llama_model * model, llama_context * lc
ret.push_back(logits_ith[j]);
}
}
- llama_batch_free(batch);
return ret;
}
// decode one token at a time; returns the logits of the last token of each step
static std::vector<float> decode_gen(llama_model * model, llama_context * lctx, const std::vector<llama_token> & tokens) {
const uint32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model));
- llama_batch batch = llama_batch_init(1, 0, 1);
+ common_batch batch(lctx);
std::vector<float> ret;
for (size_t i = 0; i < tokens.size(); i++) {
- common_batch_clear(batch);
- common_batch_add(batch, tokens[i], i, { 0 }, true);
- if (llama_decode(lctx, batch)) {
- llama_batch_free(batch);
+ batch.clear();
+ batch.add(tokens[i], i, 0, true);
+ if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
throw std::runtime_error("decode failed");
}
const float * logits = llama_get_logits_ith(lctx, 0);
@@ -184,7 +180,6 @@ static std::vector<float> decode_gen(llama_model * model, llama_context * lctx,
ret.push_back(logits[j]);
}
}
- llama_batch_free(batch);
return ret;
}
diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp
index 298073d06..015e3414b 100644
--- a/tests/test-llama-archs.cpp
+++ b/tests/test-llama-archs.cpp
@@ -510,20 +510,17 @@ static std::vector<float> get_logits(
const uint32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model));
const uint32_t n_ctx = llama_n_ctx(lctx);
const uint32_t n_tokens = tokens.size();
- llama_batch batch = llama_batch_init(n_ctx, 0, 1);
+ common_batch batch(lctx);
GGML_ASSERT(n_tokens <= n_ctx);
for (uint32_t pos = 0; pos < n_tokens; pos++) {
- common_batch_add(batch, tokens[pos], pos, {0}, true);
+ batch.add(tokens[pos], pos, 0, true);
}
- batch.n_tokens = n_tokens;
if (encode) {
- if (llama_encode(lctx, batch)) {
- llama_batch_free(batch);
+ if (llama_process(lctx, LLAMA_PROCESS_TYPE_ENCODE, batch.get())) {
throw std::runtime_error("failed to encode batch");
}
}
- if (llama_decode(lctx, batch)) {
- llama_batch_free(batch);
+ if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
throw std::runtime_error("failed to decode batch");
}
@@ -535,7 +532,6 @@ static std::vector<float> get_logits(
ret.push_back(logits_ith[j]);
}
}
- llama_batch_free(batch);
return ret;
}
diff --git a/tests/test-recurrent-state-rollback.cpp b/tests/test-recurrent-state-rollback.cpp
index 4ad0d6f9e..f8eda55c8 100644
--- a/tests/test-recurrent-state-rollback.cpp
+++ b/tests/test-recurrent-state-rollback.cpp
@@ -35,21 +35,17 @@ static const char * test_status_str(test_status status) {
}
static bool decode_tokens(llama_context * ctx, const std::vector<llama_token> & tokens, uint32_t count) {
- llama_batch batch = llama_batch_init(count, 0, 1);
+ common_batch batch(ctx);
for (uint32_t pos = 0; pos < count; ++pos) {
- common_batch_add(batch, tokens[pos], pos, { 0 }, pos + 1 == count);
+ batch.add(tokens[pos], pos, 0, pos + 1 == count);
}
- const bool ok = llama_decode(ctx, batch) == 0;
- llama_batch_free(batch);
- return ok;
+ return llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0;
}
-static bool decode_one(llama_context * ctx, llama_token tok, llama_pos pos) {
- llama_batch batch = llama_batch_init(1, 0, 1);
- common_batch_add(batch, tok, pos, { 0 }, true);
- const bool ok = llama_decode(ctx, batch) == 0;
- llama_batch_free(batch);
- return ok;
+static bool decode_one(llama_context * ctx, llama_token tok, llama_pos pos, llama_seq_id seq = 0) {
+ common_batch batch(ctx);
+ batch.add(tok, pos, seq, true);
+ return llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0;
}
struct cache_buffer_collector : llama_io_write_i {
@@ -166,22 +162,22 @@ static test_status test_multi_seq_split_replay(const common_params & params, lla
bool ok = true;
+ // decode tokens [p_begin, p_end) of seq s, a batch belongs to one context so it is built per call
+ const auto decode_range = [&](llama_context * ctx, uint32_t s, llama_pos p_begin, llama_pos p_end) {
+ common_batch batch(ctx);
+ for (llama_pos pos = p_begin; pos < p_end; ++pos) {
+ batch.add(tok(s, pos), pos, (llama_seq_id) s, false);
+ }
+ return llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0;
+ };
+
// both contexts decode the identical [0, p0) prefill; only ctx_roll decodes
// the tail, which is then rolled back so its restore is pending at replay
for (uint32_t s = 0; s < n_seqs && ok; ++s) {
- llama_batch batch = llama_batch_init(n_prompt, 0, 1);
- for (llama_pos pos = 0; pos < (llama_pos) p0; ++pos) {
- common_batch_add(batch, tok(s, pos), pos, { (llama_seq_id) s }, false);
- }
- ok = ok && llama_decode(ctx_roll.get(), batch) == 0;
- ok = ok && llama_decode(ctx_ref.get(), batch) == 0;
+ ok = ok && decode_range(ctx_roll.get(), s, 0, (llama_pos) p0);
+ ok = ok && decode_range(ctx_ref.get(), s, 0, (llama_pos) p0);
- common_batch_clear(batch);
- for (llama_pos pos = p0; pos < (llama_pos) n_prompt; ++pos) {
- common_batch_add(batch, tok(s, pos), pos, { (llama_seq_id) s }, false);
- }
- ok = ok && llama_decode(ctx_roll.get(), batch) == 0;
- llama_batch_free(batch);
+ ok = ok && decode_range(ctx_roll.get(), s, (llama_pos) p0, (llama_pos) n_prompt);
ok = ok && llama_memory_seq_rm(llama_get_memory(ctx_roll.get()), (llama_seq_id) s, p0, -1);
@@ -193,16 +189,19 @@ static test_status test_multi_seq_split_replay(const common_params & params, lla
return test_status::FAIL;
}
- llama_batch batch = llama_batch_init(n_seqs*n_replay, 0, 1);
- for (uint32_t s = 0; s < n_seqs; ++s) {
- for (uint32_t i = 0; i < n_replay; ++i) {
- const llama_pos pos = p0 + (llama_pos) i;
- common_batch_add(batch, tok(s, pos), pos, { (llama_seq_id) s }, true);
+ // all seqs replay in a single batch
+ const auto decode_replay = [&](llama_context * ctx) {
+ common_batch batch(ctx);
+ for (uint32_t s = 0; s < n_seqs; ++s) {
+ for (uint32_t i = 0; i < n_replay; ++i) {
+ const llama_pos pos = p0 + (llama_pos) i;
+ batch.add(tok(s, pos), pos, (llama_seq_id) s, true);
+ }
}
- }
- ok = llama_decode(ctx_roll.get(), batch) == 0;
- ok = ok && llama_decode(ctx_ref.get(), batch) == 0;
- llama_batch_free(batch);
+ return llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0;
+ };
+ ok = decode_replay(ctx_roll.get());
+ ok = ok && decode_replay(ctx_ref.get());
if (!ok) {
LOG_ERR("%s: multi-seq replay decode failed\n", __func__);
return test_status::FAIL;
@@ -258,13 +257,12 @@ static test_status test_multi_seq_split_replay(const common_params & params, lla
constexpr uint32_t n_tail = 4;
{
- llama_batch batch_tail = llama_batch_init(n_tail, 0, 1);
+ common_batch batch_tail(ctx_ref.get());
for (uint32_t i = 0; i < n_tail; ++i) {
const llama_pos pos = p0 + (llama_pos) (n_replay + i);
- common_batch_add(batch_tail, tok(0, pos + 7), pos, { 0 }, false);
+ batch_tail.add(tok(0, pos + 7), pos, 0, false);
}
- ok = llama_decode(ctx_ref.get(), batch_tail) == 0;
- llama_batch_free(batch_tail);
+ ok = llama_process(ctx_ref.get(), LLAMA_PROCESS_TYPE_DECODE, batch_tail.get()) == 0;
}
float diff_tail = 0.0f;
@@ -272,11 +270,8 @@ static test_status test_multi_seq_split_replay(const common_params & params, lla
double nmse_tail_a0 = 0.0;
for (uint32_t i = 0; i < n_tail && ok; ++i) {
const llama_pos pos = p0 + (llama_pos) (n_replay + i);
- llama_batch batch_one = llama_batch_init(1, 0, 1);
- common_batch_add(batch_one, tok(1, pos), pos, { 1 }, true);
- ok = llama_decode(ctx_roll.get(), batch_one) == 0;
- ok = ok && llama_decode(ctx_ref.get(), batch_one) == 0;
- llama_batch_free(batch_one);
+ ok = decode_one(ctx_roll.get(), tok(1, pos), pos, 1);
+ ok = ok && decode_one(ctx_ref.get(), tok(1, pos), pos, 1);
if (!ok) {
break;
}
diff --git a/tests/test-save-load-state.cpp b/tests/test-save-load-state.cpp
index dee5e17be..02984d185 100644
--- a/tests/test-save-load-state.cpp
+++ b/tests/test-save-load-state.cpp
@@ -69,26 +69,9 @@ static bool get_current_logits(llama_context * ctx, std::vector<float> & out) {
return true;
}
-struct llama_batch_ptr {
- llama_batch batch;
-
- llama_batch_ptr(int32_t n_tokens, int32_t embd, int32_t n_seq_max)
- : batch{llama_batch_init(n_tokens, embd, n_seq_max)} {}
-
- ~llama_batch_ptr() { llama_batch_free(batch); }
-
- llama_batch_ptr(const llama_batch_ptr &) = delete;
- llama_batch_ptr & operator=(const llama_batch_ptr &) = delete;
- llama_batch_ptr(llama_batch_ptr &&) = default;
- llama_batch_ptr & operator=(llama_batch_ptr &&) = default;
-
- llama_batch & get() { return batch; }
- const llama_batch & get() const { return batch; }
-};
-
static generation_result generate_tokens(llama_context * ctx, llama_sampler * smpl, int & n_past, int32_t n_predict, llama_seq_id seq_id) {
generation_result result;
- llama_batch_ptr batch(1, 0, 1);
+ common_batch batch(ctx);
for (int i = 0; i < n_predict; i++) {
std::vector<float> logits;
@@ -104,10 +87,10 @@ static generation_result generate_tokens(llama_context * ctx, llama_sampler * sm
result.tokens.push_back(next_token);
result.logits.push_back(std::move(logits));
- common_batch_clear(batch.get());
- common_batch_add(batch.get(), next_token, n_past, {seq_id}, true);
+ batch.clear();
+ batch.add(next_token, n_past, seq_id, true);
- if (llama_decode(ctx, batch.get())) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("\n%s: failed to evaluate\n", __func__);
return {};
}
@@ -125,7 +108,7 @@ static bool generate_tokens_compare(
return false;
}
- llama_batch_ptr batch(1, 0, 1);
+ common_batch batch(ctx);
for (int i = 0; i < n_predict; i++) {
std::vector<float> logits;
@@ -153,10 +136,10 @@ static bool generate_tokens_compare(
LOG_TRC("%s: sampled token %d differs from expected %d, using expected token\n", __func__, next_token, expected_token);
}
- common_batch_clear(batch.get());
- common_batch_add(batch.get(), expected_token, n_past, {seq_id}, true);
+ batch.clear();
+ batch.add(expected_token, n_past, seq_id, true);
- if (llama_decode(ctx, batch.get())) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("\n%s: failed to evaluate\n", __func__);
return false;
}
@@ -222,12 +205,12 @@ static bool test_seq_rm_isolated(
const size_t n_tokens = tokens.size() < 128 ? tokens.size() : 128;
for (llama_seq_id seq_id = 0; seq_id < 2; ++seq_id) {
- llama_batch_ptr batch(n_tokens, 0, 1);
+ common_batch batch(ctx.get());
for (size_t i = 0; i < n_tokens; ++i) {
- common_batch_add(batch.get(), tokens[i], i, { seq_id }, i == n_tokens - 1);
+ batch.add(tokens[i], i, seq_id, i == n_tokens - 1);
}
- if (llama_decode(ctx.get(), batch.get())) {
+ if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("%s: failed to decode prompt for sequence %d\n", __func__, seq_id);
return false;
}
@@ -469,9 +452,9 @@ static bool test_seq_cp_scatter(struct llama_model * model, const struct common_
const uint32_t flags = on_device ? LLAMA_STATE_SEQ_FLAGS_ON_DEVICE : LLAMA_STATE_SEQ_FLAGS_NONE;
auto decode_one = [&](llama_token tok, int pos, llama_seq_id seq) {
- llama_batch_ptr batch(1, 0, 1);
- common_batch_add(batch.get(), tok, pos, { seq }, true);
- return llama_decode(ctx.get(), batch.get()) == 0;
+ common_batch batch(ctx.get());
+ batch.add(tok, pos, seq, true);
+ return llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0;
};
// seq 0 cells 0,1,4 interleave the seq 1 cells 2,3,5
@@ -554,7 +537,8 @@ static bool test_state_roundtrip(struct llama_model * model, const struct common
LOGV(LOG_LEVEL_INFO, "\n=== Test 8: state blob round-trip ===\n");
- if (llama_decode(ctx.get(), llama_batch_get_one(const_cast<llama_token *>(tokens.data()), (int32_t) tokens.size()))) {
+ common_batch batch = common_batch_get_one(ctx.get(), tokens);
+ if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("\n%s: failed to decode prompt\n", __func__);
return false;
}
@@ -643,12 +627,12 @@ static bool test_state_restore_failure(struct llama_model * model, const struct
}
const auto decode = [&](const llama_tokens & inp, llama_seq_id seq_id, std::vector<float> * logits_out) {
- llama_batch_ptr batch(inp.size(), 0, 1);
+ common_batch batch(ctx.get());
for (size_t i = 0; i < inp.size(); ++i) {
- common_batch_add(batch.get(), inp[i], i, { seq_id }, i == inp.size() - 1);
+ batch.add(inp[i], i, seq_id, i == inp.size() - 1);
}
- if (llama_decode(ctx.get(), batch.get())) {
+ if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("%s: failed to decode on sequence %d\n", __func__, seq_id);
return false;
}
diff --git a/tests/test-state-restore-fragmented.cpp b/tests/test-state-restore-fragmented.cpp
index 33ce6f276..ea3006949 100644
--- a/tests/test-state-restore-fragmented.cpp
+++ b/tests/test-state-restore-fragmented.cpp
@@ -49,15 +49,15 @@ int main(int argc, char ** argv) {
// interleave the 3 sequences:
// 01201230123...
- llama_batch batch = llama_batch_init(params.n_parallel*tokens.size(), 0, 1);
+ common_batch batch(ctx);
for (size_t i = 0; i < tokens.size(); i++) {
for (int s = 0; s < params.n_parallel; ++s) {
- common_batch_add(batch, tokens[i], i, {s}, false);
+ batch.add(tokens[i], i, s, false);
}
}
- batch.logits[batch.n_tokens - 1] = true;
+ batch.set_output(batch.size() - 1, true);
- if (llama_decode(ctx, batch)) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
fprintf(stderr, "%s : failed to decode seq 0\n", __func__);
return 1;
}
@@ -91,7 +91,6 @@ int main(int argc, char ** argv) {
fprintf(stderr, "%s : FAILED to restore seq state into fragmented cache (got %zu, expected %zu)\n",
__func__, nset, seq_state.size());
fprintf(stderr, "%s : This is the bug - state restore fails with fragmented KV cache\n", __func__);
- llama_batch_free(batch);
return 1;
}
fprintf(stderr, "%s : restored state into seq 1, %zu bytes\n", __func__, nset);
@@ -105,13 +104,12 @@ int main(int argc, char ** argv) {
auto next_token = llama_sampler_sample(smpl, ctx, -1);
auto next_token_str = common_token_to_piece(ctx, next_token);
- common_batch_clear(batch);
- common_batch_add(batch, next_token, (int)tokens.size(), {1}, true);
+ batch.clear();
+ batch.add(next_token, (int)tokens.size(), 1, true);
- if (llama_decode(ctx, batch)) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
fprintf(stderr, "%s : failed to decode with restored state\n", __func__);
llama_sampler_free(smpl);
- llama_batch_free(batch);
return 1;
}
@@ -119,7 +117,6 @@ int main(int argc, char ** argv) {
fprintf(stderr, "%s : SUCCESS - state restore works with fragmented KV cache\n", __func__);
llama_sampler_free(smpl);
- llama_batch_free(batch);
return 0;
}
diff --git a/tests/test-thread-safety.cpp b/tests/test-thread-safety.cpp
index d0b5946e2..4fbb1a270 100644
--- a/tests/test-thread-safety.cpp
+++ b/tests/test-thread-safety.cpp
@@ -97,7 +97,6 @@ int main(int argc, char ** argv) {
return;
}
- llama_batch batch = {};
{
auto prompt = common_tokenize(ctx.get(), params.prompt, true);
if (prompt.empty()) {
@@ -105,8 +104,8 @@ int main(int argc, char ** argv) {
failed.store(true);
return;
}
- batch = llama_batch_get_one(prompt.data(), prompt.size());
- if (llama_decode(ctx.get(), batch)) {
+ common_batch batch = common_batch_get_one(ctx.get(), prompt);
+ if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("failed to decode prompt\n");
failed.store(true);
return;
@@ -117,12 +116,7 @@ int main(int argc, char ** argv) {
std::string result = params.prompt;
for (int i = 0; i < params.n_predict; i++) {
- llama_token token;
- if (batch.n_tokens > 0) {
- token = common_sampler_sample(sampler.get(), ctx.get(), batch.n_tokens - 1);
- } else {
- token = llama_vocab_bos(vocab);
- }
+ llama_token token = common_sampler_sample(sampler.get(), ctx.get(), -1);
result += common_token_to_piece(ctx.get(), token);
@@ -130,9 +124,9 @@ int main(int argc, char ** argv) {
break;
}
- batch = llama_batch_get_one(&token, 1);
+ common_batch batch = common_batch_get_one(ctx.get(), &token, 1);
- int ret = llama_decode(ctx.get(), batch);
+ int ret = llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (ret == 1 && i > 0) {
LOG_INF("Context full, stopping generation.\n");
break;
diff --git a/tools/batched-bench/batched-bench.cpp b/tools/batched-bench/batched-bench.cpp
index e2dcd0b2e..c260425d8 100644
--- a/tools/batched-bench/batched-bench.cpp
+++ b/tools/batched-bench/batched-bench.cpp
@@ -76,24 +76,14 @@ int llama_batched_bench(int argc, char ** argv) {
const int32_t n_kv_max = llama_n_ctx(ctx);
- llama_batch batch = llama_batch_init(n_kv_max, 0, 1);
+ common_batch batch(ctx);
// decode in batches of ctx_params.n_batch tokens
- auto decode_helper = [](llama_context * ctx, llama_batch & batch, int32_t n_batch, bool synchronize) {
- for (int32_t i = 0; i < batch.n_tokens; i += n_batch) {
- const int32_t n_tokens = std::min(n_batch, batch.n_tokens - i);
-
- llama_batch batch_view = {
- n_tokens,
- batch.token + i,
- nullptr,
- batch.pos + i,
- batch.n_seq_id + i,
- batch.seq_id + i,
- batch.logits + i,
- };
-
- const int ret = llama_decode(ctx, batch_view);
+ auto decode_helper = [](llama_context * ctx, common_batch & batch, int32_t n_batch, bool synchronize) {
+ for (int32_t i = 0; i < batch.size(); i += n_batch) {
+ const int32_t n_tokens = std::min(n_batch, batch.size() - i);
+
+ const int ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get_sub_batch(i, n_tokens));
if (ret != 0) {
LOG_ERR("failed to decode the batch, n_batch = %d, ret = %d\n", n_batch, ret);
return false;
@@ -110,7 +100,7 @@ int llama_batched_bench(int argc, char ** argv) {
// warm up
{
for (int i = 0; i < 16; ++i) {
- common_batch_add(batch, get_token_rand(), i, { 0 }, false);
+ batch.add(get_token_rand(), i, 0, false);
}
if (!decode_helper(ctx, batch, ctx_params.n_batch, true)) {
@@ -142,11 +132,11 @@ int llama_batched_bench(int argc, char ** argv) {
continue;
}
- common_batch_clear(batch);
+ batch.clear();
for (int j = 0; j < (is_pp_shared ? 1 : pl); ++j) {
for (int i = 0; i < pp; ++i) {
- common_batch_add(batch, get_token_rand(), i, { j }, i == pp - 1);
+ batch.add(get_token_rand(), i, j, i == pp - 1);
}
}
@@ -172,8 +162,8 @@ int llama_batched_bench(int argc, char ** argv) {
if (!params.kv_unified) {
// run one dummy token to apply the memory copy
- common_batch_clear(batch);
- common_batch_add(batch, get_token_rand(), pp + 0, { 0 }, true);
+ batch.clear();
+ batch.add(get_token_rand(), pp + 0, 0, true);
if (!decode_helper(ctx, batch, ctx_params.n_batch, true)) {
LOG_ERR("%s: llama_decode() failed\n", __func__);
llama_free(ctx);
@@ -191,9 +181,9 @@ int llama_batched_bench(int argc, char ** argv) {
// 0 0 0 ... 1 1 1 ... 2 2 2 ... 3 3 3 ...
for (int j = 0; j < pl; ++j) {
for (int i = 0; i < tg; ++i) {
- common_batch_clear(batch);
+ batch.clear();
- common_batch_add(batch, get_token_rand(), pp + i, { j }, true);
+ batch.add(get_token_rand(), pp + i, j, true);
if (!decode_helper(ctx, batch, ctx_params.n_batch, true)) {
LOG_ERR("%s: llama_decode() failed\n", __func__);
@@ -207,10 +197,10 @@ int llama_batched_bench(int argc, char ** argv) {
// decode pattern:
// 0123 0123 0123 ...
for (int i = 0; i < tg; ++i) {
- common_batch_clear(batch);
+ batch.clear();
for (int j = 0; j < pl; ++j) {
- common_batch_add(batch, get_token_rand(), pp + i, { j }, true);
+ batch.add(get_token_rand(), pp + i, j, true);
}
if (!decode_helper(ctx, batch, ctx_params.n_batch, true)) {
@@ -251,7 +241,6 @@ int llama_batched_bench(int argc, char ** argv) {
LOG("\n");
llama_perf_context_print(ctx);
- llama_batch_free(batch);
llama_free(ctx);
llama_model_free(model);
diff --git a/tools/completion/completion.cpp b/tools/completion/completion.cpp
index 941b7399b..718438ef8 100644
--- a/tools/completion/completion.cpp
+++ b/tools/completion/completion.cpp
@@ -525,10 +525,9 @@ int llama_completion(int argc, char ** argv) {
}
if (llama_model_has_encoder(model)) {
- int enc_input_size = embd_inp.size();
- llama_token * enc_input_buf = embd_inp.data();
+ common_batch batch = common_batch_get_one(ctx, embd_inp);
- if (llama_encode(ctx, llama_batch_get_one(enc_input_buf, enc_input_size))) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_ENCODE, batch.get())) {
LOG_ERR("%s : failed to eval\n", __func__);
return 1;
}
diff --git a/tools/cvector-generator/cvector-generator.cpp b/tools/cvector-generator/cvector-generator.cpp
index 558c37e61..af05031f2 100644
--- a/tools/cvector-generator/cvector-generator.cpp
+++ b/tools/cvector-generator/cvector-generator.cpp
@@ -346,7 +346,8 @@ static bool cb_eval(struct ggml_tensor * t, bool ask, void * user_data) {
static bool get_hidden_layers(llama_context * ctx, std::vector<llama_token> & tokens) {
llama_memory_clear(llama_get_memory(ctx), true);
- if (llama_decode(ctx, llama_batch_get_one(tokens.data(), tokens.size()))) {
+ common_batch batch = common_batch_get_one(ctx, tokens);
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
fprintf(stderr, "%s : failed to eval\n", __func__);
return false;
}
diff --git a/tools/imatrix/imatrix.cpp b/tools/imatrix/imatrix.cpp
index f5fee6218..7baa46968 100644
--- a/tools/imatrix/imatrix.cpp
+++ b/tools/imatrix/imatrix.cpp
@@ -845,7 +845,7 @@ static bool compute_imatrix(llama_context * ctx, const common_params & params, c
GGML_ASSERT(n_batch < n_ctx || n_batch % n_ctx == 0);
GGML_ASSERT(params.n_ctx == n_seq * n_ctx);
- llama_batch batch = llama_batch_init(std::min(n_batch, n_ctx*n_seq), 0, 1);
+ common_batch batch(ctx);
std::vector<float> logits;
if (params.compute_ppl && num_batches > 1) {
@@ -872,7 +872,7 @@ static bool compute_imatrix(llama_context * ctx, const common_params & params, c
const int batch_size = std::min(end - batch_start, n_batch);
// clear the batch
- common_batch_clear(batch);
+ batch.clear();
for (int seq = 0; seq < n_seq_batch; seq++) {
int seq_start = batch_start + seq*n_ctx;
@@ -889,16 +889,15 @@ static bool compute_imatrix(llama_context * ctx, const common_params & params, c
// and also for the perplexity calculation.
// TODO: only get outputs when (params.process_output || params.compute_ppl)
// (not possible when this skips FFN computation of the last layer)
- common_batch_add(batch, tokens[seq_start + k], j*n_batch + k, { seq }, true);
+ batch.add(tokens[seq_start + k], j*n_batch + k, seq, true);
}
// restore the original token in case it was set to BOS
tokens[seq_start] = token_org;
}
- if (llama_decode(ctx, batch)) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("%s : failed to eval\n", __func__);
- llama_batch_free(batch);
return false;
}
@@ -960,7 +959,6 @@ static bool compute_imatrix(llama_context * ctx, const common_params & params, c
}
}
- llama_batch_free(batch);
return true;
}
diff --git a/tools/llama-bench/llama-bench.cpp b/tools/llama-bench/llama-bench.cpp
index 70e15d044..fd64b4c44 100644
--- a/tools/llama-bench/llama-bench.cpp
+++ b/tools/llama-bench/llama-bench.cpp
@@ -2182,7 +2182,8 @@ static bool test_prompt(llama_context * ctx, int n_prompt, int n_batch, int n_th
for (int i = 1; i < n_tokens; i++) {
tokens[i] = std::rand() % n_vocab;
}
- int res = llama_decode(ctx, llama_batch_get_one(tokens.data(), n_tokens));
+ common_batch batch = common_batch_get_one(ctx, tokens.data(), n_tokens);
+ int res = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (res != 0) {
fprintf(stderr, "%s: failed to decode prompt batch, res = %d\n", __func__, res);
return false;
@@ -2203,8 +2204,13 @@ static bool test_gen(llama_context * ctx, int n_gen, int n_threads) {
llama_token token = llama_vocab_get_add_bos(vocab) ? llama_vocab_bos(vocab) : std::rand() % n_vocab;
+ common_batch batch(ctx);
+ llama_pos pos = llama_memory_seq_pos_max(llama_get_memory(ctx), 0) + 1;
+
for (int i = 0; i < n_gen; i++) {
- int res = llama_decode(ctx, llama_batch_get_one(&token, 1));
+ batch.clear();
+ batch.add(token, pos++, 0, true);
+ int res = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
if (res != 0) {
fprintf(stderr, "%s: failed to decode generation batch, res = %d\n", __func__, res);
return false;
diff --git a/tools/perplexity/perplexity.cpp b/tools/perplexity/perplexity.cpp
index ba41287d8..601361acd 100644
--- a/tools/perplexity/perplexity.cpp
+++ b/tools/perplexity/perplexity.cpp
@@ -366,21 +366,20 @@ static results_perplexity perplexity_v2(llama_context * ctx, const common_params
// clear the KV cache
llama_memory_clear(llama_get_memory(ctx), true);
- llama_batch batch = llama_batch_init(n_batch, 0, 1);
+ common_batch batch(ctx);
for (int j = 0; j < num_batches; ++j) {
const int batch_start = start + j * n_batch;
const int batch_size = std::min(end - batch_start, n_batch);
- common_batch_clear(batch);
+ batch.clear();
for (int i = 0; i < batch_size; i++) {
- common_batch_add(batch, tokens[batch_start + i], j*n_batch + i, {0}, true);
+ batch.add(tokens[batch_start + i], j*n_batch + i, 0, true);
}
//LOG_DBG(" Batch %d: starts at %d, size is %d, n_past is %d\n",j,batch_start,batch_size,j * n_batch);
- if (llama_decode(ctx, batch)) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
//LOG_ERR("%s : failed to eval\n", __func__);
- llama_batch_free(batch);
return {tokens, -1, logit_history, prob_history};
}
@@ -400,7 +399,6 @@ static results_perplexity perplexity_v2(llama_context * ctx, const common_params
}
}
- llama_batch_free(batch);
const auto t_end = std::chrono::high_resolution_clock::now();
@@ -507,7 +505,7 @@ static results_perplexity perplexity(llama_context * ctx, const common_params &
GGML_ASSERT(n_batch < n_ctx || n_batch % n_ctx == 0);
GGML_ASSERT(params.n_ctx == n_seq * n_ctx);
- llama_batch batch = llama_batch_init(std::min(n_batch, n_ctx*n_seq), 0, 1);
+ common_batch batch(ctx);
std::vector<float> logits;
if (num_batches > 1) {
@@ -558,7 +556,7 @@ static results_perplexity perplexity(llama_context * ctx, const common_params &
int n_outputs = 0;
- batch.n_tokens = 0;
+ batch.clear();
for (int seq = 0; seq < n_seq_batch; seq++) {
int seq_start = batch_start + seq*n_ctx;
@@ -571,22 +569,17 @@ static results_perplexity perplexity(llama_context * ctx, const common_params &
}
for (int k = 0; k < batch_size; ++k) {
- const int idx = seq*n_ctx + k;
- batch.token [idx] = tokens[seq_start + k];
- batch.pos [idx] = j*n_batch + k;
- batch.n_seq_id[idx] = 1;
- batch.seq_id [idx][0] = seq;
- batch.logits [idx] = batch.pos[idx] >= first ? 1 : 0;
-
- n_outputs += batch.logits[idx] != 0;
+ const llama_pos pos = j*n_batch + k;
+ const bool need_logits = pos >= first;
+ batch.add(tokens[seq_start + k], pos, seq, need_logits);
+ n_outputs += need_logits;
}
- batch.n_tokens += batch_size;
// restore the original token in case it was set to BOS
tokens[seq_start] = token_org;
}
- if (llama_decode(ctx, batch)) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_INF("%s : failed to decode\n", __func__);
return {tokens, -1, logit_history, prob_history};
}
@@ -656,35 +649,24 @@ static results_perplexity perplexity(llama_context * ctx, const common_params &
LOG_ERR("Unexpected negative standard deviation of log(prob)\n");
}
- llama_batch_free(batch);
return {tokens, ppl, logit_history, prob_history};
}
-static bool decode_helper(llama_context * ctx, llama_batch & batch, std::vector<float> & batch_logits, int n_batch, int n_vocab) {
+static bool decode_helper(llama_context * ctx, common_batch & batch, std::vector<float> & batch_logits, int n_batch, int n_vocab) {
int prev_outputs = 0;
- for (int i = 0; i < (int) batch.n_tokens; i += n_batch) {
- const int n_tokens = std::min<int>(n_batch, batch.n_tokens - i);
-
- llama_batch batch_view = {
- n_tokens,
- batch.token + i,
- nullptr,
- batch.pos + i,
- batch.n_seq_id + i,
- batch.seq_id + i,
- batch.logits + i,
- };
+ for (int i = 0; i < batch.size(); i += n_batch) {
+ const int n_tokens = std::min<int>(n_batch, batch.size() - i);
- const int ret = llama_decode(ctx, batch_view);
+ const int ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get_sub_batch(i, n_tokens));
if (ret != 0) {
LOG_ERR("failed to decode the batch, n_batch = %d, ret = %d\n", n_batch, ret);
return false;
}
int n_outputs = 0;
- for (int i = 0; i < n_tokens; ++i) {
- n_outputs += batch_view.logits[i] != 0;
+ for (int j = i; j < i + n_tokens; ++j) {
+ n_outputs += batch.tokens[j].output;
}
memcpy(batch_logits.data() + size_t(prev_outputs)*n_vocab, llama_get_logits(ctx), size_t(n_outputs)*n_vocab*sizeof(float));
@@ -866,7 +848,7 @@ static void hellaswag_score(llama_context * ctx, const common_params & params) {
const int max_tasks_per_batch = 32;
const int max_seq = std::min(4*max_tasks_per_batch, (int) llama_n_seq_max(ctx));
- llama_batch batch = llama_batch_init(n_ctx, 0, 4);
+ common_batch batch(ctx);
std::vector<float> tok_logits(n_vocab);
// TODO: this could be made smaller; it's currently the worst-case size
@@ -882,7 +864,7 @@ static void hellaswag_score(llama_context * ctx, const common_params & params) {
size_t i1 = i0;
size_t i_logits = 0; // this tells us how many logits were needed before this point in the batch
- common_batch_clear(batch);
+ batch.clear();
// batch as much tasks as possible into the available context
// each task has 4 unique sequence ids - one for each ending
@@ -898,9 +880,9 @@ static void hellaswag_score(llama_context * ctx, const common_params & params) {
}
for (size_t i = 0; i < hs_cur.common_prefix; ++i) {
- common_batch_add(batch, hs_cur.seq_tokens[0][i], i, { s0 + 0, s0 + 1, s0 + 2, s0 + 3 }, false);
+ batch.add(hs_cur.seq_tokens[0][i], i, { s0 + 0, s0 + 1, s0 + 2, s0 + 3 }, false);
}
- batch.logits[batch.n_tokens - 1] = true; // we need logits for the last token of the common prefix
+ batch.set_output(batch.size() - 1, true); // we need logits for the last token of the common prefix
n_logits += 1;
for (int s = 0; s < 4; ++s) {
@@ -908,7 +890,7 @@ static void hellaswag_score(llama_context * ctx, const common_params & params) {
// TODO: don't evaluate the last token of each sequence
for (size_t i = hs_cur.common_prefix; i < seq_tokens_size; ++i) {
const bool needs_logits = i < seq_tokens_size - 1;
- common_batch_add(batch, hs_cur.seq_tokens[s][i], i, { s0 + s }, needs_logits);
+ batch.add(hs_cur.seq_tokens[s][i], i, s0 + s, needs_logits);
n_logits += needs_logits;
}
}
@@ -1009,7 +991,6 @@ static void hellaswag_score(llama_context * ctx, const common_params & params) {
i0 = i1 - 1;
}
- llama_batch_free(batch);
LOG("\n");
}
@@ -1164,7 +1145,7 @@ static void winogrande_score(llama_context * ctx, const common_params & params)
const int max_tasks_per_batch = 128;
const int max_seq = std::min(2*max_tasks_per_batch, (int) llama_n_seq_max(ctx));
- llama_batch batch = llama_batch_init(n_ctx, 0, 2);
+ common_batch batch(ctx);
std::vector<float> tok_logits(n_vocab);
// TODO: this could be made smaller; it's currently the worst-case size
@@ -1183,7 +1164,7 @@ static void winogrande_score(llama_context * ctx, const common_params & params)
size_t i1 = i0;
size_t i_logits = 0;
- common_batch_clear(batch);
+ batch.clear();
while (n_cur + (int) data[i1].required_tokens <= n_ctx) {
int n_logits = 0;
@@ -1193,15 +1174,15 @@ static void winogrande_score(llama_context * ctx, const common_params & params)
}
for (size_t i = 0; i < data[i1].common_prefix; ++i) {
- common_batch_add(batch, data[i1].seq_tokens[0][i], i, { s0 + 0, s0 + 1 }, false);
+ batch.add(data[i1].seq_tokens[0][i], i, { s0 + 0, s0 + 1 }, false);
}
- batch.logits[batch.n_tokens - 1] = true;
+ batch.set_output(batch.size() - 1, true);
n_logits += 1;
for (int s = 0; s < 2; ++s) {
// TODO: end before the last token, no need to predict past the end of the sequences
for (size_t i = data[i1].common_prefix; i < data[i1].seq_tokens[s].size(); ++i) {
- common_batch_add(batch, data[i1].seq_tokens[s][i], i, { s0 + s }, true);
+ batch.add(data[i1].seq_tokens[s][i], i, s0 + s, true);
n_logits += 1;
}
}
@@ -1518,7 +1499,7 @@ static void multiple_choice_score(llama_context * ctx, const common_params & par
const int max_tasks_per_batch = 32;
const int max_seq = std::min(4*max_tasks_per_batch, (int) llama_n_seq_max(ctx));
- llama_batch batch = llama_batch_init(n_ctx, 0, max_seq);
+ common_batch batch(ctx);
std::vector<float> tok_logits(n_vocab);
std::vector<float> batch_logits(size_t(n_ctx)*n_vocab);
@@ -1538,7 +1519,7 @@ static void multiple_choice_score(llama_context * ctx, const common_params & par
size_t i1 = i0;
size_t i_logits = 0; // this tells us how many logits were needed before this point in the batch
- common_batch_clear(batch);
+ batch.clear();
// batch as much tasks as possible into the available context
// each task has 4 unique sequence ids - one for each ending
@@ -1568,9 +1549,9 @@ static void multiple_choice_score(llama_context * ctx, const common_params & par
for (size_t i = 0; i < cur_task.common_prefix; ++i) {
//llama_batch_add(batch, cur_task.seq_tokens[0][i], i, { s0 + 0, s0 + 1, s0 + 2, s0 + 3}, false);
- common_batch_add(batch, cur_task.seq_tokens[0][i], i, batch_indeces, false);
+ batch.add(cur_task.seq_tokens[0][i], i, batch_indeces, false);
}
- batch.logits[batch.n_tokens - 1] = true; // we need logits for the last token of the common prefix
+ batch.set_output(batch.size() - 1, true); // we need logits for the last token of the common prefix
n_logits += 1;
for (int s = 0; s < int(cur_task.seq_tokens.size()); ++s) {
@@ -1578,7 +1559,7 @@ static void multiple_choice_score(llama_context * ctx, const common_params & par
// TODO: don't evaluate the last token of each sequence
for (size_t i = cur_task.common_prefix; i < seq_tokens_size; ++i) {
const bool needs_logits = i < seq_tokens_size - 1;
- common_batch_add(batch, cur_task.seq_tokens[s][i], i, { s0 + s }, needs_logits);
+ batch.add(cur_task.seq_tokens[s][i], i, s0 + s, needs_logits);
n_logits += needs_logits;
}
}
@@ -1677,7 +1658,6 @@ static void multiple_choice_score(llama_context * ctx, const common_params & par
i0 = i1 - 1;
}
- llama_batch_free(batch);
if (n_done < 100 && (params.multiple_choice_tasks != 0 && params.multiple_choice_tasks < (size_t)n_task)) return;
@@ -1753,7 +1733,7 @@ static void kl_divergence(llama_context * ctx, const common_params & params) {
const bool add_bos = llama_vocab_get_add_bos(vocab);
GGML_ASSERT(!llama_vocab_get_add_eos(vocab));
- llama_batch batch = llama_batch_init(std::min(n_batch, static_cast<int>(n_ctx)*n_seq), 0, 1);
+ common_batch batch(ctx);
std::vector<uint16_t> log_probs_uint16(size_t(n_ctx - 1 - n_ctx/2) * nv);
std::vector<float> kld_values(size_t(n_ctx - 1 - n_ctx/2)*n_chunk);
@@ -1808,7 +1788,7 @@ static void kl_divergence(llama_context * ctx, const common_params & params) {
int n_outputs = 0;
- common_batch_clear(batch);
+ batch.clear();
for (int seq = 0; seq < n_seq_batch; seq++) {
int seq_start = batch_start + seq*n_ctx;
@@ -1823,7 +1803,7 @@ static void kl_divergence(llama_context * ctx, const common_params & params) {
for (int k = 0; k < batch_size; ++k) {
const int pos = j*n_batch + k;
const bool need_logits = pos >= first;
- common_batch_add(batch, tokens[seq_start + k], pos, { seq }, need_logits);
+ batch.add(tokens[seq_start + k], pos, seq, need_logits);
n_outputs += need_logits;
}
@@ -1831,9 +1811,8 @@ static void kl_divergence(llama_context * ctx, const common_params & params) {
tokens[seq_start] = token_org;
}
- if (llama_decode(ctx, batch)) {
+ if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("%s : failed to decode\n", __func__);
- llama_batch_free(batch);
return;
}
@@ -1862,7 +1841,6 @@ static void kl_divergence(llama_context * ctx, const common_params & params) {
for (int seq = 0; seq < n_seq_batch; seq++) {
if (in.read((char *)log_probs_uint16.data(), log_probs_uint16.size()*sizeof(uint16_t)).fail()) {
LOG_ERR("%s: failed reading log-probs for chunk %d\n", __func__, i + seq);
- llama_batch_free(batch);
return;
}
@@ -1904,7 +1882,6 @@ static void kl_divergence(llama_context * ctx, const common_params & params) {
logits.clear();
}
- llama_batch_free(batch);
LOG("\n");
if (kld.count < 100) return; // we do not wish to do statistics on so few values
diff --git a/tools/results/results.cpp b/tools/results/results.cpp
index f2179ed27..2d6479482 100644
--- a/tools/results/results.cpp
+++ b/tools/results/results.cpp
@@ -32,14 +32,12 @@ static std::vector<float> get_logits(
const uint32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model));
const uint32_t n_ctx = llama_n_ctx(lctx);
const uint32_t n_tokens = tokens.size();
- llama_batch batch = llama_batch_init(n_ctx, 0, 1);
+ common_batch batch(lctx);
GGML_ASSERT(n_tokens <= n_ctx);
for (uint32_t pos = 0; pos < n_tokens; pos++) {
- common_batch_add(batch, tokens[pos], pos, {0}, true);
+ batch.add(tokens[pos], pos, 0, true);
}
- batch.n_tokens = n_tokens;
- if (llama_decode(lctx, batch)) {
- llama_batch_free(batch);
+ if (llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
throw std::runtime_error("failed to decode batch");
}
@@ -51,7 +49,6 @@ static std::vector<float> get_logits(
ret.push_back(logits_ith[j]);
}
}
- llama_batch_free(batch);
return ret;
}