support shared prompt prefix

This commit is contained in:
Xuan Son Nguyen
2026-10-01 20:06:04 +02:00
parent 85eae8aeee
commit d244cdcf3d
5 changed files with 90 additions and 0 deletions
+39
View File
@@ -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));
}
+36
View File
@@ -389,3 +389,39 @@ json server_decision_context::format_answer(const server_decision_question & que
}
return answer;
}
//
// shared prompt prefix
//
std::vector<server_task> server_decision_group_tasks(std::vector<server_task> && tasks, size_t n_slots) {
n_slots = std::max(n_slots, (size_t) 1);
std::vector<server_task> 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;
}
+8
View File
@@ -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<server_decision_question> 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_task> server_decision_group_tasks(std::vector<server_task> && tasks, size_t n_slots);
+2
View File
@@ -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<server_task> 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;
@@ -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