Repository navigation
llama: add llama_batch_ext #24669
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
+1,201
−253
Merged
llama: add llama_batch_ext #24669
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 c7c5468
wip
ngxson c0fb071
updated design
ngxson 3132371
updated impl
ngxson 32343b2
Merge branch 'master' into xsn/llama_batch_ext
ngxson bf372b3
change signature
ngxson 901ed68
unused var
ngxson 231af77
demo common_prompt_batch_decode
ngxson 4cd8c26
fix pos
ngxson e4c474f
tmp disable test-batch-alloc
ngxson 4ba39e5
fix compat
ngxson 7db50e9
Merge remote-tracking branch 'upstream/master' into xsn/llama_batch_ext
ngxson 7d626f5
nits: add const
ngxson 3b86110
no more pos_max
ngxson 05a8c11
Merge branch 'master' into xsn/llama_batch_ext
ngxson 51b471f
add comment about llama_batch_ext_set_embd_state
ngxson 95956ee
handle n_embd_out properly
ngxson e640410
rename api --> embd_token
ngxson 6d00861
llama_embd
ngxson 69258fb
stub llama_batch_ext_set_embd_state
ngxson 2b50dab
support both token + embd + state in batch
ngxson 4bcbdff
Merge branch 'master' into xsn/llama_batch_ext
ngxson de88a4a
llama_batch_ext_add_embd
ngxson 4dd6864
Merge remote-tracking branch 'upstream/master' into xsn/llama_batch_ext
ngxson e67da54
upstream some changes
ngxson d48a14d
Merge branch 'master' into xsn/llama_batch_ext
ngxson 12691fb
nits
ngxson 82a63e3
fix test-batch-alloc
ngxson 4070433
add test for compat
ngxson File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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); | ||
|
|
||
|
Comment on lines
+1052
to
+1060
Collaborator
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. note: |
||
| // 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) | ||
|
|
||
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
to be discussed:
n_rowsorn_tokens?There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Rows seems better.