diff --git a/common/common.cpp b/common/common.cpp index 598a97d105..401de1dc2a 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 & 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()); + 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 & 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]); } - 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 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); +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 bool has_token = batch.token != nullptr; - const bool has_embd = batch.embd != nullptr; + llama_batch_ext * res = batch.get(); + llama_batch_ext_clear(res); - const size_t n_embd = llama_model_n_embd_inp(llama_get_model(ctx)); - - // positions continue from the memory when none are given - auto * mem = llama_get_memory(ctx); - std::vector 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; - - 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]; - } 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]; - } - } - - 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 }; + for (int32_t i = off; i < off + n; ++i) { + const token & t = tokens[i]; 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); + 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 { - idx = res.add_embd(embd, pos, seq_id, output); + 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); - for (int32_t s = 1; s < n_sid; ++s) { - llama_batch_ext_add_seq(res.get(), idx, batch.seq_id[i][s]); + 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); + } + } + 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 e95eb2fd0e..dcc5ec1aef 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 & 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 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 seq_ids_extra; // see add_seq() }; std::vector 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 & 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 82e9e92238..b1244d9a67 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> 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(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 211fcdabd1..d46b21eb71 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 830e45f5af..9da844c7e1 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 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 i_batch(n_parallel, batch.n_tokens - 1); + std::vector 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 761e7a2db5..18a264b63a 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 97d6b69449..0ff57d953c 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 @@ -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 f6a20ef9d0..a59f04cae9 100644 --- a/examples/embedding/embedding.cpp +++ b/examples/embedding/embedding.cpp @@ -27,27 +27,27 @@ static std::vector split_lines(const std::string & s, const std::st return lines; } -static void batch_add_seq(llama_batch & batch, const std::vector & tokens, llama_seq_id seq_id) { +static void batch_add_seq(common_batch & batch, const std::vector & 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 4ce8d600b1..703ce130b3 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 409fd25c18..ddbda79939 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 03ab96cfd8..ca6f7172d9 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 b7f5c6de86..62772814fa 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 6621058655..004062c576 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 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 4b74540f07..de877ca365 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); + const int32_t n_tokens = std::min(n_batch, batch.size() - 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); + 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 8440a2bf77..9ac8a01702 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 7d93ab1172..8793e57519 100644 --- a/examples/retrieval/retrieval.cpp +++ b/examples/retrieval/retrieval.cpp @@ -75,30 +75,30 @@ static std::vector chunk_file(const std::string & filename, int chunk_siz return chunks; } -static void batch_add_seq(llama_batch & batch, const std::vector & tokens, llama_seq_id seq_id) { +static void batch_add_seq(common_batch & batch, const std::vector & 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 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 30a0966e07..0cad9652c0 100644 --- a/examples/simple-chat/simple-chat.cpp +++ b/examples/simple-chat/simple-chat.cpp @@ -6,6 +6,17 @@ #include #include +// 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(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 982a4d8600..ebb1969fb5 100644 --- a/examples/simple/simple.cpp +++ b/examples/simple/simple.cpp @@ -5,6 +5,17 @@ #include #include +// 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 81aa106f14..08a1f2a887 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 17071aa054..1bb47594b9 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()); 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 c23e7248d5..56736ac468 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 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 & 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 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 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 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 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 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 65d444ecf4..55e4f9df46 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 decode_prefill(llama_model * model, llama_context * lctx, const std::vector & 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 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 decode_gen(llama_model * model, llama_context * lctx, const std::vector & 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 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 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 298073d06c..015e3414b7 100644 --- a/tests/test-llama-archs.cpp +++ b/tests/test-llama-archs.cpp @@ -510,20 +510,17 @@ static std::vector 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 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 4ad0d6f9ef..f8eda55c8b 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 & 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 dee5e17be6..02984d185a 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 & 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 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 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(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 * 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 33ce6f2763..ea30069499 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 d0b5946e2c..4fbb1a2700 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 e2dcd0b2e7..c260425d83 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); + 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); - 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); + 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 941b7399b2..718438ef86 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 558c37e612..af05031f23 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 & 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 f5fee62184..7baa46968a 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 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 70e15d0449..fd64b4c442 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 ba41287d8e..601361acdc 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 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 & batch_logits, int n_batch, int n_vocab) { +static bool decode_helper(llama_context * ctx, common_batch & batch, std::vector & 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(n_batch, batch.n_tokens - i); + for (int i = 0; i < batch.size(); i += n_batch) { + const int n_tokens = std::min(n_batch, batch.size() - 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); + 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 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 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 tok_logits(n_vocab); std::vector 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(n_ctx)*n_seq), 0, 1); + common_batch batch(ctx); std::vector log_probs_uint16(size_t(n_ctx - 1 - n_ctx/2) * nv); std::vector 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 f2179ed275..2d6479482e 100644 --- a/tools/results/results.cpp +++ b/tools/results/results.cpp @@ -32,14 +32,12 @@ static std::vector 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 get_logits( ret.push_back(logits_ith[j]); } } - llama_batch_free(batch); return ret; }