-
Notifications
You must be signed in to change notification settings - Fork 21.5k
server: fix n_cmpl not skipping processing prompt #18663
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from 7 commits
a9d7bcb
d7c27d4
59dda88
439c3b5
91fd50b
f2d988d
a4854f0
9ceb268
aef22e7
cc5cafe
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -153,7 +153,7 @@ struct server_slot { | |
|
|
||
| common_sampler_ptr smpl; | ||
|
|
||
| llama_token sampled; // in speculative mode, this is the last accepted token | ||
| llama_token sampled; // in speculative mode, this is the last accepted token | ||
| llama_tokens drafted; | ||
|
|
||
| // stats | ||
|
|
@@ -201,6 +201,31 @@ struct server_slot { | |
| alora_invocation_start = -1; | ||
| } | ||
|
|
||
| void clear() { | ||
| llama_memory_seq_rm(llama_get_memory(ctx), id, -1, -1); | ||
| prompt.tokens.clear(); | ||
| } | ||
|
|
||
| void init_sampler() const { | ||
| const int64_t t_start = ggml_time_us(); | ||
|
|
||
| common_sampler_reset(smpl.get()); | ||
|
|
||
| int n_text = 0; | ||
|
|
||
| for (int i = 0; i < (int) prompt.tokens.size(); i++) { | ||
| const llama_token id = prompt.tokens[i]; | ||
|
|
||
| if (id != LLAMA_TOKEN_NULL) { | ||
| common_sampler_accept(smpl.get(), id, false); | ||
| n_text++; | ||
| } | ||
| } | ||
|
|
||
| SLT_INF(*this, "init sampler, took %0.2f ms, tokens: text = %d, total = %d\n", | ||
| (ggml_time_us() - t_start) / 1000.0, n_text, (int) prompt.tokens.size()); | ||
| } | ||
|
|
||
| bool need_embd() const { | ||
| GGML_ASSERT(task); | ||
|
|
||
|
|
@@ -288,11 +313,11 @@ struct server_slot { | |
|
|
||
| // note: a slot can also be either a parent or a child | ||
| bool is_parent() const { | ||
| return is_processing() && task->n_children > 0; | ||
| return task->n_children > 0; | ||
| } | ||
|
|
||
| bool is_child() const { | ||
| return is_processing() && task->id_parent >= 0; | ||
| return task->id_parent >= 0; | ||
| } | ||
|
|
||
| void release() { | ||
|
|
@@ -301,10 +326,18 @@ struct server_slot { | |
|
|
||
| SLT_INF(*this, "stop processing: n_tokens = %d, truncated = %d\n", prompt.n_tokens(), truncated); | ||
|
|
||
| t_last_used = ggml_time_us(); | ||
| t_last_used = ggml_time_us(); | ||
| t_token_generation = (ggml_time_us() - t_start_generation) / 1e3; | ||
|
|
||
| state = SLOT_STATE_IDLE; | ||
|
|
||
| // do not keep context of the child slots - the parent's context is enough | ||
| if (is_child()) { | ||
| SLT_INF(*this, "clearing child slot with %zu tokens\n", prompt.tokens.size()); | ||
|
|
||
| clear(); | ||
| } | ||
|
|
||
| task_prev = std::move(task); | ||
| task.reset(); | ||
|
|
||
|
|
@@ -425,14 +458,22 @@ struct server_slot { | |
| } | ||
|
|
||
| void copy_state_to(server_slot & other) const { | ||
| llama_memory_seq_rm(llama_get_memory(ctx), other.id, 0, -1); | ||
| llama_memory_seq_cp(llama_get_memory(ctx), id, other.id, 0, -1); | ||
| GGML_ASSERT(state == SLOT_STATE_DONE_PROMPT); | ||
|
|
||
| llama_memory_seq_rm(llama_get_memory(ctx), other.id, -1, -1); | ||
| llama_memory_seq_cp(llama_get_memory(ctx), id, other.id, -1, -1); | ||
|
|
||
| other.n_decoded = n_decoded; | ||
| other.n_remaining = n_remaining; | ||
| other.i_batch = i_batch; | ||
|
|
||
| other.t_start_process_prompt = t_start_process_prompt; | ||
| other.t_prompt_processing = t_prompt_processing; | ||
| other.n_prompt_tokens_cache = n_prompt_tokens_cache; | ||
| other.n_prompt_tokens_processed = n_prompt_tokens_processed; | ||
|
|
||
| other.prompt = prompt.clone(); | ||
| other.init_sampler(); | ||
| } | ||
| }; | ||
|
|
||
|
|
@@ -1005,15 +1046,14 @@ struct server_context_impl { | |
| return ret; | ||
| } | ||
|
|
||
| void clear_slot(server_slot & slot, bool allow_processing = false) const { | ||
| static void clear_slot(server_slot & slot, bool allow_processing = false) { | ||
|
Collaborator
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I think this function be removed now, as the logic can be moved to A static function with the first argument being class instance can always be converted to a class method
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Hm, I think I broke something in 9ceb268 - server tests are failing locally. |
||
| if (!allow_processing) { | ||
| GGML_ASSERT(!slot.is_processing()); | ||
| } | ||
|
|
||
| SLT_WRN(slot, "clearing slot with %zu tokens\n", slot.prompt.tokens.size()); | ||
|
|
||
| llama_memory_seq_rm(llama_get_memory(ctx), slot.id, -1, -1); | ||
| slot.prompt.tokens.clear(); | ||
| slot.clear(); | ||
| } | ||
|
|
||
| // return true if at least one slot has been cleared | ||
|
|
@@ -1182,7 +1222,7 @@ struct server_context_impl { | |
| ? SLOT_STATE_WAIT_OTHER // wait for the parent to process prompt | ||
| : SLOT_STATE_STARTED; | ||
|
|
||
| SLT_INF(slot, "%s", "processing task\n"); | ||
| SLT_INF(slot, "processing task, is_child = %d\n", slot.is_child()); | ||
|
|
||
| return true; | ||
| } | ||
|
|
@@ -2053,6 +2093,27 @@ struct server_context_impl { | |
| continue; | ||
| } | ||
|
|
||
| // check if this is a child slot | ||
| if (slot.state == SLOT_STATE_WAIT_OTHER) { | ||
| SLT_DBG(slot, "%s", "waiting for parent slot to complete\n"); | ||
| continue; | ||
| } | ||
|
|
||
| // wait for all children to be launched | ||
| if (slot.is_parent()) { | ||
| int n_launched = 0; | ||
| for (auto & other : slots) { | ||
|
Collaborator
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I'm a little bit worry that this nested loop will be invoked on each new token of the parent slot. Probably move this inside the The idea is that transition from SLOT_STATE_STARTED to SLOT_STATE_PROCESSING_PROMPT is only permitted if all child slots are launched
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I have a few more changes to improve this logic, but will push them in a follow-up PR because they restructure the
Collaborator
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Yes that sounds good to me. After your refactoring, I will attempt to break the transition down inside small segments. Actually I talked about this point earlier via DM: server_slot::run_pre_decode(...) {
if (state == SLOT_STATE_A) {
// do transition_A_to_B
return SLOT_STATE_B;
}
...
}And inside for (auto & slot : slots) {
slot.state = slot.run_pre_decode(batch, ...);
}
llama_decode(batch);
for (auto & slot : slots) {
slot.state = slot.run_post_decode(batch, ...);
}The main benefit would be that (most) state transitions will be bound to / isolated to one slot, since transition functions will now be slot's member. The biggest benefit of this will be to define error boundary inside a transition function. So if one slot got an exception (currently, grammar system can throw one), we can shutdown one single slot instead of letting the server crash. |
||
| if (other.is_processing() && other.is_child() && other.task->id_parent == slot.task->id) { | ||
| ++n_launched; | ||
| } | ||
| } | ||
|
|
||
| if (n_launched < slot.task->n_children) { | ||
| SLT_DBG(slot, "waiting for children to be launched, n_children = %d, n_launched = %d\n", slot.task->n_children, n_launched); | ||
| continue; | ||
| } | ||
| } | ||
|
|
||
| // this slot still has a prompt to be processed | ||
| if (slot.state == SLOT_STATE_PROCESSING_PROMPT || slot.state == SLOT_STATE_STARTED) { | ||
| const auto & input_tokens = slot.task->tokens; | ||
|
|
@@ -2455,16 +2516,6 @@ struct server_context_impl { | |
|
|
||
| GGML_ASSERT(batch.n_tokens > 0); | ||
|
|
||
| common_sampler_reset(slot.smpl.get()); | ||
|
|
||
| // Process all prompt tokens through sampler system | ||
| for (int i = 0; i < slot.task->n_tokens(); ++i) { | ||
| llama_token id = input_tokens[i]; | ||
| if (id != LLAMA_TOKEN_NULL) { | ||
| common_sampler_accept(slot.smpl.get(), id, false); | ||
| } | ||
| } | ||
|
|
||
| // extract the logits only for the last token | ||
| batch.logits[batch.n_tokens - 1] = true; | ||
|
|
||
|
|
@@ -2473,6 +2524,8 @@ struct server_context_impl { | |
|
|
||
| SLT_INF(slot, "prompt done, n_tokens = %d, batch.n_tokens = %d\n", slot.prompt.n_tokens(), batch.n_tokens); | ||
|
|
||
| slot.init_sampler(); | ||
|
|
||
| const auto pos_min = llama_memory_seq_pos_min(llama_get_memory(ctx), slot.id); | ||
| const auto pos_max = llama_memory_seq_pos_max(llama_get_memory(ctx), slot.id); | ||
|
|
||
|
|
@@ -2519,11 +2572,6 @@ struct server_context_impl { | |
| } | ||
| } | ||
|
|
||
| if (batch.n_tokens == 0) { | ||
| SRV_WRN("%s", "no tokens to decode\n"); | ||
| return; | ||
| } | ||
|
|
||
| SRV_DBG("decoding batch, n_tokens = %d\n", batch.n_tokens); | ||
|
|
||
| if (slot_batched) { | ||
|
|
@@ -2540,6 +2588,10 @@ struct server_context_impl { | |
| llama_set_embeddings(ctx, slot_batched->need_embd()); | ||
| } | ||
|
|
||
| if (batch.n_tokens == 0) { | ||
| SRV_WRN("%s", "no tokens to decode\n"); | ||
| } | ||
|
|
||
| int32_t i_next = 0; | ||
|
|
||
| // process the created batch of tokens | ||
|
|
@@ -2615,27 +2667,34 @@ struct server_context_impl { | |
| // on successful decode, restore the original batch size | ||
| n_batch = llama_n_batch(ctx); | ||
|
|
||
| // handle `n_cmpl > 1` tasks - when the main prompt is processed, activate all child tasks too | ||
| for (auto & slot : slots) { | ||
| // may need to copy state to other slots | ||
| if (slot.state == SLOT_STATE_DONE_PROMPT && slot.is_parent()) { | ||
| std::vector<server_slot *> child_slots; | ||
| SLT_INF(slot, "parent task prompt done, n_children = %d\n", slot.task->n_children); | ||
|
|
||
| std::vector<server_slot *> children; | ||
| for (auto & other : slots) { | ||
| if (other.state == SLOT_STATE_WAIT_OTHER && slot.task->id == other.task->id_parent) { | ||
| child_slots.push_back(&other); | ||
| children.push_back(&other); | ||
| } | ||
| } | ||
|
|
||
| // we can only proceed if all child slots are having the correct tasks | ||
| if (child_slots.size() == slot.task->n_children) { | ||
| if (slot.task->n_children == (int) children.size()) { | ||
| // copy state to the child slots | ||
| for (auto & child : child_slots) { | ||
| SLT_INF(slot, "copying state to child %d\n", child->id); | ||
| for (auto & child : children) { | ||
| SLT_INF(slot, " - copying state to child %d\n", child->id); | ||
|
|
||
| GGML_ASSERT(child->state == SLOT_STATE_WAIT_OTHER); | ||
|
|
||
| slot.copy_state_to(*child); | ||
| child->state = SLOT_STATE_DONE_PROMPT; | ||
| } | ||
| } | ||
| } | ||
| } | ||
|
|
||
| for (auto & slot : slots) { | ||
| // optionally send prompt processing progress | ||
| if (slot.state == SLOT_STATE_PROCESSING_PROMPT || slot.state == SLOT_STATE_DONE_PROMPT) { | ||
| if (slot.task->params.stream && slot.task->params.return_progress) { | ||
|
|
@@ -2720,7 +2779,7 @@ struct server_context_impl { | |
| continue; | ||
| } | ||
|
|
||
| size_t n_draft = slot.drafted.size(); | ||
| const size_t n_draft = slot.drafted.size(); | ||
|
|
||
| // the accepted tokens from the speculation | ||
| const auto ids = common_sampler_sample_and_accept_n(slot.smpl.get(), ctx, slot.i_batch_dft, slot.drafted); | ||
|
|
@@ -2923,9 +2982,11 @@ std::unique_ptr<server_res_generator> server_routes::handle_completions_impl( | |
| task.params.oaicompat_cmpl_id = completion_id; | ||
| task.params.oaicompat_model = meta->model_name; | ||
|
|
||
| // prepare child tasks | ||
| if (task.params.n_cmpl > 1) { | ||
| task.n_children = task.params.n_cmpl - 1; | ||
| for (size_t j = 0; j < task.n_children; j++) { | ||
|
|
||
| for (int j = 0; j < task.n_children; j++) { | ||
| server_task child = task.create_child(task.id, rd.get_new_id()); | ||
|
|
||
| // use different sampling seed for each child | ||
|
|
@@ -2938,7 +2999,8 @@ std::unique_ptr<server_res_generator> server_routes::handle_completions_impl( | |
| } | ||
| } | ||
|
|
||
| tasks.push_back(std::move(task)); | ||
| // note: the parent task always launches first | ||
| tasks.insert(tasks.begin(), std::move(task)); | ||
| } | ||
|
|
||
| rd.post_tasks(std::move(tasks)); | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
This
is_processing()check also seemed redundant so removed it.