Commit 4da633776 for llama.cpp
commit 4da6337767f973e2b4d0797e5b323d77d8565e4a
Author: Tim Wang <149349643+timothywang21@users.noreply.github.com>
Date: Sun Sep 27 17:28:10 2026 -0400
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>
diff --git a/include/llama.h b/include/llama.h
index 1805ed055..ce454df52 100644
--- a/include/llama.h
+++ b/include/llama.h
@@ -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.
//
diff --git a/src/llama-context.cpp b/src/llama-context.cpp
index 99e55da68..27b9a38d7 100644
--- a/src/llama-context.cpp
+++ b/src/llama-context.cpp
@@ -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);
}
diff --git a/src/llama-context.h b/src/llama-context.h
index b403b099b..35a7071ed 100644
--- a/src/llama-context.h
+++ b/src/llama-context.h
@@ -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);
diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp
index 0b3bab612..a806126ef 100644
--- a/src/llama-graph.cpp
+++ b/src/llama-graph.cpp
@@ -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) {
diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp
index e95fb63ab..611e82a6a 100644
--- a/tools/server/server-context.cpp
+++ b/tools/server/server-context.cpp
@@ -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 {