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:
Tim Wang
2026-09-27 23:28:10 +02:00
committed by GitHub
co-authored by timothywang21
parent a97cce86a8
commit 4da6337767
5 changed files with 32 additions and 8 deletions
+3
View File
@@ -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.
//
+8
View File
@@ -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);
}
+2
View File
@@ -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
View File
@@ -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) {
+18 -7
View File
@@ -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 {