Commit 0bb496dbd for llama.cpp
commit 0bb496dbd3af0add77ff82c406a915b41e839d56
Author: Xuan-Son Nguyen <son@huggingface.co>
Date: Mon Oct 5 01:35:49 2026 +0200
llama: support both embd + raw tokens in batch (#29622)
* llama: support both embd + raw tokens in batch
* add to test-llama-archs
* also check case llm_arch_supports_mixed_batch = false
* constant graph topology
* have dedicated input for mixed case
* rm set_tensor_backend
* is_embd --> type
* consolidate m-rope pos handling into one place
* nits
diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp
index 2af5445df..850737eca 100644
--- a/src/llama-arch.cpp
+++ b/src/llama-arch.cpp
@@ -1178,6 +1178,21 @@ bool llm_arch_supports_rs_rollback(const llm_arch & arch) {
}
}
+// these models pick weights, routing or input meaning per ubatch based on token vs embd input
+bool llm_arch_supports_mixed_batch(const llm_arch & arch) {
+ switch (arch) {
+ case LLM_ARCH_COGVLM:
+ case LLM_ARCH_DEEPSEEK4:
+ case LLM_ARCH_GRANITE_SWITCH:
+ case LLM_ARCH_EAGLE3:
+ case LLM_ARCH_DFLASH:
+ case LLM_ARCH_GEMMA4_ASSISTANT:
+ return false;
+ default:
+ return true;
+ }
+}
+
bool llm_arch_supports_sm_tensor(const llm_arch & arch) {
switch (arch) {
case LLM_ARCH_GROK:
diff --git a/src/llama-arch.h b/src/llama-arch.h
index 148d293ce..24068fb7a 100644
--- a/src/llama-arch.h
+++ b/src/llama-arch.h
@@ -825,3 +825,4 @@ bool llm_arch_is_hybrid (const llm_arch & arch);
bool llm_arch_is_diffusion (const llm_arch & arch);
bool llm_arch_supports_sm_tensor(const llm_arch & arch);
bool llm_arch_supports_rs_rollback(const llm_arch & arch);
+bool llm_arch_supports_mixed_batch(const llm_arch & arch);
diff --git a/src/llama-batch.cpp b/src/llama-batch.cpp
index 1b3d70627..ecd48dd80 100644
--- a/src/llama-batch.cpp
+++ b/src/llama-batch.cpp
@@ -12,7 +12,7 @@
#include <algorithm>
#include <sstream>
-llama_batch_allocr::llama_batch_allocr(uint32_t n_pos_per_embd) : n_pos_per_embd(n_pos_per_embd) {
+llama_batch_allocr::llama_batch_allocr(uint32_t n_pos_per_embd, bool allow_mixed) : n_pos_per_embd(n_pos_per_embd), allow_mixed(allow_mixed) {
const char * LLAMA_BATCH_DEBUG = getenv("LLAMA_BATCH_DEBUG");
debug = LLAMA_BATCH_DEBUG ? atoi(LLAMA_BATCH_DEBUG) : 0;
@@ -49,25 +49,49 @@ bool llama_batch_allocr::init(
//
// determine the content types of the batch
// an entry can carry a token id, a token embedding, or both (e.g. MTP hook batches)
- // all entries must carry the same combination
+ // all entries must carry the same combination, or be a mix of token and embd entries
//
- const bool has_token = batch_inp.tokens[0].id != LLAMA_TOKEN_NULL;
- const bool has_embd = batch_inp.tokens[0].has_embd;
+ int32_t n_tok_only = 0;
+ int32_t n_embd_only = 0;
+ int32_t n_both = 0;
- for (int32_t i = 1; i < n_tok; ++i) {
- if ((batch_inp.tokens[i].id != LLAMA_TOKEN_NULL) != has_token ||
- batch_inp.tokens[i].has_embd != has_embd) {
- LLAMA_LOG_ERROR("%s: all entries in the batch must have the same content types\n", __func__);
+ for (int32_t i = 0; i < n_tok; ++i) {
+ const bool is_tok = batch_inp.tokens[i].id != LLAMA_TOKEN_NULL;
+ const bool is_emb = batch_inp.tokens[i].has_embd;
+
+ if (!is_tok && !is_emb) {
+ LLAMA_LOG_ERROR("%s: entry %d has neither a token id nor an embedding\n", __func__, i);
return false;
}
+
+ n_tok_only += is_tok && !is_emb;
+ n_embd_only += is_emb && !is_tok;
+ n_both += is_tok && is_emb;
}
- if (!has_token && !has_embd) {
- LLAMA_LOG_ERROR("%s: batch has neither token ids nor embeddings\n", __func__);
+ if (n_both > 0 && n_both != n_tok) {
+ LLAMA_LOG_ERROR("%s: entries with both a token id and an embedding cannot be mixed with other entries\n", __func__);
return false;
}
+ const bool mixed = n_tok_only > 0 && n_embd_only > 0;
+
+ if (mixed && !allow_mixed) {
+ LLAMA_LOG_ERROR("%s: this model or context does not support batches mixing token and embedding entries\n", __func__);
+ return false;
+ }
+
+ const bool has_token = n_tok_only > 0 || n_both > 0;
+ const bool has_embd = n_embd_only > 0 || n_both > 0;
+
+ if (mixed) {
+ is_embd_vec.resize(n_tok);
+ for (int32_t i = 0; i < n_tok; ++i) {
+ is_embd_vec[i] = batch_inp.tokens[i].has_embd ? 1 : 0;
+ }
+ }
+
//
// build flat token/embd array
//
@@ -75,6 +99,10 @@ bool llama_batch_allocr::init(
if (has_token) {
token_vec.resize(n_tok);
for (int32_t i = 0; i < n_tok; ++i) {
+ if (mixed && is_embd_vec[i]) {
+ token_vec[i] = 0; // placeholder
+ continue;
+ }
const llama_token id = batch_inp.tokens[i].id;
if (id < 0 || id >= batch_inp.n_vocab) {
LLAMA_LOG_ERROR("%s: invalid token[%d] = %d\n", __func__, i, id);
@@ -84,29 +112,35 @@ bool llama_batch_allocr::init(
}
}
- if (has_embd) {
+ if (mixed) {
+ embd_vec.assign((size_t) n_tok*n_embd, 0.0f);
+ for (int32_t i = 0; i < n_tok; ++i) {
+ if (is_embd_vec[i]) {
+ const float * src = batch_inp.embd.data() + batch_inp.tokens[i].embd_off;
+ std::copy(src, src + n_embd, embd_vec.data() + (size_t) i*n_embd);
+ }
+ }
+ } else if (has_embd) {
embd_vec = batch_inp.embd;
}
//
- // build flat pos array
- // token batch: pos[i] = tokens[i].pos[0]
- // embedding batch: pos[j*n_tok + i] = tokens[i].pos[j] (section-major)
+ // build flat pos array, section-major: pos[j*n_tok + i] = section j of entry i
+ // token entry: [p, p, p, 0] (M-RoPE text position)
+ // embd entry: tokens[i].pos as-is
//
- {
- const int32_t n_pos_total = has_token ? n_tok : n_tok * (int32_t) n_pos_per_embd;
- pos.resize(n_pos_total);
- if (has_token) {
- for (int32_t i = 0; i < n_tok; ++i) {
- pos[i] = batch_inp.tokens[i].pos[0];
- }
- } else {
- for (int32_t i = 0; i < n_tok; ++i) {
- for (uint32_t j = 0; j < n_pos_per_embd; ++j) {
- pos[(int32_t) j * n_tok + i] = batch_inp.tokens[i].pos[j];
- }
+ pos.resize((size_t) n_tok*n_pos_per_embd);
+ for (int32_t i = 0; i < n_tok; ++i) {
+ const auto & tok = batch_inp.tokens[i];
+ const bool expand = tok.id != LLAMA_TOKEN_NULL;
+ for (uint32_t j = 0; j < n_pos_per_embd; ++j) {
+ llama_pos p = tok.pos[j];
+ if (expand) {
+ // expand [p] to [p, p, p, 0] for M-RoPE
+ p = j < 3 ? tok.pos[0] : 0;
}
+ pos[(size_t) j*n_tok + i] = p;
}
}
@@ -264,6 +298,7 @@ bool llama_batch_allocr::init(
/*.seq_id_unq =*/ this->seq_id_unq.data(),
/*.seq_idx =*/ this->seq_idx.data(),
/*.output =*/ batch.logits,
+ /*.type =*/ is_embd_vec.empty() ? nullptr : is_embd_vec.data(),
/*.decision_order =*/ decision_order.empty() ? nullptr : decision_order.data(),
/*.data =*/ {},
};
@@ -294,6 +329,21 @@ bool llama_batch_allocr::init(
//
if (n_pos_per_embd > 1) {
+ // in a mixed batch, the first entry of each seq picks the rule
+ std::vector<int8_t> seq_first_embd(n_seq_max, batch.token ? 0 : 1);
+ if (mixed) {
+ std::vector<bool> seen(n_seq_max, false);
+ for (int32_t i = 0; i < batch.n_tokens; ++i) {
+ for (int32_t s = 0; s < batch.n_seq_id[i]; ++s) {
+ const llama_seq_id sid = batch.seq_id[i][s];
+ if (!seen[sid]) {
+ seen[sid] = true;
+ seq_first_embd[sid] = is_embd_vec[i];
+ }
+ }
+ }
+ }
+
// M-RoPE case: allow position to "jump" forward only (non-continuous positions are allowed)
for (uint32_t s = 0; s < n_seq_max; ++s) {
if (seq_pos[s].empty()) {
@@ -302,7 +352,7 @@ bool llama_batch_allocr::init(
const llama_pos p0 = mem ? mem->seq_pos_max(s) : -1;
- if (batch.token) {
+ if (!seq_first_embd[s]) {
if (p0 >= 0 && p0 >= seq_pos_min(s)) {
LLAMA_LOG_ERROR(
"%s: the tokens of sequence %d in the input batch have inconsistent sequence positions:\n"
@@ -471,6 +521,7 @@ llama_ubatch llama_batch_allocr::ubatch_reserve(uint32_t n_seq_tokens, uint32_t
/*.seq_id_unq =*/ udata->seq_id_unq.data(),
/*.seq_idx =*/ udata->seq_idx.data(),
/*.output =*/ udata->output.data(),
+ /*.type =*/ nullptr,
/*.decision_order =*/ nullptr,
/*.data =*/ std::move(udata),
};
@@ -769,6 +820,7 @@ void llama_batch_allocr::clear() {
token_vec .clear();
embd_vec .clear();
+ is_embd_vec .clear();
seq_id_data .clear();
pos .clear();
n_seq_id .clear();
@@ -799,7 +851,20 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u
auto udata = std::make_shared<llama_ubatch::data_t>();
- const int64_t n_embd_all = batch.embd ? (int64_t) n_tokens*n_embd : 0;
+ const bool mixed_batch = !is_embd_vec.empty();
+
+ // a ubatch with a single kind of rows is emitted as a plain token or embd ubatch
+ uint32_t n_embd_rows = 0;
+ if (mixed_batch) {
+ for (int32_t idx : idxs) {
+ n_embd_rows += is_embd_vec[idx];
+ }
+ }
+ const bool mixed = mixed_batch && n_embd_rows > 0 && n_embd_rows < n_tokens;
+ const bool use_token = batch.token && !(mixed_batch && n_embd_rows == n_tokens);
+ const bool use_embd = batch.embd && !(mixed_batch && n_embd_rows == 0);
+
+ const int64_t n_embd_all = use_embd ? (int64_t) n_tokens*n_embd : 0;
const int64_t n_pos_all = (int64_t) n_tokens*n_pos_per_embd;
udata->token .resize(n_tokens);
@@ -810,6 +875,7 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u
udata->seq_id_unq.resize(0);
udata->seq_idx .resize(LLAMA_MAX_SEQ, -1);
udata->output .resize(n_tokens);
+ udata->type .resize(mixed ? n_tokens : 0);
udata->decision_order.resize(decision_order.empty() ? 0 : n_tokens);
udata->batch_idxs = idxs;
@@ -818,21 +884,20 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u
seq_set_t seq_set_unq;
for (size_t i = 0; i < idxs.size(); ++i) {
- if (batch.token) {
+ if (use_token) {
udata->token[i] = batch.token[idxs[i]];
}
- if (batch.embd) {
+ if (use_embd) {
memcpy(udata->embd.data() + i*n_embd, batch.embd + (int64_t) idxs[i]*n_embd, n_embd*sizeof(float));
}
+ if (mixed) {
+ udata->type[i] = is_embd_vec[idxs[i]];
+ }
+
for (size_t j = 0; j < (size_t)n_pos_per_embd; ++j) {
- // if we are using M-RoPE
- // if the current batch is text, we need to broadcast the same position across all RoPE sections
- // otherwise, the input batch is image embeddings, we copy the positions as-is
- // if we are not using M-RoPE, there is only one position per token (this loop runs only once)
- size_t src_off = batch.token ? 0 : j*batch.n_tokens;
- udata->pos[j*n_tokens + i] = batch.pos[src_off + idxs[i]];
+ udata->pos[j*n_tokens + i] = batch.pos[j*batch.n_tokens + idxs[i]];
}
udata->n_seq_id[i] = batch.n_seq_id[idxs[i]];
@@ -875,14 +940,15 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector<int32_t> & idxs, u
/*.n_seqs_unq =*/ (uint32_t) udata->seq_id_unq.size(),
/*.n_pos =*/ n_pos_per_embd,
- /*.token =*/ batch.token ? udata->token.data() : nullptr,
- /*.embd =*/ batch.embd ? udata->embd.data() : nullptr,
+ /*.token =*/ use_token ? udata->token.data() : nullptr,
+ /*.embd =*/ use_embd ? udata->embd.data() : nullptr,
/*.pos =*/ udata->pos.data(),
/*.n_seq_id =*/ udata->n_seq_id.data(),
/*.seq_id =*/ udata->seq_id.data(),
/*.seq_id_unq =*/ udata->seq_id_unq.data(),
/*.seq_idx =*/ udata->seq_idx.data(),
/*.output =*/ udata->output.data(),
+ /*.type =*/ mixed ? udata->type.data() : nullptr,
/*.decision_order =*/ udata->decision_order.empty() ? nullptr : udata->decision_order.data(),
/*.data =*/ std::move(udata),
};
@@ -933,6 +999,7 @@ void llama_batch_allocr::ubatch_print(const llama_ubatch & ubatch, int debug) {
LLAMA_LOG_DEBUG("%s: seq_id_unq = %s\n", __func__, ss_seq_id_unq.str().c_str());
LLAMA_LOG_DEBUG("%s: seq_idx = %s\n", __func__, ss_seq_idx.str().c_str());
LLAMA_LOG_DEBUG("%s: output = %p\n", __func__, (void *) ubatch.output);
+ LLAMA_LOG_DEBUG("%s: type = %p\n", __func__, (void *) ubatch.type);
LLAMA_LOG_DEBUG("%s: n_outputs = %d\n", __func__, n_outputs);
if (debug > 0) {
@@ -963,7 +1030,7 @@ void llama_batch_allocr::ubatch_print(const llama_ubatch & ubatch, int debug) {
}
}
- if (ubatch.token) {
+ if (ubatch.token && !(ubatch.is_mixed() && ubatch.type[i])) {
LLAMA_LOG_DEBUG("%s: %4d: id = %6d (%16s), pos = %4d, n_seq_id = %2d, seq_id = [%s], output = %d\n",
__func__, i, ubatch.token[i], vocab->token_to_piece(ubatch.token[i]).c_str(),
ubatch.pos[i], ubatch.n_seq_id[i], ss.str().c_str(), ubatch.output[i]);
diff --git a/src/llama-batch.h b/src/llama-batch.h
index 32f103bfc..ff62e26f7 100644
--- a/src/llama-batch.h
+++ b/src/llama-batch.h
@@ -29,6 +29,11 @@ struct llama_ubatch {
return n_pos >= 3;
}
+ // mixed: type picks token or embd per row, pos has n_pos sections for all rows
+ bool is_mixed() const {
+ return type != nullptr;
+ }
+
uint32_t b_equal_seqs; // note: this is a boolean, but we use an int32_t for alignment
// otherwise address sanitizer complains
// TODO: whole_seqs for embeddings?
@@ -52,6 +57,7 @@ struct llama_ubatch {
llama_seq_id * seq_id_unq; // [n_seqs_unq] | s | seq_id
int32_t * seq_idx; // [LLAMA_MAX_SEQ] | - | seq_idx
int8_t * output; // [n_tokens] | i | -
+ int8_t * type; // [n_tokens] | i | - (mixed ubatch only, 0 - token, 1 - embd)
int32_t * decision_order; // [n_tokens], NULL if no entry has one, see llama_batch_ext_set_decision_order()
struct data_t {
@@ -63,6 +69,7 @@ struct llama_ubatch {
std::vector<llama_seq_id> seq_id_unq;
std::vector<int32_t> seq_idx;
std::vector<int8_t> output;
+ std::vector<int8_t> type;
std::vector<int32_t> batch_idxs; // original batch index for each token
std::vector<int32_t> decision_order;
@@ -73,6 +80,9 @@ struct llama_ubatch {
std::shared_ptr<data_t> data;
};
+// crash if a mixed ubatch reaches code that expects only tokens or only embd
+#define ASSERT_EMBD_OR_TOKEN(ubatch) GGML_ASSERT(!(ubatch).is_mixed() && "mixed token/embd ubatch is not supported here")
+
struct llama_hparams;
// MTP hook batches carry the target model's hidden state (n_embd_out size).
@@ -134,7 +144,7 @@ struct llama_batch_ext {
// a helper for sanitizing, fulfilling and splitting a batch
class llama_batch_allocr {
public:
- llama_batch_allocr(uint32_t n_pos_per_embd);
+ llama_batch_allocr(uint32_t n_pos_per_embd, bool allow_mixed = false);
// convert a llama_batch_ext to internal llama_batch and sanitize it
bool init(
@@ -192,12 +202,15 @@ private:
// ref: https://github.com/ggml-org/llama.cpp/issues/13694#issuecomment-2983871762
const uint32_t n_pos_per_embd;
+ const bool allow_mixed;
+
uint32_t n_embd;
uint32_t n_seq_max;
uint32_t n_outputs;
std::vector<llama_token> token_vec; // owned token IDs built from llama_batch_ext
std::vector<float> embd_vec; // owned embeddings built from llama_batch_ext
+ std::vector<int8_t> is_embd_vec; // mixed batch only (= 1 if embd, 0 if text token)
std::vector<llama_seq_id> seq_id_data; // flat storage for seq_id pointers below
std::vector<llama_pos> pos;
diff --git a/src/llama-context.cpp b/src/llama-context.cpp
index 96b5464e6..07b4c6148 100644
--- a/src/llama-context.cpp
+++ b/src/llama-context.cpp
@@ -87,7 +87,9 @@ llama_context::llama_context(
model(model),
cvec(std::make_unique<llama_adapter_cvec>()),
loras(std::make_unique<llama_adapter_loras>()),
- balloc(std::make_unique<llama_batch_allocr>(model.hparams.n_pos_per_embd())) {
+ // MTP uses the embd input for the hidden state
+ balloc(std::make_unique<llama_batch_allocr>(model.hparams.n_pos_per_embd(),
+ llm_arch_supports_mixed_batch(model.arch) && params.ctx_type == LLAMA_CONTEXT_TYPE_DEFAULT)) {
// TODO warning when creating llama_context with awkward ctx size that is not a power of 2,
// may need to be backend-dependent
LLAMA_LOG_INFO("%s: constructing llama_context\n", __func__);
diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp
index 16a3ed2ab..1112ad885 100644
--- a/src/llama-graph.cpp
+++ b/src/llama-graph.cpp
@@ -73,13 +73,55 @@ void llm_graph_input_embd::set_input(const llama_ubatch * ubatch) {
ggml_backend_tensor_set(tokens, ubatch->token, 0, n_tokens*ggml_element_size(tokens));
}
- if (ubatch->embd) {
+ if (ubatch->embd && embd && !ubatch->is_mixed()) {
GGML_ASSERT(n_embd == embd->ne[0]);
const int64_t n_tokens = ubatch->n_tokens;
ggml_backend_tensor_set(embd, ubatch->embd, 0, n_tokens*n_embd*ggml_element_size(embd));
}
+
+ if (ubatch->is_mixed() && embd) {
+ GGML_ASSERT(mixed_tokens && mixed_slots && mixed_embd && "mixed token/embd ubatch is not supported here");
+
+ std::vector<int32_t> ids;
+ std::vector<int64_t> slots;
+ for (uint32_t i = 0; i < ubatch->n_tokens; ++i) {
+ if (!ubatch->type[i]) {
+ ids.push_back(ubatch->token[i]);
+ slots.push_back(i);
+ }
+ }
+ GGML_ASSERT((int64_t) ids.size() == mixed_tokens->ne[0]);
+ GGML_ASSERT(n_embd == mixed_embd->ne[0]);
+
+ ggml_backend_tensor_set(mixed_tokens, ids.data(), 0, ggml_nbytes(mixed_tokens));
+ ggml_backend_tensor_set(mixed_slots, slots.data(), 0, ggml_nbytes(mixed_slots));
+ ggml_backend_tensor_set(mixed_embd, ubatch->embd, 0, ggml_nbytes(mixed_embd));
+ }
+
+ if (scale_rows) {
+ const int64_t n_tokens = ubatch->n_tokens;
+
+ std::vector<float> data(n_tokens);
+ for (int64_t i = 0; i < n_tokens; ++i) {
+ const bool is_embd = !ubatch->token || (ubatch->is_mixed() && ubatch->type[i]);
+ data[i] = is_embd ? 1.0f : scale_tok;
+ }
+ ggml_backend_tensor_set(scale_rows, data.data(), 0, ggml_nbytes(scale_rows));
+ }
+}
+
+// number of token rows of the mixed path, a non-mixed ubatch is sized for the worst case
+static int64_t llm_graph_n_tok_rows(const llama_ubatch & ubatch) {
+ if (!ubatch.is_mixed()) {
+ return ubatch.n_tokens;
+ }
+ int64_t n = 0;
+ for (uint32_t i = 0; i < ubatch.n_tokens; ++i) {
+ n += !ubatch.type[i];
+ }
+ return n;
}
bool llm_graph_input_embd::can_reuse(const llm_graph_params & params) {
@@ -87,11 +129,16 @@ bool llm_graph_input_embd::can_reuse(const llm_graph_params & params) {
res &= (!params.ubatch.token) || (tokens && tokens->ne[0] == params.ubatch.n_tokens);
res &= (!params.ubatch.embd) || (embd && embd->ne[1] == params.ubatch.n_tokens);
+ res &= (!mixed_tokens) || mixed_tokens->ne[0] == llm_graph_n_tok_rows(params.ubatch);
+ res &= (!mixed_embd) || mixed_embd->ne[1] == params.ubatch.n_tokens;
+ res &= (!scale_rows) || scale_rows->ne[1] == params.ubatch.n_tokens;
return res;
}
void llm_graph_input_embd_h::set_input(const llama_ubatch * ubatch) {
+ ASSERT_EMBD_OR_TOKEN(*ubatch);
+
const int64_t n_tokens = ubatch->n_tokens;
if (ubatch->token) {
@@ -128,21 +175,7 @@ void llm_graph_input_pos::set_input(const llama_ubatch * ubatch) {
if (ubatch->pos && pos) {
const int64_t n_tokens = ubatch->n_tokens;
- if (ubatch->token && n_pos_per_embd == 4) {
- // in case we're using M-RoPE with text tokens, convert the 1D positions to 4D
- // the 3 first dims are the same, and 4th dim is all 0
- std::vector<llama_pos> pos_data(n_tokens*n_pos_per_embd);
- // copy the first dimension
- for (int i = 0; i < n_tokens; ++i) {
- pos_data[ i] = ubatch->pos[i];
- pos_data[ n_tokens + i] = ubatch->pos[i];
- pos_data[2 * n_tokens + i] = ubatch->pos[i];
- pos_data[3 * n_tokens + i] = 0; // 4th dim is 0
- }
- ggml_backend_tensor_set(pos, pos_data.data(), 0, pos_data.size()*ggml_element_size(pos));
- } else {
- ggml_backend_tensor_set(pos, ubatch->pos, 0, n_tokens*n_pos_per_embd*ggml_element_size(pos));
- }
+ ggml_backend_tensor_set(pos, ubatch->pos, 0, n_tokens*n_pos_per_embd*ggml_element_size(pos));
}
}
@@ -2377,7 +2410,7 @@ ggml_tensor * llm_graph_context::build_moe_ffn(
}
// input embeddings with optional lora
-ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd) const {
+ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd, float tok_scale) const {
const int64_t n_embd_inp = hparams.n_embd_inp();
const int64_t n_embd = hparams.n_embd;
@@ -2394,15 +2427,9 @@ ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd) const {
cb(inp->embd, "inp_embd", -1);
ggml_set_input(inp->embd);
- // select one of the 2 inputs, based on the batch contents
- // ref: https://github.com/ggml-org/llama.cpp/pull/18550
- std::array<ggml_tensor *, 2> inps;
-
- // token embeddings path (ubatch.token != nullptr)
- {
- auto & cur = inps[0];
-
- cur = ggml_get_rows(ctx0, tok_embd, inp->tokens);
+ // token embeddings with lora and padding
+ auto build_tok = [&](ggml_tensor * ids) {
+ ggml_tensor * cur = ggml_get_rows(ctx0, tok_embd, ids);
// apply lora for embedding tokens if needed
for (const auto & lora : *loras) {
@@ -2416,7 +2443,7 @@ ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd) const {
ggml_tensor * inpL_delta = ggml_scale(ctx0, ggml_mul_mat(
ctx0, lw->b, // non-transposed lora_b
- ggml_get_rows(ctx0, lw->a, inp->tokens)
+ ggml_get_rows(ctx0, lw->a, ids)
), scale);
cur = ggml_add(ctx0, cur, inpL_delta);
@@ -2425,19 +2452,48 @@ ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd) const {
if (n_embd_inp != n_embd) {
cur = ggml_pad(ctx0, cur, hparams.n_embd_inp() - n_embd, 0, 0, 0);
}
- }
+
+ return cur;
+ };
+
+ // select one of the 3 inputs, based on the batch contents
+ // ref: https://github.com/ggml-org/llama.cpp/pull/18550
+ std::array<ggml_tensor *, 3> inps = {};
+
+ // token embeddings path (ubatch.token != nullptr)
+ inps[0] = build_tok(inp->tokens);
// vector embeddings path (ubatch.embd != nullptr)
- {
- auto & cur = inps[1];
+ inps[1] = inp->embd;
+
+ // mixed path (ubatch.is_mixed()): set_rows the token rows into a copy of the embd rows, with its own inputs as select branches must not share tensors
+ // TODO: use inp->tokens and inp->embd once ggml_build_forward_select allows it
+ const bool has_mixed = llm_arch_supports_mixed_batch(arch) && cparams.ctx_type == LLAMA_CONTEXT_TYPE_DEFAULT;
+ if (has_mixed) {
+ const int64_t n_tok_rows = llm_graph_n_tok_rows(ubatch);
+
+ inp->mixed_tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tok_rows);
+ cb(inp->mixed_tokens, "inp_mixed_tokens", -1);
+ ggml_set_input(inp->mixed_tokens);
- cur = inp->embd;
+ inp->mixed_slots = ggml_new_tensor_1d(ctx0, GGML_TYPE_I64, n_tok_rows);
+ cb(inp->mixed_slots, "inp_mixed_slots", -1);
+ ggml_set_input(inp->mixed_slots);
+
+ inp->mixed_embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_embd_inp, ubatch.n_tokens);
+ cb(inp->mixed_embd, "inp_mixed_embd", -1);
+ ggml_set_input(inp->mixed_embd);
+
+ // note: set_rows writes into its destination, so it gets a copy of the input
+ inps[2] = ggml_set_rows(ctx0, ggml_dup(ctx0, inp->mixed_embd), build_tok(inp->mixed_tokens), inp->mixed_slots);
}
assert(ggml_are_same_shape (inps[0], inps[1]));
assert(ggml_are_same_stride(inps[0], inps[1]));
- ggml_tensor * cur = ggml_build_forward_select(gf, inps.data(), inps.size(), ubatch.token ? 0 : 1);
+ const int idx = ubatch.is_mixed() ? 2 : ubatch.token ? 0 : 1;
+
+ ggml_tensor * cur = ggml_build_forward_select(gf, inps.data(), has_mixed ? 3 : 2, idx);
if (n_embd_inp != n_embd) {
cur = ggml_view_2d(ctx0, cur, n_embd, n_tokens, cur->nb[1], 0);
@@ -2445,16 +2501,31 @@ ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd) const {
res->t_inp_embd = cur;
- // For Granite architecture
// NOTE: For deepstack models, only apply scale to token inputs (ie text-only input).
// Raw embeddings are assumed to be multimodal inputs that should not be scaled.
- if (hparams.f_embedding_scale != 0.0f && (ubatch.token || hparams.n_deepstack_layers == 0)) {
+ const bool scale_tok_only = hparams.f_embedding_scale != 0.0f && hparams.n_deepstack_layers > 0;
+
+ // For Granite architecture
+ if (hparams.f_embedding_scale != 0.0f && !scale_tok_only) {
if (!ggml_is_contiguous(cur)) {
cur = ggml_cont(ctx0, cur);
}
cur = ggml_scale(ctx0, cur, hparams.f_embedding_scale);
}
+ // scale the token rows only, applied after the select so that the graph is the same for any batch contents
+ inp->scale_tok = tok_scale*(scale_tok_only ? hparams.f_embedding_scale : 1.0f);
+ if (inp->scale_tok != 1.0f) {
+ inp->scale_rows = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, 1, ubatch.n_tokens);
+ cb(inp->scale_rows, "inp_scale_rows", -1);
+ ggml_set_input(inp->scale_rows);
+
+ if (!ggml_is_contiguous(cur)) {
+ cur = ggml_cont(ctx0, cur);
+ }
+ cur = ggml_mul(ctx0, cur, inp->scale_rows);
+ }
+
cb(cur, "embd", -1);
res->add_input(std::move(inp));
diff --git a/src/llama-graph.h b/src/llama-graph.h
index 366eccbcd..5cf74c997 100644
--- a/src/llama-graph.h
+++ b/src/llama-graph.h
@@ -134,8 +134,14 @@ public:
bool can_reuse(const llm_graph_params & params) override;
- ggml_tensor * tokens = nullptr; // I32 [n_batch]
- ggml_tensor * embd = nullptr; // F32 [n_embd, n_batch]
+ ggml_tensor * tokens = nullptr; // I32 [n_batch]
+ ggml_tensor * embd = nullptr; // F32 [n_embd, n_batch]
+ ggml_tensor * mixed_tokens = nullptr; // I32 [n_tok_rows], mixed path: ids of the token rows
+ ggml_tensor * mixed_slots = nullptr; // I64 [n_tok_rows], mixed path: batch index of the token rows
+ ggml_tensor * mixed_embd = nullptr; // F32 [n_embd, n_batch], mixed path: embd rows, token rows are overwritten
+ ggml_tensor * scale_rows = nullptr; // F32 [1, n_batch], per-row scale: scale_tok for token rows, 1 for embd rows
+
+ float scale_tok = 1.0f;
const int64_t n_embd = 0;
};
@@ -823,6 +829,7 @@ struct llm_graph_params {
ubatch.n_seq_tokens == other.ubatch.n_seq_tokens &&
ubatch.n_seqs == other.ubatch.n_seqs &&
ubatch.n_seqs_unq == other.ubatch.n_seqs_unq &&
+ ubatch.is_mixed() == other.ubatch.is_mixed() &&
(
(!ubatch.token && !other.ubatch.token) ||
(!ubatch.embd && !other.ubatch.embd) ||
@@ -1166,7 +1173,8 @@ struct llm_graph_context {
// inputs
//
- ggml_tensor * build_inp_embd(ggml_tensor * tok_embd) const;
+ // tok_scale: applied to token rows only
+ ggml_tensor * build_inp_embd(ggml_tensor * tok_embd, float tok_scale = 1.0f) const;
ggml_tensor * build_inp_pos() const;
ggml_tensor * build_inp_attn_scale() const;
ggml_tensor * build_inp_out_ids() const;
diff --git a/src/llama-kv-cache-dsv4.cpp b/src/llama-kv-cache-dsv4.cpp
index 4dbcfb8db..5dd0b4591 100644
--- a/src/llama-kv-cache-dsv4.cpp
+++ b/src/llama-kv-cache-dsv4.cpp
@@ -83,6 +83,7 @@ static llama_ubatch dsv4_build_raw_write_ubatch(const llama_ubatch & ubatch) {
if (!dsv4_ubatch_has_coupled(ubatch)) {
return ubatch;
}
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.embd) {
throw std::runtime_error("DSV4 coupled embedding ubatches are not supported");
}
@@ -164,6 +165,7 @@ static llama_ubatch dsv4_build_raw_write_ubatch(const llama_ubatch & ubatch) {
/*.seq_id_unq =*/ data->seq_id_unq.data(),
/*.seq_idx =*/ data->seq_idx.data(),
/*.output =*/ data->output.data(),
+ /*.type =*/ nullptr,
/*.decision_order =*/ nullptr,
/*.data =*/ data,
};
diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp
index 1d0554fa1..8891526cd 100644
--- a/src/llama-kv-cache.cpp
+++ b/src/llama-kv-cache.cpp
@@ -1138,7 +1138,9 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch &
ext.y = ubatch.pos[i + ubatch.n_tokens];
}
- if (ubatch.token) {
+ const bool is_embd = !ubatch.token || (ubatch.is_mixed() && ubatch.type[i]);
+
+ if (!is_embd) {
ext.tok = ubatch.token[i];
} else if (hparams.ple_n_heads > 0) {
// embd batch (multimodal input) has no token ids, need to pad it with the correct ID for PLE layers
@@ -1862,10 +1864,13 @@ void llama_kv_cache::get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, st
// an embd (multimodal) ubatch can repeat one position for a whole image, so positions
// do not encode the token order; resolve its predecessors by ubatch order instead
+ // same for a mixed ubatch
+ const bool by_order = !ubatch.token || ubatch.is_mixed();
+
std::vector<uint32_t> ord; // index among the ubatch tokens of the same seq
std::unordered_map<llama_seq_id, std::vector<uint32_t>> seq_idx;
- if (!ubatch.token) {
+ if (by_order) {
ord.resize(n_tokens);
for (uint32_t i = 0; i < n_tokens; ++i) {
auto & v = seq_idx[ubatch.seq_id[i][0]];
@@ -1883,7 +1888,7 @@ void llama_kv_cache::get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, st
const llama_pos d = (llama_pos) (n - j);
llama_pos p;
- if (!ubatch.token) {
+ if (by_order) {
const auto & v = seq_idx[seq_id];
const int64_t k = (int64_t) ord[i] - d;
// k >= 0: an earlier token of this very ubatch; k < 0: before the chunk
diff --git a/src/models/cogvlm.cpp b/src/models/cogvlm.cpp
index 750f57a39..bc394b3c1 100644
--- a/src/models/cogvlm.cpp
+++ b/src/models/cogvlm.cpp
@@ -70,6 +70,7 @@ llama_model_cogvlm::graph::graph(const llama_model & model, const llm_graph_para
// check ubatch to see if we have input tokens (text)
// or an input embedding vector (image)
bool is_text;
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.token) {
is_text = true;
} else {
diff --git a/src/models/cohere2moe.cpp b/src/models/cohere2moe.cpp
index 7704cbb87..cf2af012d 100644
--- a/src/models/cohere2moe.cpp
+++ b/src/models/cohere2moe.cpp
@@ -317,6 +317,7 @@ llama_model_cohere2moe::graph_mtp::graph_mtp(const llama_model & model, const ll
// TODO: make static using `ggml_build_forward_select()`
// see llm_graph_context::build_inp_embd() for reference
ggml_tensor * tok_embd;
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.token) {
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens);
diff --git a/src/models/deepseek2.cpp b/src/models/deepseek2.cpp
index deca86527..6d217ea0c 100644
--- a/src/models/deepseek2.cpp
+++ b/src/models/deepseek2.cpp
@@ -221,6 +221,7 @@ llama_model_deepseek2::graph_mtp::graph_mtp(const llama_model & model, const llm
ggml_set_input(inp->embd);
ggml_tensor * tok_embd;
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.token) {
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens
? layer.nextn.embed_tokens
diff --git a/src/models/deepseek32.cpp b/src/models/deepseek32.cpp
index 60cc17c49..849c7a9a0 100644
--- a/src/models/deepseek32.cpp
+++ b/src/models/deepseek32.cpp
@@ -535,6 +535,7 @@ llama_model_deepseek32::graph_mtp::graph_mtp(const llama_model & model, const ll
ggml_set_input(inp->embd);
ggml_tensor * tok_embd;
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.token) {
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
diff --git a/src/models/deepseek4.cpp b/src/models/deepseek4.cpp
index a14388725..223336d49 100644
--- a/src/models/deepseek4.cpp
+++ b/src/models/deepseek4.cpp
@@ -1281,6 +1281,7 @@ llama_model_deepseek4::graph::graph(const llama_model & model, const llm_graph_p
ggml_tensor * exp_probs_b = layer.ffn_exp_probs_b;
// may apply exp_probs_b_vl is input is from mtmd
+ ASSERT_EMBD_OR_TOKEN(ubatch);
const bool is_media = ubatch.embd != nullptr;
if (is_media) {
if (layer.ffn_exp_probs_b_vl) {
@@ -1366,6 +1367,7 @@ llama_model_deepseek4::graph_mtp::graph_mtp(const llama_model & model, const llm
GGML_ASSERT(cparams.nextn_layer_offset >= 0 &&
cparams.nextn_layer_offset < (int) hparams.n_layer_nextn &&
"nextn_layer_offset out of range [0, n_layer_nextn)");
+ ASSERT_EMBD_OR_TOKEN(ubatch);
GGML_ASSERT(ubatch.token && "DEEPSEEK4 MTP requires token input");
const int64_t hc = hparams.dsv4_hc_mult;
diff --git a/src/models/dflash.cpp b/src/models/dflash.cpp
index c448e63f3..c8c2895b6 100644
--- a/src/models/dflash.cpp
+++ b/src/models/dflash.cpp
@@ -604,6 +604,7 @@ llama_model_dflash::graph<false>::graph(const llama_model & model, const llm_gra
};
// KV cache injection
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.embd) {
auto inp = std::make_unique<llm_graph_input_embd>(n_embd_inp);
@@ -870,6 +871,7 @@ llama_model_dflash::graph_dsv4::graph_dsv4(const llama_model & model, const llm_
llm_graph_input_attn_k_iswa * inp_attn = build_attn_inp_k_iswa();
// KV cache injection: fused target features from the encoder
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.embd) {
auto inp = std::make_unique<llm_graph_input_embd>(n_embd_inp);
diff --git a/src/models/gemma-embedding.cpp b/src/models/gemma-embedding.cpp
index 6c97883d8..692e19ee9 100644
--- a/src/models/gemma-embedding.cpp
+++ b/src/models/gemma-embedding.cpp
@@ -77,10 +77,8 @@ llama_model_gemma_embedding::graph::graph(const llama_model & model, const llm_g
ggml_tensor * cur;
ggml_tensor * inpL;
- inpL = build_inp_embd(model.tok_embd);
-
// important: do not normalize weights for raw embeddings input (i.e. encoded image embeddings)
- inpL = ggml_scale(ctx0, inpL, ubatch.token ? sqrtf(n_embd) : 1.0f);
+ inpL = build_inp_embd(model.tok_embd, sqrtf(n_embd));
cb(inpL, "inp_scaled", -1);
// inp_pos - contains the positions
diff --git a/src/models/gemma3.cpp b/src/models/gemma3.cpp
index f99bbaacd..83cb57f99 100644
--- a/src/models/gemma3.cpp
+++ b/src/models/gemma3.cpp
@@ -85,10 +85,8 @@ llama_model_gemma3::graph<iswa>::graph(const llama_model & model, const llm_grap
ggml_tensor * cur;
ggml_tensor * inpL;
- inpL = build_inp_embd(model.tok_embd);
-
// important: do not normalize weights for raw embeddings input (i.e. encoded image embeddings)
- inpL = ggml_scale(ctx0, inpL, ubatch.token ? sqrtf(n_embd) : 1.0f);
+ inpL = build_inp_embd(model.tok_embd, sqrtf(n_embd));
cb(inpL, "inp_scaled", -1);
// inp_pos - contains the positions
diff --git a/src/models/gemma3n.cpp b/src/models/gemma3n.cpp
index 4d47ddc62..50d30c63f 100644
--- a/src/models/gemma3n.cpp
+++ b/src/models/gemma3n.cpp
@@ -96,10 +96,8 @@ llama_model_gemma3n::graph::graph(const llama_model & model, const llm_graph_par
ggml_tensor * cur;
ggml_tensor * inpL;
- inpL = build_inp_embd(model.tok_embd);
-
// important: do not normalize weights for raw embeddings input (i.e. encoded image embeddings)
- inpL = ggml_scale(ctx0, inpL, ubatch.token ? sqrtf(n_embd) : 1.0f);
+ inpL = build_inp_embd(model.tok_embd, sqrtf(n_embd));
cb(inpL, "inp_scaled", -1);
// inp_pos - contains the positions
@@ -325,6 +323,7 @@ ggml_tensor * llama_model_gemma3n::graph::build_inp_per_layer() {
auto inp = std::make_unique<llm_graph_input_embd>(n_embd);
ggml_tensor * inp_per_layer;
float tok_embd_scale = sqrtf((float) n_embd_altup);
+ // mixed ubatch: embd rows have token id 0, same padding row as below
if (ubatch.token) {
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, ubatch.n_tokens);
ggml_set_input(inp->tokens);
diff --git a/src/models/gemma4.cpp b/src/models/gemma4.cpp
index 65fc7623d..38239eba0 100644
--- a/src/models/gemma4.cpp
+++ b/src/models/gemma4.cpp
@@ -159,10 +159,8 @@ llama_model_gemma4::graph::graph(const llama_model & model, const llm_graph_para
ggml_tensor * cur;
ggml_tensor * inpL;
- inpL = build_inp_embd(model.tok_embd);
-
// important: do not normalize weights for raw embeddings input (i.e. encoded image emdeddings)
- inpL = ggml_scale(ctx0, inpL, ubatch.token ? sqrtf(n_embd) : 1.0f);
+ inpL = build_inp_embd(model.tok_embd, sqrtf(n_embd));
cb(inpL, "inp_scaled", -1);
// inp_pos - contains the positions
@@ -473,6 +471,7 @@ ggml_tensor * llama_model_gemma4::graph::build_inp_per_layer() {
ggml_tensor * inp_per_layer;
float tok_embd_scale = sqrtf((float) n_embd_per_layer);
+ // mixed ubatch: embd rows have token id 0, same padding row as below
if (ubatch.token) {
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, ubatch.n_tokens);
ggml_set_input(inp->tokens);
diff --git a/src/models/glm-dsa.cpp b/src/models/glm-dsa.cpp
index 44d883274..6a5132cf6 100644
--- a/src/models/glm-dsa.cpp
+++ b/src/models/glm-dsa.cpp
@@ -579,6 +579,7 @@ llama_model_glm_dsa::graph_mtp::graph_mtp(const llama_model & model, const llm_g
ggml_set_input(inp->embd);
ggml_tensor * tok_embd;
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.token) {
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
diff --git a/src/models/glm4-moe.cpp b/src/models/glm4-moe.cpp
index d6ae5783c..8cdbe10ad 100644
--- a/src/models/glm4-moe.cpp
+++ b/src/models/glm4-moe.cpp
@@ -158,6 +158,7 @@ llama_model_glm4_moe::graph_mtp::graph_mtp(const llama_model & model, const llm_
ggml_set_input(inp->embd);
ggml_tensor * tok_embd;
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.token) {
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens);
diff --git a/src/models/granite-switch.cpp b/src/models/granite-switch.cpp
index 7c9a901c8..044c8e699 100644
--- a/src/models/granite-switch.cpp
+++ b/src/models/granite-switch.cpp
@@ -149,6 +149,7 @@ public:
// K dim-0 is +gain for an adapter token, -gain otherwise; the causal softmax then
// lets a single visible adapter token dominate so the readback recovers its slot.
void llm_graph_input_switch::set_input(const llama_ubatch * ubatch) {
+ ASSERT_EMBD_OR_TOKEN(*ubatch);
if (!ubatch->token) {
return;
}
@@ -226,6 +227,7 @@ llama_model_granite_switch::graph::graph(
const auto & smodel = static_cast<const llama_model_granite_switch &>(model);
// TODO: support raw embedding input (multimodal / pre-embedded tokens) when needed
+ ASSERT_EMBD_OR_TOKEN(ubatch);
GGML_ASSERT(ubatch.token && "granite-switch requires token input");
const int64_t n_embd_head = hparams.n_embd_head_v();
diff --git a/src/models/nemotron-h-moe.cpp b/src/models/nemotron-h-moe.cpp
index b4fb25430..b9b4fdcf1 100644
--- a/src/models/nemotron-h-moe.cpp
+++ b/src/models/nemotron-h-moe.cpp
@@ -34,6 +34,7 @@ llama_model_nemotron_h_moe::graph_mtp::graph_mtp(const llama_model & model, cons
ggml_set_input(inp->embd);
ggml_tensor * tok_embd;
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.token) {
tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens);
} else {
diff --git a/src/models/qwen35.cpp b/src/models/qwen35.cpp
index d50f067a5..ab98744a5 100644
--- a/src/models/qwen35.cpp
+++ b/src/models/qwen35.cpp
@@ -529,6 +529,7 @@ llama_model_qwen35::graph_mtp::graph_mtp(const llama_model & model, const llm_gr
// TODO: make static using `ggml_build_forward_select()`
// see llm_graph_context::build_inp_embd() for reference
ggml_tensor * tok_embd;
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.token) {
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
diff --git a/src/models/qwen35moe.cpp b/src/models/qwen35moe.cpp
index bdf772625..f0f917af7 100644
--- a/src/models/qwen35moe.cpp
+++ b/src/models/qwen35moe.cpp
@@ -579,6 +579,7 @@ llama_model_qwen35moe::graph_mtp::graph_mtp(const llama_model & model, const llm
// TODO: make static using `ggml_build_forward_select()`
// see llm_graph_context::build_inp_embd() for reference
ggml_tensor * tok_embd;
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.token) {
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
diff --git a/src/models/qwen3next.cpp b/src/models/qwen3next.cpp
index b63fc9c6a..340ef28f7 100644
--- a/src/models/qwen3next.cpp
+++ b/src/models/qwen3next.cpp
@@ -653,6 +653,7 @@ llama_model_qwen3next::graph_mtp::graph_mtp(const llama_model & model, const llm
// TODO: make static using `ggml_build_forward_select()`
// see llm_graph_context::build_inp_embd() for reference
ggml_tensor * tok_embd;
+ ASSERT_EMBD_OR_TOKEN(ubatch);
if (ubatch.token) {
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd;
diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp
index f4df6a5c2..416c9263b 100644
--- a/src/models/qwen4exp.cpp
+++ b/src/models/qwen4exp.cpp
@@ -1240,7 +1240,8 @@ void llm_graph_input_qwen4exp_ple::set_input(const llama_ubatch * ubatch) {
? (llama_token) hparams.ple_image_token_id
: (llama_token) hparams.ple_eos_token_id;
auto tok_of = [&](int64_t k) -> llama_token {
- return ubatch->token ? ubatch->token[k] : img_tok;
+ const bool is_embd = !ubatch->token || (ubatch->is_mixed() && ubatch->type[k]);
+ return is_embd ? img_tok : ubatch->token[k];
};
const int64_t n_tokens = ubatch->n_tokens;
diff --git a/tests/test-batch-alloc.cpp b/tests/test-batch-alloc.cpp
index ad186c693..b085917cf 100644
--- a/tests/test-batch-alloc.cpp
+++ b/tests/test-batch-alloc.cpp
@@ -99,6 +99,15 @@ struct batch_builder {
const llama_pos pos[GGML_MROPE_SECTIONS] = { p, 0, 0, 0 };
return add_embd(pos, seq_ids, output);
}
+
+ int32_t add_tok(llama_token id, llama_pos p, llama_seq_id seq_id, bool output) {
+ const int32_t idx = b.add_token(seq_id);
+ GGML_ASSERT(idx >= 0);
+ GGML_ASSERT(b.set_token_id(idx, id));
+ GGML_ASSERT(b.set_token_pos(idx, &p));
+ GGML_ASSERT(b.set_output(idx, output));
+ return idx;
+ }
};
static void test_init(testing & t) {
@@ -431,6 +440,125 @@ static void test_content_types(testing & t) {
});
}
+static void test_mixed(testing & t) {
+ llama_vocab vocab;
+
+ t.test("rejected_unless_allowed", [&](testing & t) {
+ batch_builder bb(2, nullptr, 4, 1, /*n_vocab*/ 10);
+ bb.add_tok(3, 0, 0, false);
+ bb.add(1, {0}, true);
+
+ llama_batch_allocr ba_default(1);
+ t.assert_true("rejected by default", !ba_default.init(bb.b, vocab, false));
+
+ llama_batch_allocr ba_mixed(1, true);
+ t.assert_true("accepted when allowed", ba_mixed.init(bb.b, vocab, false));
+ });
+
+ t.test("layout_and_split", [&](testing & t) {
+ batch_builder bb(2, nullptr, 4, 1, /*n_vocab*/ 10);
+ bb.add_tok(3, 0, 0, false);
+ bb.add(1, {0}, false);
+ bb.add(2, {0}, false);
+ bb.add_tok(5, 3, 0, true);
+
+ llama_batch_allocr ba(1, true);
+ t.assert_true(ba.init(bb.b, vocab, false));
+
+ const llama_batch & batch = ba.get_batch();
+ t.assert_true(batch.token != nullptr && batch.embd != nullptr);
+
+ const llama_token exp_tok[4] = { 3, 0, 0, 5 };
+ const float exp_embd[8] = { 0, 0, 100, 101, 200, 201, 0, 0 };
+ for (int i = 0; i < 4; ++i) {
+ t.assert_equal(exp_tok[i], batch.token[i]);
+ }
+ for (int i = 0; i < 8; ++i) {
+ t.assert_equal(exp_embd[i], batch.embd[i]);
+ }
+
+ llama_ubatch ub0 = ba.split_simple(3);
+ t.assert_equal(3u, ub0.n_tokens);
+ t.assert_true(ub0.is_mixed());
+ const int8_t exp_is_embd[3] = { 0, 1, 1 };
+ for (int i = 0; i < 3; ++i) {
+ t.assert_equal(exp_is_embd[i], ub0.type[i]);
+ t.assert_equal((llama_pos) i, ub0.pos[i]);
+ }
+ t.assert_equal(3, ub0.token[0]);
+ t.assert_equal(0.0f, ub0.embd[0]);
+ t.assert_equal(100.0f, ub0.embd[2]);
+
+ // token rows only: a plain token ubatch
+ llama_ubatch ub1 = ba.split_simple(3);
+ t.assert_equal(1u, ub1.n_tokens);
+ t.assert_true(!ub1.is_mixed());
+ t.assert_true(ub1.embd == nullptr);
+ t.assert_equal(5, ub1.token[0]);
+ t.assert_equal((llama_pos) 3, ub1.pos[0]);
+
+ t.assert_equal(0u, ba.split_simple(3).n_tokens);
+ });
+
+ t.test("rejects_entry_with_both", [&](testing & t) {
+ // token + embd on one entry is the MTP layout
+ batch_builder bb(2, nullptr, 4, 1, /*n_vocab*/ 10);
+ bb.add_tok(3, 0, 0, false);
+ bb.add(1, {0}, false);
+ const int32_t idx = bb.add_tok(4, 2, 0, true);
+ const auto r = bb.row(idx, bb.n_embd);
+ t.assert_true(bb.b.set_token_embd(idx, { r.data(), 1, bb.n_embd }));
+
+ llama_batch_allocr ba(1, true);
+ t.assert_true(!ba.init(bb.b, vocab, false));
+ });
+
+ t.test("mrope_pos_expanded", [&](testing & t) {
+ const uint32_t n_pos = 4;
+ batch_builder bb(2, nullptr, 4, n_pos, /*n_vocab*/ 10);
+
+ bb.add_tok(3, 10, 0, false);
+ const llama_pos pos1[n_pos] = { 11, 5, 7, 0 };
+ bb.add_embd(pos1, {0}, true);
+
+ llama_batch_allocr ba(n_pos, true);
+ t.assert_true(ba.init(bb.b, vocab, false));
+
+ llama_ubatch ub = ba.split_simple(2);
+ const llama_pos expected[8] = { 10, 11, 10, 5, 10, 7, 0, 0 };
+ for (int i = 0; i < 8; ++i) {
+ t.assert_equal(expected[i], ub.pos[i]);
+ }
+ });
+
+ t.test("mrope_rule_follows_first_entry", [&](testing & t) {
+ const uint32_t n_pos = 4;
+
+ mock_memory mem;
+ mem.ranges[0] = {0, 9};
+
+ llama_batch_allocr ba(n_pos, true);
+
+ // token first: must start after the memory
+ {
+ batch_builder bb(2, &mem, 4, n_pos, /*n_vocab*/ 10);
+ bb.add_tok(3, 9, 0, false);
+ const llama_pos pos[n_pos] = { 10, 1, 1, 0 };
+ bb.add_embd(pos, {0}, true);
+ t.assert_true("token overlapping the memory is rejected", !ba.init(bb.b, vocab, false));
+ }
+
+ // embd first: can overlap the memory
+ {
+ batch_builder bb(2, &mem, 4, n_pos, /*n_vocab*/ 10);
+ const llama_pos pos[n_pos] = { 9, 1, 1, 0 };
+ bb.add_embd(pos, {0}, false);
+ bb.add_tok(3, 10, 0, true);
+ t.assert_true("embd overlapping the memory is allowed", ba.init(bb.b, vocab, false));
+ }
+ });
+}
+
static void test_split(testing & t) {
llama_vocab vocab;
@@ -1059,6 +1187,7 @@ int main(int argc, char ** argv) {
t.test("init", test_init);
t.test("content_types", test_content_types);
+ t.test("mixed", test_mixed);
t.test("compat", test_compat);
t.test("split", test_split);
t.test("keep_tail", test_keep_tail);
diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp
index 93bf47baa..7a76efccc 100644
--- a/tests/test-llama-archs.cpp
+++ b/tests/test-llama-archs.cpp
@@ -535,6 +535,50 @@ static std::vector<float> get_logits(
return ret;
}
+// entries [n/4, n/2) are embd rows, decoded either as token/embd/token chunks or as one mixed batch
+// returns the llama_process() error code
+static int32_t get_logits_mixed(
+ llama_model * model, llama_context * lctx, const std::vector<llama_token> & tokens, const std::vector<float> & embd, bool mixed,
+ std::vector<float> & ret) {
+ const uint32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model));
+ const uint32_t n_embd = llama_model_n_embd_inp(model);
+ const uint32_t n_tokens = tokens.size();
+ const uint32_t i_embd_0 = n_tokens/4;
+ const uint32_t i_embd_1 = n_tokens/2;
+
+ const std::vector<uint32_t> bounds = mixed
+ ? std::vector<uint32_t>{0, n_tokens}
+ : std::vector<uint32_t>{0, i_embd_0, i_embd_1, n_tokens};
+
+ llama_memory_clear(llama_get_memory(lctx), true);
+ llama_batch_ext_ptr batch(llama_batch_ext_init(lctx));
+
+ ret.clear();
+ ret.reserve(n_tokens*n_vocab);
+ for (size_t c = 0; c + 1 < bounds.size(); c++) {
+ llama_batch_ext_clear(batch.get());
+ for (uint32_t i = bounds[c]; i < bounds[c + 1]; i++) {
+ const bool is_embd = i >= i_embd_0 && i < i_embd_1;
+ const int32_t idx = is_embd
+ ? llama_batch_ext_add_embd(batch.get(), 0, { embd.data() + (size_t) (i - i_embd_0)*n_embd, 1, n_embd })
+ : llama_batch_ext_add_token(batch.get(), 0, tokens[i]);
+ GGML_ASSERT(idx >= 0);
+ const llama_pos pos[4] = { (llama_pos) i, (llama_pos) i, (llama_pos) i, 0 };
+ llama_batch_ext_set_pos(batch.get(), idx, pos);
+ llama_batch_ext_set_output_logits(batch.get(), idx, true);
+ }
+ const int32_t err = llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
+ if (err != 0) {
+ return err;
+ }
+ for (uint32_t i = 0; i < bounds[c + 1] - bounds[c]; i++) {
+ const float * logits_ith = llama_get_logits_ith(lctx, i);
+ ret.insert(ret.end(), logits_ith, logits_ith + n_vocab);
+ }
+ }
+ return 0;
+}
+
static bool check_causal_attn_toggle(
llama_model * model, llama_context * lctx, const std::vector<llama_token> & tokens) {
const uint32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model));
@@ -833,15 +877,15 @@ static int test_backends(const std::string & arch_filter, const size_t seed, con
max_arch_name_length = std::max(max_arch_name_length, strlen(llm_arch_name(arch)));
}
- const std::string template_header = std::string("|%" + std::to_string(max_arch_name_length) + "s|%") + std::to_string(max_device_label_length) + "s|%6s|%15s|%9s|\n";
+ const std::string template_header = std::string("|%" + std::to_string(max_arch_name_length) + "s|%") + std::to_string(max_device_label_length) + "s|%6s|%15s|%9s|%15s|\n";
const std::string template_row_cfg = std::string("|%" + std::to_string(max_arch_name_length) + "s|%") + std::to_string(max_device_label_length) + "s|%6s|";
- const std::string template_row_res = "%15s %10s|%20s|\n";
+ const std::string template_row_res = "%15s %10s|%20s|%15s %10s|\n";
bool all_ok = true;
size_t n_tests = 0;
size_t n_failed = 0;
common_log_flush(common_log_main());
- LOG(template_header.c_str(), "Model arch.", "Device", "Config", "NMSE vs. CPU", "Roundtrip");
+ LOG(template_header.c_str(), "Model arch.", "Device", "Config", "NMSE vs. CPU", "Roundtrip", "Mixed batch");
LOG("|");
for (size_t i = 0; i < max_arch_name_length; i++) {
LOG("-");
@@ -850,7 +894,7 @@ static int test_backends(const std::string & arch_filter, const size_t seed, con
for (size_t i = 0; i < max_device_label_length; i++) {
LOG("-");
}
- LOG("|------|---------------|---------|\n");
+ LOG("|------|---------------|---------|---------------|\n");
for (const llm_arch & arch : llm_arch_all()) {
if (arch == LLM_ARCH_UNKNOWN) {
continue;
@@ -889,7 +933,9 @@ static int test_backends(const std::string & arch_filter, const size_t seed, con
std::vector<float> logits_dev;
std::string status_nmse = "\033[1;33mSKIP\033[0m";
std::string status_roundtrip = "\033[1;33mSKIP\033[0m";
+ std::string status_mixed = "\033[1;33mSKIP\033[0m";
char nmse_str[12] = {0};
+ char mixed_str[12] = {0};
bool skip = !arch_supported(arch) || (dc.split_mode == LLAMA_SPLIT_MODE_TENSOR && dc.devs.empty());
bool test_executed = false;
@@ -910,6 +956,43 @@ static int test_backends(const std::string & arch_filter, const size_t seed, con
test_ok = false;
status_nmse = "\033[1;31mFAIL\033[0m";
}
+
+ // chunked decode matches a single batch only with causal attention over a memory
+ llama_context * lctx_dev = model_and_ctx_dev.second.get();
+ if (!encode && llama_get_memory(lctx_dev) != nullptr) {
+ std::vector<float> embd_mixed((size_t) llama_model_n_embd_inp(model_and_ctx_dev.first.get())*tokens.size()/4);
+ std::mt19937 gen(seed);
+ std::normal_distribution<float> dis(0.0f, stdev);
+ for (float & v : embd_mixed) {
+ v = dis(gen);
+ }
+ std::vector<float> logits_mixed;
+ std::vector<float> logits_chunks;
+ if (llm_arch_supports_mixed_batch(arch)) {
+ if (get_logits_mixed(model_and_ctx_dev.first.get(), lctx_dev, tokens, embd_mixed, false, logits_chunks) != 0 ||
+ get_logits_mixed(model_and_ctx_dev.first.get(), lctx_dev, tokens, embd_mixed, true, logits_mixed) != 0) {
+ throw std::runtime_error("failed to decode mixed batch");
+ }
+ const double nmse_mixed = nmse(logits_chunks, logits_mixed);
+ snprintf(mixed_str, sizeof(mixed_str), "(%.2e)", nmse_mixed);
+ status_mixed = "\033[1;32mOK\033[0m";
+ if (nmse_mixed > 1e-4) {
+ test_ok = false;
+ status_mixed = "\033[1;31mFAIL\033[0m";
+ }
+ } else {
+ // must be rejected as an invalid batch, mute the expected error log
+ ud.verbosity = LOG_LEVEL_OUTPUT;
+ const int32_t err = get_logits_mixed(model_and_ctx_cpu.first.get(), model_and_ctx_cpu.second.get(), tokens, embd_mixed, true, logits_mixed);
+ ud.verbosity = verbosity;
+ if (err != -1) {
+ test_ok = false;
+ status_mixed = "\033[1;31mFAIL\033[0m";
+ }
+ }
+ }
+
+ // runs after the mixed batch check, as it leaves the context with non-causal attention
if (!encode && !check_causal_attn_toggle(model_and_ctx_dev.first.get(), model_and_ctx_dev.second.get(), tokens)) {
if (test_ok) {
status_nmse = "\033[1;31mFAIL\033[0m (toggle)";
@@ -954,7 +1037,7 @@ static int test_backends(const std::string & arch_filter, const size_t seed, con
}
// log the results for this test case
- LOG(template_row_res.c_str(), status_nmse.c_str(), nmse_str, status_roundtrip.c_str());
+ LOG(template_row_res.c_str(), status_nmse.c_str(), nmse_str, status_roundtrip.c_str(), status_mixed.c_str(), mixed_str);
}
}
}