From d244cdcf3d292ab833169558b2b0a6cbf67bb746 Mon Sep 17 00:00:00 2001 From: Xuan Son Nguyen Date: Thu, 1 Oct 2026 20:06:04 +0200 Subject: [PATCH] support shared prompt prefix --- tools/server/server-context.cpp | 39 +++++++++++++++++++++++ tools/server/server-decision.cpp | 36 +++++++++++++++++++++ tools/server/server-decision.h | 8 +++++ tools/server/server-task.h | 2 ++ tools/server/tests/unit/test_systemone.py | 5 +++ 5 files changed, 90 insertions(+) diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index da719cdf43..c070f3c806 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -690,6 +690,14 @@ struct server_slot { return res; } + // the other slot continues from the tokens processed so far + void copy_prompt_to(server_slot & other) const { + mem.seq_rm(other.id, -1, -1); + mem.seq_cp(id, other.id, -1, -1); + + other.prompt = prompt.clone(); + } + void copy_state_to(server_slot & other) const { GGML_ASSERT(state == SLOT_STATE_DONE_PROMPT); @@ -3272,6 +3280,11 @@ private: // reuse any previously computed tokens that are common with the new prompt n_past = slot.prompt.tokens.get_common_prefix(input_tokens); + // the children start from the shared prefix, do not go past it + if (slot.task->n_tokens_shared > 0) { + n_past = std::min(n_past, slot.task->n_tokens_shared); + } + // if there is an alora invoked, don't cache after the invocation start if (slot.alora_invocation_start > 0) { SLT_DBG(slot, "only caching to alora invocation start (n_past = %d, alora_invocation_start = %d)\n", n_past, slot.alora_invocation_start); @@ -3497,6 +3510,24 @@ private: slot.mem.seq_rm(slot.id, p0, -1); + // shared prompt prefix: once it is processed, the children continue from it with their own prompt + bool wait_shared = false; + if (slot.task->n_tokens_shared > 0) { + const bool is_shared_done = slot.prompt.n_tokens() == slot.task->n_tokens_shared; + for (auto & other : slots) { + if (other.state != SLOT_STATE_WAIT_OTHER || other.task->id_parent != slot.task->id) { + continue; + } + if (is_shared_done) { + SLT_TRC(slot, " - copying shared prompt (%d tokens) to child %d\n", slot.prompt.n_tokens(), other.id); + slot.copy_prompt_to(other); + other.state = SLOT_STATE_STARTED; + } else { + wait_shared = true; + } + } + } + // If using an alora, there may be uncached tokens that come // before the invocation sequence. When this happens, the // tokens before the invocation sequence need to be @@ -3582,6 +3613,11 @@ private: break; // end of text chunk } + // stop at the end of the shared prefix, the children are started from this state + if (wait_shared && slot.prompt.n_tokens() == slot.task->n_tokens_shared) { + break; + } + // if this is an alora request with pre-invocation // tokens that are not cached, we need to stop filling // this batch at those pre-invocation tokens. @@ -5311,6 +5347,9 @@ void server_routes::init_routes() { decision.fill_task(body.at("state"), question, task); tasks.push_back(std::move(task)); } + if (decision.can_share_prompt()) { + tasks = server_decision_group_tasks(std::move(tasks), params.n_parallel); + } rd.post_tasks(std::move(tasks)); } diff --git a/tools/server/server-decision.cpp b/tools/server/server-decision.cpp index f113ab7048..49da30c352 100644 --- a/tools/server/server-decision.cpp +++ b/tools/server/server-decision.cpp @@ -389,3 +389,39 @@ json server_decision_context::format_answer(const server_decision_question & que } return answer; } + +// +// shared prompt prefix +// + +std::vector server_decision_group_tasks(std::vector && tasks, size_t n_slots) { + n_slots = std::max(n_slots, (size_t) 1); + + std::vector groups; + for (size_t i = 0; i < tasks.size(); i += n_slots) { + const size_t end = std::min(tasks.size(), i + n_slots); + server_task & parent = tasks[i]; + + // every task must have at least one token of its own to evaluate + size_t n_shared = parent.tokens.size() - 1; + for (size_t j = i + 1; j < end; j++) { + n_shared = std::min(n_shared, parent.tokens.get_common_prefix(tasks[j].tokens)); + n_shared = std::min(n_shared, tasks[j].tokens.size() - 1); + } + + if (end - i < 2 || n_shared == 0) { + for (size_t j = i; j < end; j++) { + groups.push_back(std::move(tasks[j])); + } + continue; + } + + parent.n_tokens_shared = n_shared; + for (size_t j = i + 1; j < end; j++) { + tasks[j].id_parent = parent.id; + parent.child_tasks.push_back(std::move(tasks[j])); + } + groups.push_back(std::move(parent)); + } + return groups; +} diff --git a/tools/server/server-decision.h b/tools/server/server-decision.h index 3dedb053df..99f2a08dc6 100644 --- a/tools/server/server-decision.h +++ b/tools/server/server-decision.h @@ -44,6 +44,9 @@ struct server_decision_context { // true if the result is read from the embeddings of each token bool need_embd() const { return type == SERVER_DECISION_TYPE_LAYA; } + // true if the questions of a request start with the same tokens, and the model can continue from them + bool can_share_prompt() const { return type == SERVER_DECISION_TYPE_OPENJEV; } + // throw std::invalid_argument on bad input std::vector parse_questions(const json & body) const; @@ -76,3 +79,8 @@ private: float get_temperature(const server_decision_question & question) const; }; + +// group the tasks so that the common prefix of their prompts is evaluated only once +// each group is one parent and its children, it takes at most n_slots slots +// note: the order of the tasks is preserved +std::vector server_decision_group_tasks(std::vector && tasks, size_t n_slots); diff --git a/tools/server/server-task.h b/tools/server/server-task.h index 56a2eec5b6..05994ab5e5 100644 --- a/tools/server/server-task.h +++ b/tools/server/server-task.h @@ -149,6 +149,8 @@ struct server_task { // temporary store of child tasks for scheduling // note: accessing to elements is invalid after the task is moved to server_slot std::vector child_tasks; + // if set on a parent, the children have their own prompt and only share its first n_tokens_shared tokens + int32_t n_tokens_shared = 0; // used by SERVER_TASK_TYPE_INFERENCE task_params params; diff --git a/tools/server/tests/unit/test_systemone.py b/tools/server/tests/unit/test_systemone.py index cf983cce91..fd7a698c8f 100644 --- a/tools/server/tests/unit/test_systemone.py +++ b/tools/server/tests/unit/test_systemone.py @@ -116,3 +116,8 @@ def test_systemone_requires_embedding(): "questions": TEST_QUESTIONS, }) assert res.status_code == 501 + + +# TODO: test the shared prompt prefix, it needs a small model of a type that supports it (e.g. openjev) +# it can be checked with GET /metrics: for one request, prompt_tokens_cached_total must grow by +# (shared tokens * number of child tasks) and prompt_tokens_total + prompt_tokens_cached_total == usage.input_tokens