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));
}