mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-02 19:07:25 -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));
|
||||
}
|
||||
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user