diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index b5aad4f79c..d7a2315faa 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -193,16 +193,20 @@ struct server_batch { }; struct server_slot_stats { + uint64_t n_prompt_cached = 0; + uint64_t n_prompt_processed = 0; + uint64_t n_predict = 0; + // Speculative decoding stats (mirror server_metrics) - int32_t n_draft_tokens = 0; - int32_t n_draft_accepted = 0; - int32_t n_draft_verif_steps = 0; - std::vector n_accepted_per_pos; + uint64_t n_draft_tokens = 0; + uint64_t n_draft_accepted = 0; + uint64_t n_draft_verif_steps = 0; + std::vector n_accepted_per_pos; // these are absolute timestamps (in us) - int64_t t_start = 0; - int64_t t_prompt_last = 0; - int64_t t_gen_last = 0; + uint64_t t_start = 0; + uint64_t t_prompt_last = 0; + uint64_t t_gen_last = 0; // can only move one direction: start -> prompt -> gen void update_prompt_start() { @@ -271,14 +275,12 @@ struct server_slot { int64_t t_last_used = -1; // generation props - int32_t n_ctx = 0; // context size per slot - int32_t n_keep = 0; - int32_t n_decoded = 0; - int32_t n_remaining = -1; - int32_t i_batch = -1; + int32_t n_ctx = 0; // context size per slot + int32_t n_keep = 0; + int32_t i_batch = -1; - int32_t n_prompt_tokens_cache = 0; - int32_t n_prompt_tokens_processed = 0; + // effective generation limit for the current task, -1 means unlimited + int32_t n_predict_max = -1; size_t last_nl_pos = 0; @@ -372,8 +374,6 @@ struct server_slot { spec_is_replay = false; - n_prompt_tokens_cache = 0; - last_nl_pos = 0; generated_text = ""; has_new_line = false; @@ -458,22 +458,13 @@ struct server_slot { && are_lora_equal(lora, other_slot.lora); } - bool has_budget(const common_params & global_params) { - GGML_ASSERT(task); + // returns -1 if the generation is limitless + int32_t n_remaining() const { + return n_predict_max == -1 ? -1 : n_predict_max - (int32_t) stats.n_predict; + } - if (task->params.n_predict == -1 && global_params.n_predict == -1) { - return true; // limitless - } - - n_remaining = -1; - - if (task->params.n_predict != -1) { - n_remaining = task->params.n_predict - n_decoded; - } else if (global_params.n_predict != -1) { - n_remaining = global_params.n_predict - n_decoded; - } - - return n_remaining > 0; // no budget + bool has_budget() const { + return n_predict_max == -1 || n_remaining() > 0; } bool is_processing() const { @@ -505,8 +496,8 @@ struct server_slot { // also, need to leave space for 1 extra token to allow context shifts int n_draft_max = n_ctx - prompt.n_tokens() - 2; - if (n_remaining > 0) { - n_draft_max = std::min(n_draft_max, n_remaining - 1); + if (n_remaining() > 0) { + n_draft_max = std::min(n_draft_max, n_remaining() - 1); } SLT_DBG(*this, "max possible draft: %d\n", n_draft_max); @@ -580,17 +571,17 @@ struct server_slot { const double t_token_generation = stats.t_gen_ms(); result_timings timings; - timings.cache_n = n_prompt_tokens_cache; + timings.cache_n = stats.n_prompt_cached; - timings.prompt_n = n_prompt_tokens_processed; + timings.prompt_n = stats.n_prompt_processed; timings.prompt_ms = t_prompt_processing; - timings.prompt_per_token_ms = t_prompt_processing / n_prompt_tokens_processed; - timings.prompt_per_second = stats.n_prompt_tps(n_prompt_tokens_processed); + timings.prompt_per_token_ms = t_prompt_processing / stats.n_prompt_processed; + timings.prompt_per_second = stats.n_prompt_tps(stats.n_prompt_processed); - timings.predicted_n = n_decoded; + timings.predicted_n = stats.n_predict; timings.predicted_ms = t_token_generation; - timings.predicted_per_token_ms = t_token_generation / n_decoded; - timings.predicted_per_second = stats.n_gen_tps(n_decoded); + timings.predicted_per_token_ms = t_token_generation / stats.n_predict; + timings.predicted_per_second = stats.n_gen_tps(stats.n_predict); // Add speculative metrics if (stats.n_draft_tokens > 0) { @@ -633,7 +624,7 @@ struct server_slot { } void print_timings_tg() { - if (n_decoded < 100) { + if (stats.n_predict < 100) { return; } @@ -643,19 +634,19 @@ struct server_slot { return; } - const double n_gen_second = stats.n_gen_tps(n_decoded); - const double n_gen_second_win = 1e6 / (t_now - t_print_last) * (n_decoded - n_decoded_last); + const double n_gen_second = stats.n_gen_tps(stats.n_predict); + const double n_gen_second_win = 1e6 / (t_now - t_print_last) * (stats.n_predict - n_decoded_last); t_print_last = t_now; - n_decoded_last = n_decoded; + n_decoded_last = stats.n_predict; - SLT_INF(*this, "n_decoded = %6d, tg = %6.2f t/s, tg_3s = %6.2f t/s\n", n_decoded, n_gen_second, n_gen_second_win); + SLT_INF(*this, "n_decoded = %6d, tg = %6.2f t/s, tg_3s = %6.2f t/s\n", (int) stats.n_predict, n_gen_second, n_gen_second_win); } void print_timings_pp() const { const double t_prompt_processing = stats.t_prompt_ms(); - const double n_prompt_second = stats.n_prompt_tps(n_prompt_tokens_processed); + const double n_prompt_second = stats.n_prompt_tps(stats.n_prompt_processed); const double f_progress = (float) prompt.n_tokens() / task->n_tokens(); if (t_prompt_processing < 3000.0) { @@ -663,30 +654,30 @@ struct server_slot { } SLT_INF(*this, "prompt processing, n_tokens = %6d, progress = %.2f, t = %6.2f s / %.2f tokens per second\n", - n_prompt_tokens_processed, f_progress, t_prompt_processing / 1e3, n_prompt_second); + (int) stats.n_prompt_processed, f_progress, t_prompt_processing / 1e3, n_prompt_second); } void print_timings() const { const double t_prompt_processing = stats.t_prompt_ms(); const double t_token_generation = stats.t_gen_ms(); - const double t_prompt = t_prompt_processing / n_prompt_tokens_processed; - const double n_prompt_second = stats.n_prompt_tps(n_prompt_tokens_processed); + const double t_prompt = t_prompt_processing / stats.n_prompt_processed; + const double n_prompt_second = stats.n_prompt_tps(stats.n_prompt_processed); - const double t_gen = t_token_generation / n_decoded; - const double n_gen_second = stats.n_gen_tps(n_decoded); + const double t_gen = t_token_generation / stats.n_predict; + const double n_gen_second = stats.n_gen_tps(stats.n_predict); SLT_INF(*this, "prompt eval time = %10.2f ms / %5d tokens (%8.2f ms per token, %8.2f tokens per second)\n", - t_prompt_processing, n_prompt_tokens_processed, t_prompt, n_prompt_second); + t_prompt_processing, (int) stats.n_prompt_processed, t_prompt, n_prompt_second); SLT_INF(*this, " eval time = %10.2f ms / %5d tokens (%8.2f ms per token, %8.2f tokens per second)\n", - t_token_generation, n_decoded, t_gen, n_gen_second); + t_token_generation, (int) stats.n_predict, t_gen, n_gen_second); SLT_INF(*this, " total time = %10.2f ms / %5d tokens\n", - t_prompt_processing + t_token_generation, n_prompt_tokens_processed + n_decoded); + t_prompt_processing + t_token_generation, (int) (stats.n_prompt_processed + stats.n_predict)); SLT_INF(*this, " graphs reused = %10d\n", @@ -735,15 +726,15 @@ struct server_slot { if (ptask) { res["id_task"] = ptask->id; res["n_prompt_tokens"] = (int32_t) prompt.tokens.size(); - res["n_prompt_tokens_processed"] = n_prompt_tokens_processed; - res["n_prompt_tokens_cache"] = n_prompt_tokens_cache; + res["n_prompt_tokens_processed"] = stats.n_prompt_processed; + res["n_prompt_tokens_cache"] = stats.n_prompt_cached; res["params"] = ptask->params.to_json(only_metrics); res["next_token"] = { { {"has_next_token", has_next_token}, {"has_new_line", has_new_line}, - {"n_remain", n_remaining}, - {"n_decoded", n_decoded}, + {"n_remain", n_remaining()}, + {"n_decoded", stats.n_predict}, } }; @@ -762,13 +753,9 @@ struct server_slot { mem.seq_rm(other.id, -1, -1); mem.seq_cp(id, other.id, -1, -1); - other.n_decoded = n_decoded; - other.n_remaining = n_remaining; - other.i_batch = i_batch; + other.i_batch = i_batch; - other.stats = stats; - other.n_prompt_tokens_cache = n_prompt_tokens_cache; - other.n_prompt_tokens_processed = n_prompt_tokens_processed; + other.stats = stats; other.prompt = prompt.clone(); other.init_sampler(); @@ -871,8 +858,8 @@ struct server_slot { void server_metrics::on_prompt_eval(const server_slot & slot) { const double t_ms = slot.stats.t_prompt_ms(); - prompt .add(slot.n_prompt_tokens_processed, t_ms); - prompt_bucket.add(slot.n_prompt_tokens_processed, t_ms); + prompt .add(slot.stats.n_prompt_processed, t_ms); + prompt_bucket.add(slot.stats.n_prompt_processed, t_ms); n_tokens_max = std::max(n_tokens_max, (uint64_t) slot.prompt.n_tokens()); } @@ -880,8 +867,8 @@ void server_metrics::on_prompt_eval(const server_slot & slot) { void server_metrics::on_prediction(const server_slot & slot) { const double t_ms = slot.stats.t_gen_ms(); - predict .add(slot.n_decoded, t_ms); - predict_bucket.add(slot.n_decoded, t_ms); + predict .add(slot.stats.n_predict, t_ms); + predict_bucket.add(slot.stats.n_predict, t_ms); n_draft_tokens += slot.stats.n_draft_tokens; n_draft_accepted += slot.stats.n_draft_accepted; @@ -1858,6 +1845,9 @@ private: slot.smpl.reset(); } + // the per-request limit takes priority over the global one + slot.n_predict_max = task.params.n_predict != -1 ? task.params.n_predict : params_base.n_predict; + slot.task = std::make_unique(std::move(task)); slot.state = slot.task->is_child() @@ -1930,15 +1920,15 @@ private: slot.has_next_token = false; SLT_DBG(slot, "stopped due to running out of context capacity, prompt.n_tokens() = %d, task.n_tokens = %d, n_decoded = %d, n_ctx = %d\n", - slot.prompt.n_tokens(), slot.task->n_tokens(), slot.n_decoded, slot.n_ctx); + slot.prompt.n_tokens(), slot.task->n_tokens(), (int) slot.stats.n_predict, slot.n_ctx); } // check the limits - if (slot.n_decoded > 0 && slot.has_next_token && !slot.has_budget(params_base)) { + if (slot.stats.n_predict > 0 && slot.has_next_token && !slot.has_budget()) { slot.stop = STOP_TYPE_LIMIT; slot.has_next_token = false; - SLT_DBG(slot, "stopped by limit, n_decoded = %d, n_predict = %d\n", slot.n_decoded, slot.task->params.n_predict); + SLT_DBG(slot, "stopped by limit, n_decoded = %d, n_predict = %d\n", (int) slot.stats.n_predict, slot.task->params.n_predict); } if (slot.has_new_line) { @@ -1962,7 +1952,7 @@ private: // cut the last line slot.generated_text.erase(pos, std::string::npos); - SLT_DBG(slot, "stopped by indentation limit, n_decoded = %d, n_indent = %d\n", slot.n_decoded, n_indent); + SLT_DBG(slot, "stopped by indentation limit, n_decoded = %d, n_indent = %d\n", (int) slot.stats.n_predict, n_indent); } } @@ -1986,7 +1976,7 @@ private: slot.stop = STOP_TYPE_LIMIT; slot.has_next_token = false; - SLT_DBG(slot, "stopped by time limit, n_decoded = %d, t_max_predict_ms = %d ms\n", slot.n_decoded, (int) slot.task->params.t_max_predict_ms); + SLT_DBG(slot, "stopped by time limit, n_decoded = %d, t_max_predict_ms = %d ms\n", (int) slot.stats.n_predict, (int) slot.task->params.t_max_predict_ms); } } @@ -1997,7 +1987,7 @@ private: SLT_DBG(slot, "%s", "stopped by EOS\n"); } - SLT_DBG(slot, "n_decoded = %d, n_remaining = %d, next token: %5d '%s'\n", slot.n_decoded, slot.n_remaining, result.tok, token_str.c_str()); + SLT_DBG(slot, "n_decoded = %d, n_remaining = %d, next token: %5d '%s'\n", (int) slot.stats.n_predict, slot.n_remaining(), result.tok, token_str.c_str()); return slot.has_next_token; // continue } @@ -2105,7 +2095,7 @@ private: if (is_progress) { res->is_progress = true; res->progress.total = slot.task->n_tokens(); - res->progress.cache = slot.n_prompt_tokens_cache; + res->progress.cache = slot.stats.n_prompt_cached; res->progress.processed = slot.prompt.tokens.size(); res->progress.time_ms = slot.stats.t_ellapsed_us() / 1000; } @@ -2116,9 +2106,9 @@ private: res->tokens = { tkn.tok }; } - res->n_decoded = slot.n_decoded; + res->n_decoded = slot.stats.n_predict; res->n_prompt_tokens = slot.task->n_tokens(); - res->n_prompt_tokens_cache = slot.n_prompt_tokens_cache; + res->n_prompt_tokens_cache = slot.stats.n_prompt_cached; res->post_sampling_probs = slot.task->params.post_sampling_probs; res->verbose = slot.task->params.verbose; @@ -2165,9 +2155,9 @@ private: res->response_fields = std::move(slot.task->params.response_fields); res->truncated = slot.truncated; - res->n_decoded = slot.n_decoded; + res->n_decoded = slot.stats.n_predict; res->n_prompt_tokens = slot.task->n_tokens(); - res->n_prompt_tokens_cache = slot.n_prompt_tokens_cache; + res->n_prompt_tokens_cache = slot.stats.n_prompt_cached; res->n_tokens_cached = slot.prompt.n_tokens(); res->has_new_line = slot.has_new_line; res->stopping_word = slot.stopping_word; @@ -3406,8 +3396,8 @@ private: SLT_WRN(slot, "n_past was set to %d\n", n_past); } - slot.n_prompt_tokens_cache = n_past; - slot.n_prompt_tokens_processed = 0; + slot.stats.n_prompt_cached = n_past; + slot.stats.n_prompt_processed = 0; slot.prompt.tokens.keep_first(n_past); @@ -3491,7 +3481,7 @@ private: continue; } - slot.n_prompt_tokens_processed += n_tokens_out; + slot.stats.n_prompt_processed += n_tokens_out; // add the image chunk to cache { @@ -3530,7 +3520,7 @@ private: slot.need_embd()); slot.prompt.tokens.push_back(cur_tok); - slot.n_prompt_tokens_processed++; + slot.stats.n_prompt_processed++; // break at the last user message, or at user messages at least min step past the last checkpoint if (do_checkpoint && spans.is_user_start(slot.prompt.n_tokens())) { @@ -3583,8 +3573,8 @@ private: // extract the logits only for the last token batch.set_output(batch.size() - 1, true); - slot.n_decoded = 0; - slot.i_batch = batch.size() - 1; + slot.stats.n_predict = 0; + slot.i_batch = batch.size() - 1; slot.init_sampler(); } else { @@ -3826,9 +3816,9 @@ private: // here we have synchronized the llama_context (due to the sampling above), so we can do time measurement const int64_t t_now = ggml_time_us(); - slot.n_decoded += 1; + slot.stats.n_predict += 1; - if (slot.n_decoded == 1) { + if (slot.stats.n_predict == 1) { slot.stats.update_prompt_last(); slot.t_print_last = t_now; slot.n_decoded_last = 0; @@ -3966,7 +3956,7 @@ private: // TODO: set result.probs - slot.n_decoded += 1; + slot.stats.n_predict += 1; if (!process_token(result, slot)) { slot.print_timings();