mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-27 21:46:57 +02:00
llama: add llama_batch_ext (#24669)
* (wip) add llama_batch_ext * wip * updated design * updated impl * change signature * unused var * demo common_prompt_batch_decode * fix pos * tmp disable test-batch-alloc * fix compat * nits: add const * no more pos_max * add comment about llama_batch_ext_set_embd_state * handle n_embd_out properly * rename api --> embd_token * llama_embd * stub llama_batch_ext_set_embd_state * support both token + embd + state in batch * llama_batch_ext_add_embd * upstream some changes * nits * fix test-batch-alloc * add test for compat
This commit is contained in:
@@ -24,7 +24,12 @@ struct llama_adapter_lora_deleter {
|
||||
void operator()(llama_adapter_lora * adapter) { llama_adapter_lora_free(adapter); }
|
||||
};
|
||||
|
||||
struct llama_batch_ext_deleter {
|
||||
void operator()(llama_batch_ext * batch) { llama_batch_ext_free(batch); }
|
||||
};
|
||||
|
||||
typedef std::unique_ptr<llama_model, llama_model_deleter> llama_model_ptr;
|
||||
typedef std::unique_ptr<llama_context, llama_context_deleter> llama_context_ptr;
|
||||
typedef std::unique_ptr<llama_sampler, llama_sampler_deleter> llama_sampler_ptr;
|
||||
typedef std::unique_ptr<llama_adapter_lora, llama_adapter_lora_deleter> llama_adapter_lora_ptr;
|
||||
typedef std::unique_ptr<llama_batch_ext, llama_batch_ext_deleter> llama_batch_ext_ptr;
|
||||
|
||||
@@ -293,6 +293,11 @@ extern "C" {
|
||||
LLAMA_MODEL_META_KEY_SAMPLING_MIROSTAT_ETA,
|
||||
};
|
||||
|
||||
enum llama_process_type {
|
||||
LLAMA_PROCESS_TYPE_ENCODE,
|
||||
LLAMA_PROCESS_TYPE_DECODE,
|
||||
};
|
||||
|
||||
struct llama_model_kv_override {
|
||||
enum llama_model_kv_override_type tag;
|
||||
|
||||
@@ -999,6 +1004,91 @@ extern "C" {
|
||||
struct llama_context * ctx,
|
||||
struct llama_batch batch);
|
||||
|
||||
//
|
||||
// Extended batch API
|
||||
//
|
||||
|
||||
struct llama_batch_ext;
|
||||
|
||||
struct llama_embd {
|
||||
const float * data;
|
||||
size_t n_rows; // number of embedding rows in data
|
||||
size_t n_embd; // size of one row
|
||||
};
|
||||
|
||||
LLAMA_API struct llama_batch_ext * llama_batch_ext_init (struct llama_context * ctx);
|
||||
LLAMA_API void llama_batch_ext_free (struct llama_batch_ext * batch);
|
||||
LLAMA_API void llama_batch_ext_clear(struct llama_batch_ext * batch);
|
||||
|
||||
// Add an input token to the batch, with default values:
|
||||
// id = LLAMA_TOKEN_NULL
|
||||
// embd = nullptr
|
||||
// pos = not set, the caller must set it with llama_batch_ext_set_pos()
|
||||
// Returns the batch index (>= 0)
|
||||
// On error:
|
||||
// -1: batch is full
|
||||
// -2: token is invalid (id == LLAMA_TOKEN_NULL or invalid embd)
|
||||
// -3: invalid sequence id
|
||||
LLAMA_API int32_t llama_batch_ext_add (struct llama_batch_ext * batch, llama_seq_id seq_id);
|
||||
|
||||
// Add an input token to the batch, with a specified token ID or token embedding
|
||||
LLAMA_API int32_t llama_batch_ext_add_token(struct llama_batch_ext * batch, llama_seq_id seq_id, llama_token id);
|
||||
LLAMA_API int32_t llama_batch_ext_add_embd (struct llama_batch_ext * batch, llama_seq_id seq_id, struct llama_embd embd);
|
||||
|
||||
// Add the token at index idx in the batch to another sequence id. The position will stays the same.
|
||||
// Note: this should be called before other _set() functions
|
||||
LLAMA_API bool llama_batch_ext_add_seq(
|
||||
struct llama_batch_ext * batch,
|
||||
int32_t idx,
|
||||
llama_seq_id seq_id);
|
||||
|
||||
// Set the token embedding for the token at index idx in the batch
|
||||
// use it after llama_batch_ext_add_token() to have an entry with both a token id and an embedding
|
||||
LLAMA_API bool llama_batch_ext_set_embd_token(
|
||||
struct llama_batch_ext * batch,
|
||||
int32_t idx,
|
||||
struct llama_embd embd);
|
||||
|
||||
// Set the "state" embedding for the token at index idx in the batch
|
||||
// "state" here means extra hidden state carried over from a previous stage, e.g.:
|
||||
// - MTP: state from N layers of the target model
|
||||
// - Qwen3 VL (deepstack): state from N layers of the vision encoder
|
||||
LLAMA_API bool llama_batch_ext_set_embd_state(
|
||||
struct llama_batch_ext * batch,
|
||||
int32_t idx,
|
||||
struct llama_embd embd);
|
||||
|
||||
// Set if output embeddings should be available for the token at index idx in the batch
|
||||
// Note: for now, this is equivalent to setting the output logits
|
||||
LLAMA_API bool llama_batch_ext_set_output_embd(
|
||||
struct llama_batch_ext * batch,
|
||||
int32_t idx,
|
||||
bool value);
|
||||
|
||||
// Set output logits for the token at index idx in the batch
|
||||
// Note: for now, this is equivalent to setting the output embd
|
||||
LLAMA_API bool llama_batch_ext_set_output_logits(
|
||||
struct llama_batch_ext * batch,
|
||||
int32_t idx,
|
||||
bool value);
|
||||
|
||||
// Set custom position for the token at index idx in the batch
|
||||
// For M-RoPE models:
|
||||
// - Embedding tokens must have multiple positions per token
|
||||
// - Text token only requires one single position per token
|
||||
LLAMA_API bool llama_batch_ext_set_pos(
|
||||
struct llama_batch_ext * batch,
|
||||
int32_t idx,
|
||||
const llama_pos * pos);
|
||||
|
||||
// TODO: implement get_embeddings() and get_logits() for llama_batch_ext
|
||||
|
||||
// Return values are the same as llama_decode()
|
||||
LLAMA_API int32_t llama_process(
|
||||
struct llama_context * ctx,
|
||||
enum llama_process_type type,
|
||||
struct llama_batch_ext * batch);
|
||||
|
||||
// Set the number of threads used for decoding
|
||||
// n_threads is the number of threads used for generation (single token)
|
||||
// n_threads_batch is the number of threads used for prompt and batch processing (multiple tokens)
|
||||
|
||||
Reference in New Issue
Block a user