Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
29 commits
Select commit Hold shift + click to select a range
0dfa5e1
(wip) add llama_batch_ext
ngxson Jun 15, 2026
c7c5468
wip
ngxson Jun 15, 2026
c0fb071
updated design
ngxson Jun 16, 2026
3132371
updated impl
ngxson Jun 17, 2026
32343b2
Merge branch 'master' into xsn/llama_batch_ext
ngxson Jul 13, 2026
bf372b3
change signature
ngxson Jul 13, 2026
901ed68
unused var
ngxson Jul 13, 2026
231af77
demo common_prompt_batch_decode
ngxson Jul 13, 2026
4cd8c26
fix pos
ngxson Jul 13, 2026
e4c474f
tmp disable test-batch-alloc
ngxson Jul 13, 2026
4ba39e5
fix compat
ngxson Jul 13, 2026
7db50e9
Merge remote-tracking branch 'upstream/master' into xsn/llama_batch_ext
ngxson Aug 12, 2026
7d626f5
nits: add const
ngxson Aug 12, 2026
3b86110
no more pos_max
ngxson Aug 13, 2026
05a8c11
Merge branch 'master' into xsn/llama_batch_ext
ngxson Aug 17, 2026
51b471f
add comment about llama_batch_ext_set_embd_state
ngxson Aug 17, 2026
95956ee
handle n_embd_out properly
ngxson Aug 17, 2026
e640410
rename api --> embd_token
ngxson Aug 17, 2026
6d00861
llama_embd
ngxson Aug 17, 2026
69258fb
stub llama_batch_ext_set_embd_state
ngxson Aug 17, 2026
2b50dab
support both token + embd + state in batch
ngxson Aug 17, 2026
4bcbdff
Merge branch 'master' into xsn/llama_batch_ext
ngxson Aug 28, 2026
de88a4a
llama_batch_ext_add_embd
ngxson Aug 28, 2026
4dd6864
Merge remote-tracking branch 'upstream/master' into xsn/llama_batch_ext
ngxson Sep 10, 2026
e67da54
upstream some changes
ngxson Sep 12, 2026
d48a14d
Merge branch 'master' into xsn/llama_batch_ext
ngxson Sep 23, 2026
12691fb
nits
ngxson Sep 23, 2026
82a63e3
fix test-batch-alloc
ngxson Sep 23, 2026
4070433
add test for compat
ngxson Sep 23, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 30 additions & 7 deletions common/common.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2198,9 +2198,28 @@ bool common_replay_last_token(struct llama_context * ctx, llama_token last_token
return true;
}

llama_batch_ext_ptr common_batch_ext_get_one(llama_context * ctx, const llama_tokens & tokens) {
llama_batch_ext_ptr batch(llama_batch_ext_init(ctx));

auto mem = llama_get_memory(ctx);
llama_pos pos = mem ? llama_memory_seq_pos_max(mem, 0) + 1 : 0;

for (size_t i = 0; i < tokens.size(); ++i) {
const int32_t idx = llama_batch_ext_add_token(batch.get(), 0, tokens[i]);
llama_batch_ext_set_pos(batch.get(), idx, &pos);
pos++;
}

if (!tokens.empty()) {
llama_batch_ext_set_output_logits(batch.get(), (int32_t) tokens.size() - 1, true);
}

return batch;
}

bool common_prompt_batch_decode(
struct llama_context * ctx,
const std::vector<llama_token> & all_tokens,
const llama_tokens & all_tokens,
int n_new,
int & n_past,
int n_batch,
Expand All @@ -2221,7 +2240,9 @@ bool common_prompt_batch_decode(
// Memory implementations in recurrent/hybrid models don't support removing tokens from their
// memory, so we can't just remove the last token from the memory and replay the last token which
// is the reason for this logic.
if (llama_decode(ctx, llama_batch_get_one(const_cast<llama_token*>(all_tokens.data() + offset), n_tokens_before_last))) {
llama_tokens prefix_tokens(all_tokens.begin() + offset, all_tokens.begin() + offset + n_tokens_before_last);
llama_batch_ext_ptr batch_prefix = common_batch_ext_get_one(ctx, prefix_tokens);
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_prefix.get())) {
COM_ERR("%s", "failed to eval\n");
return false;
}
Expand All @@ -2231,17 +2252,19 @@ bool common_prompt_batch_decode(
COM_INF("saved session before last token to %s, n_new = %zu\n", state_path.data(), all_tokens.size());

llama_token last_token = all_tokens.back();
llama_batch batch = llama_batch_get_one(&last_token, 1);
int32_t pos = n_past;
batch.pos = &pos;
llama_batch_ext_ptr batch_last = common_batch_ext_get_one(ctx, { last_token });
llama_pos pos = n_past;
llama_batch_ext_set_pos(batch_last.get(), 0, &pos);

if (llama_decode(ctx, batch)) {
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_last.get())) {
COM_ERR("%s", "failed to eval last token\n");
return false;
}
n_past++;
} else {
if (llama_decode(ctx, llama_batch_get_one(const_cast<llama_token*>(all_tokens.data() + offset), n_new))) {
llama_tokens new_tokens(all_tokens.begin() + offset, all_tokens.begin() + offset + n_new);
llama_batch_ext_ptr batch = common_batch_ext_get_one(ctx, new_tokens);
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
COM_ERR("%s", "failed to eval\n");
return false;
}
Expand Down
6 changes: 5 additions & 1 deletion common/common.h
Original file line number Diff line number Diff line change
Expand Up @@ -1021,14 +1021,18 @@ void common_batch_add(
const std::vector<llama_seq_id> & seq_ids,
bool logits);

// create a single-sequence batch from a list of tokens
// last token always have output_logits set to true
llama_batch_ext_ptr common_batch_ext_get_one(struct llama_context * ctx, const llama_tokens & tokens);

// 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
// compatibility with all memory types. Recurrent/hybrid models cannot remove
// tokens from memory, so this approach works across all model architectures.
bool common_prompt_batch_decode(
struct llama_context * ctx,
const std::vector<llama_token> & all_tokens,
const llama_tokens & all_tokens,
int n_new,
int & n_past,
int n_batch,
Expand Down
5 changes: 5 additions & 0 deletions include/llama-cpp.h
Original file line number Diff line number Diff line change
Expand Up @@ -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;
90 changes: 90 additions & 0 deletions include/llama.h
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand Down Expand Up @@ -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
};

Comment on lines +1013 to +1018

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

to be discussed: n_rows or n_tokens ?

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Rows seems better.

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);

Comment on lines +1052 to +1060

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

note: llama_batch_ext_set_embd_state is TODO

// 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)
Expand Down
Loading