mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-02 02:47:26 -05:00
support shared prompt prefix
This commit is contained in:
@@ -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));
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user