mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-03 19:37:29 -05:00
server : allow RANK pooling batch splitting for causal LLM rerankers (ie. Qwen3 and Qwen3-VL) (#28876)
* server : allow splitting RANK pooling for causal LLM rerankers Rerank models fall into two categories: bidirectional cross-encoders (BERT, etc.) that require all tokens in a single physical batch, and causal LLMs repurposed as rerankers (Qwen3, Qwen3-VL) that can use chunked prefill like any other decoder. Previously the server rejected all RANK-pooling inputs larger than n_ubatch, and the graph builder hardcoded QWEN3/QWEN3VL arch checks to determine last-token pooling. This broke long-document and multimodal reranking for causal models. Fix: expose llama_get_causal_attn(ctx) so the server can check the effective runtime attention type (reflecting any --attention override or set_causal_attn call). Also expose llama_model_is_causal(model) for querying the static architectural property from GGUF metadata. can_split() now permits chunked prefill for RANK pooling when the context is causal. The graph builder's inline arch check is replaced with the same cparams.causal_attn predicate, removing the duplication. Assisted-by: Opencode/Qwen3.8-27B * remove unused llama_model_is_causal, fix whitespace Assisted-by: opencode --------- Co-authored-by: timothywang21 <timothywang21@users.noreply.github.com>
This commit is contained in:
co-authored by
timothywang21
parent
a97cce86a8
commit
4da6337767
@@ -1108,6 +1108,9 @@ extern "C" {
|
||||
// If set to true, the model will only attend to the past tokens
|
||||
LLAMA_API void llama_set_causal_attn(struct llama_context * ctx, bool causal_attn);
|
||||
|
||||
// Returns whether the context is currently using causal attention
|
||||
LLAMA_API bool llama_get_causal_attn(const struct llama_context * ctx);
|
||||
|
||||
// Set whether the model is in warmup mode or not
|
||||
// If true, all model tensors are activated during llama_decode() to load and cache their weights.
|
||||
//
|
||||
|
||||
@@ -1259,6 +1259,10 @@ void llama_context::set_causal_attn(bool value) {
|
||||
sched_need_reserve = true;
|
||||
}
|
||||
|
||||
bool llama_context::get_causal_attn() const {
|
||||
return cparams.causal_attn;
|
||||
}
|
||||
|
||||
void llama_context::set_warmup(bool value) {
|
||||
LLAMA_LOG_DEBUG("%s: value = %d\n", __func__, value);
|
||||
|
||||
@@ -3933,6 +3937,10 @@ void llama_set_causal_attn(llama_context * ctx, bool causal_attn) {
|
||||
ctx->set_causal_attn(causal_attn);
|
||||
}
|
||||
|
||||
bool llama_get_causal_attn(const llama_context * ctx) {
|
||||
return ctx->get_causal_attn();
|
||||
}
|
||||
|
||||
void llama_set_warmup(llama_context * ctx, bool warmup) {
|
||||
ctx->set_warmup(warmup);
|
||||
}
|
||||
|
||||
@@ -103,6 +103,8 @@ struct llama_context {
|
||||
const llama_token * get_sampled_candidates_ith(int32_t idx);
|
||||
size_t get_sampled_candidates_count(int32_t idx);
|
||||
|
||||
bool get_causal_attn() const;
|
||||
|
||||
void attach_threadpool(
|
||||
ggml_threadpool_t threadpool,
|
||||
ggml_threadpool_t threadpool_batch);
|
||||
|
||||
+1
-1
@@ -297,7 +297,7 @@ void llm_graph_input_cls::set_input(const llama_ubatch * ubatch) {
|
||||
|
||||
const bool last = (
|
||||
cparams.pooling_type == LLAMA_POOLING_TYPE_LAST ||
|
||||
(cparams.pooling_type == LLAMA_POOLING_TYPE_RANK && (arch == LLM_ARCH_QWEN3 || arch == LLM_ARCH_QWEN3VL)) // qwen3 reranking & embedding models use last token
|
||||
(cparams.pooling_type == LLAMA_POOLING_TYPE_RANK && cparams.causal_attn)
|
||||
);
|
||||
|
||||
for (int i = 0; i < n_tokens; ++i) {
|
||||
|
||||
@@ -435,15 +435,26 @@ struct server_slot {
|
||||
return task->need_embd();
|
||||
}
|
||||
|
||||
// if the context does not have a memory module then all embeddings have to be computed within a single ubatch
|
||||
// also we cannot split if the pooling would require any past tokens
|
||||
// (MTP supports splitting — uses task->need_embd() not need_embd())
|
||||
bool can_split() const {
|
||||
GGML_ASSERT(task);
|
||||
|
||||
return
|
||||
!task->need_embd() ||
|
||||
(llama_get_memory(ctx_tgt) && llama_pooling_type(ctx_tgt) == LLAMA_POOLING_TYPE_LAST);
|
||||
// MTP supports splitting - uses task->need_embd() not need_embd()
|
||||
if (!task->need_embd()) {
|
||||
return true;
|
||||
}
|
||||
// if the context does not have a memory module then all embeddings have to be computed within a single ubatch
|
||||
if (!llama_get_memory(ctx_tgt)) {
|
||||
return false;
|
||||
}
|
||||
// context can be chunked/split if the pooling type is LAST
|
||||
const auto pooling = llama_pooling_type(ctx_tgt);
|
||||
if (pooling == LLAMA_POOLING_TYPE_LAST) {
|
||||
return true;
|
||||
}
|
||||
// causal rerankers read the last token and have a KV cache, so they can also be chunked/split.
|
||||
if (pooling == LLAMA_POOLING_TYPE_RANK && llama_get_causal_attn(ctx_tgt)) {
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
bool can_batch_with(server_slot & other_slot) const {
|
||||
|
||||
Reference in New Issue
Block a user