From bf2c86ddc0685f580595954056c2e77ebabfab4f Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Tue, 14 Jul 2026 18:25:52 +0300 Subject: [PATCH] server : refactor prompt cache state ownership (#25649) * server : clear checkpoints upon prompt clear * server : move the prompt state data to the server_prompt_cache Assisted-by: pi:llama.cpp/Qwen3.6-27B * server : handle batched slot being cleared --- tools/server/server-context.cpp | 71 +++++++++++++++------------------ tools/server/server-task.cpp | 26 ++++++------ tools/server/server-task.h | 53 +++++++++++++----------- 3 files changed, 76 insertions(+), 74 deletions(-) diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 5fadc1c8d9..7564ad4e9c 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -220,8 +220,6 @@ struct server_slot { return false; } - GGML_ASSERT(prompt.data.size() == 0); - const size_t cur_size_tgt = llama_state_seq_get_size_ext(ctx_tgt, id, LLAMA_STATE_SEQ_FLAGS_NONE); const size_t cur_size_dft = ctx_dft ? llama_state_seq_get_size_ext(ctx_dft, id, LLAMA_STATE_SEQ_FLAGS_NONE) : 0; @@ -252,11 +250,7 @@ struct server_slot { return res; } - void prompt_clear(bool allow_processing) { - if (!allow_processing) { - GGML_ASSERT(!is_processing()); - } - + void prompt_clear() { SLT_TRC(*this, "clearing prompt with %zu tokens\n", prompt.tokens.size()); common_context_seq_rm(ctx_tgt, id, -1, -1); @@ -264,7 +258,7 @@ struct server_slot { common_context_seq_rm(ctx_dft, id, -1, -1); } - prompt.tokens.clear(); + prompt.clear(); } std::vector lora; @@ -493,7 +487,7 @@ struct server_slot { // do not keep context of the child slots - the parent's context is enough if (task->is_child()) { - prompt_clear(false); + prompt_clear(); } reset(); @@ -1626,7 +1620,7 @@ private: ret->prompt_save(*prompt_cache); if (!ret->prompt_load(*prompt_cache, task.tokens)) { - ret->prompt_clear(false); + ret->prompt_clear(); } prompt_cache->update(); @@ -1658,7 +1652,7 @@ private: if (slot.prompt.n_tokens() > 0) { SRV_WRN("purging slot %d with %zu tokens\n", slot.id, slot.prompt.tokens.size()); - slot.prompt_clear(false); + slot.prompt_clear(); res = true; @@ -1691,7 +1685,7 @@ private: // if lora has changed, check to see if the cache should be cleared if (lora_should_clear_cache(slot.lora, task_loras)) { SLT_TRC(slot, "clearing cache for lora change. %zu loras -> %zu loras\n", slot.lora.size(), task.params.lora.size()); - slot.prompt.tokens.clear(); + slot.prompt.clear(); } else { SLT_TRC(slot, "keeping cache for alora. %zu target loras\n", task_loras.size()); } @@ -2405,7 +2399,7 @@ private: if (params_base.kv_unified) { // [TAG_IDLE_SLOT_CLEAR] - slot.prompt_clear(false); + slot.prompt_clear(); } } } @@ -2573,12 +2567,12 @@ private: size_t token_count = 0; size_t nread = llama_state_seq_load_file(ctx_tgt, filepath.c_str(), slot->id, tokens.data(), tokens.size(), &token_count); if (nread == 0) { - slot->prompt.tokens.clear(); // KV may already been invalidated? + slot->prompt.clear(); // KV may already been invalidated? send_error(task, "Unable to restore slot, no available space in KV cache or invalid slot save file", ERROR_TYPE_INVALID_REQUEST); break; } tokens.resize(token_count); - slot->prompt.tokens.clear(); + slot->prompt.clear(); slot->prompt.tokens.insert(tokens); const int64_t t_end = ggml_time_us(); @@ -2615,7 +2609,7 @@ private: // Erase token cache const size_t n_erased = slot->prompt.tokens.size(); - slot->prompt_clear(false); + slot->prompt_clear(); auto res = std::make_unique(); res->id = task.id; @@ -2775,6 +2769,27 @@ private: abort_all_slots("pre_decode() failed: " + std::string(e.what())); } + GGML_ASSERT(batch.slot_batched || batch.size() == 0); + + if (batch.slot_batched) { + auto & slot_batched = batch.slot_batched; + auto & alora_scale = batch.alora_scale; + auto & alora_disabled_id = batch.alora_disabled_id; + + // TODO @ngxson : alora handling is too messy, need to refactor it to be more clear and maintainable + // apply lora, only need to do it once per batch + common_set_adapter_lora(ctx_tgt, slot_batched->lora); + + // if the lora is temporarily disabled for an alora, re-enable it + // for next time + if (alora_scale > 0.0f) { + SRV_DBG("re-enabling alora with scale %f\n", alora_scale); + slot_batched->lora[alora_disabled_id].scale = alora_scale; + } + + llama_set_embeddings(ctx_tgt, slot_batched->need_embd()); + } + llama_batch batch_view; int32_t off_next = 0; int32_t n_batch = llama_n_batch(ctx_tgt); @@ -2814,7 +2829,6 @@ private: abort_all_slots("post_decode() failed: " + std::string(e.what())); break; // stop any further processing } - } } @@ -2880,7 +2894,7 @@ private: new_tokens.resize(slot.prompt.tokens.size() - n_discard); - slot.prompt.tokens.clear(); + slot.prompt.clear(); slot.prompt.tokens.insert(new_tokens); } @@ -3556,25 +3570,6 @@ private: bool decode(int32_t & n_batch, int32_t off, llama_batch & batch_view) { SRV_DBG("n_batch (effective) = %d, off = %d\n", n_batch, off); - auto & slot_batched = batch.slot_batched; - auto & alora_scale = batch.alora_scale; - auto & alora_disabled_id = batch.alora_disabled_id; - - // TODO @ngxson : alora handling is too messy, need to refactor it to be more clear and maintainable - if (slot_batched) { - // apply lora, only need to do it once per batch - common_set_adapter_lora(ctx_tgt, slot_batched->lora); - - // if the lora is temporarily disabled for an alora, re-enable it - // for next time - if (alora_scale > 0.0f) { - SRV_DBG("re-enabling alora with scale %f\n", alora_scale); - slot_batched->lora[alora_disabled_id].scale = alora_scale; - } - - llama_set_embeddings(ctx_tgt, slot_batched->need_embd()); - } - if (batch.size() == 0) { SRV_WRN("%s", "no tokens to decode\n"); @@ -3622,7 +3617,7 @@ private: // note: it's complicated to keep track of how much of the current batch has been // processed before the error occurred, so we simply clear the entire context - slot.prompt_clear(false); + slot.prompt_clear(); } } diff --git a/tools/server/server-task.cpp b/tools/server/server-task.cpp index 8d611e520d..f01780100c 100644 --- a/tools/server/server-task.cpp +++ b/tools/server/server-task.cpp @@ -1646,16 +1646,16 @@ size_t server_prompt_cache::n_tokens() const { size_t res = 0; for (const auto & state : states) { - res += state.n_tokens(); + res += state.prompt.n_tokens(); } return res; } -server_prompt * server_prompt_cache::alloc(const server_prompt & prompt, size_t state_size_tgt, size_t state_size_dft) { +server_prompt_cache_state * server_prompt_cache::alloc(const server_prompt & prompt, size_t state_size_tgt, size_t state_size_dft) { // first check if the current state is contained fully in the cache for (auto it = states.begin(); it != states.end(); ++it) { - const int cur_lcp_len = it->tokens.get_common_prefix(prompt.tokens); + const int cur_lcp_len = it->prompt.tokens.get_common_prefix(prompt.tokens); if (cur_lcp_len == (int) prompt.tokens.size()) { SRV_TRC("%s", " - prompt is already in the cache, skipping\n"); @@ -1680,9 +1680,9 @@ server_prompt * server_prompt_cache::alloc(const server_prompt & prompt, size_t // remove any cached prompts that are fully contained in the current prompt for (auto it = states.begin(); it != states.end();) { - const int len = it->tokens.get_common_prefix(prompt.tokens); + const int len = it->prompt.tokens.get_common_prefix(prompt.tokens); - if (len == (int) it->tokens.size()) { + if (len == (int) it->prompt.tokens.size()) { SRV_TRC(" - removing obsolete cached prompt with length %d\n", len); it = states.erase(it); @@ -1721,12 +1721,14 @@ server_prompt * server_prompt_cache::alloc(const server_prompt & prompt, size_t } states.push_back({ - /*.tokens =*/ prompt.tokens.clone(), - /*.data =*/ { + /*.prompt =*/ { + /*.tokens =*/ prompt.tokens.clone(), + /*.checkpoints =*/ prompt.checkpoints, + }, + /*.data =*/ { /*.main =*/ std::move(state_data_tgt), /*.drft =*/ std::move(state_data_dft), }, - /*.checkpoints =*/ prompt.checkpoints, }); return &states.back(); @@ -1744,9 +1746,9 @@ bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tok // find the most similar cached prompt, that would also preserve the most context for (auto it = states.begin(); it != states.end(); ++it) { - const int lcp_cur = it->tokens.get_common_prefix(tokens_new); + const int lcp_cur = it->prompt.tokens.get_common_prefix(tokens_new); - const float f_keep_cur = float(lcp_cur) / it->tokens.size(); + const float f_keep_cur = float(lcp_cur) / it->prompt.tokens.size(); const float sim_cur = float(lcp_cur) / tokens_new.size(); // don't trash large prompts @@ -1799,7 +1801,7 @@ bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tok } } - prompt = std::move(*it_best); + prompt = std::move(it_best->prompt); states.erase(it_best); } @@ -1836,6 +1838,6 @@ void server_prompt_cache::update() { for (const auto & state : states) { SRV_TRC(" - prompt %p: %7d tokens, checkpoints: %2zu, %9.3f MiB\n", - (const void *)&state, state.n_tokens(), state.checkpoints.size(), state.size() / (1024.0 * 1024.0)); + (const void *)&state, state.prompt.n_tokens(), state.prompt.checkpoints.size(), state.size() / (1024.0 * 1024.0)); } } diff --git a/tools/server/server-task.h b/tools/server/server-task.h index dc6b2dac1e..c3eea2ecb8 100644 --- a/tools/server/server-task.h +++ b/tools/server/server-task.h @@ -584,32 +584,14 @@ struct server_task_result_apply_lora : server_task_result { virtual json to_json() override; }; -struct server_prompt_data { - std::vector main; - std::vector drft; - - size_t size() const { - return main.size() + drft.size(); - } -}; - struct server_prompt { server_tokens tokens; - server_prompt_data data; - std::list checkpoints; - size_t size() const { - size_t res = 0; - - res += data.size(); - - for (const auto & ckpt : checkpoints) { - res += ckpt.size(); - } - - return res; + void clear() { + tokens.clear(); + checkpoints.clear(); } int n_tokens() const { @@ -619,19 +601,42 @@ struct server_prompt { server_prompt clone() const { return server_prompt { tokens.clone(), - data, checkpoints, }; } }; +struct server_prompt_data { + std::vector main; + std::vector drft; + + size_t size() const { + return main.size() + drft.size(); + } +}; + +struct server_prompt_cache_state { + server_prompt prompt; + server_prompt_data data; + + size_t size() const { + size_t res = data.size(); + + for (const auto & ckpt : prompt.checkpoints) { + res += ckpt.size(); + } + + return res; + } +}; + struct server_prompt_cache { server_prompt_cache(int32_t limit_size_mib, size_t limit_tokens) { this->limit_size = 1024ull*1024ull*(limit_size_mib < 0 ? 0 : limit_size_mib); this->limit_tokens = limit_tokens; } - std::list states; + std::list states; // in bytes, 0 = no limit size_t limit_size = 0; @@ -643,7 +648,7 @@ struct server_prompt_cache { size_t n_tokens() const; - server_prompt * alloc(const server_prompt & prompt, size_t state_size_main, size_t state_size_drft); + server_prompt_cache_state * alloc(const server_prompt & prompt, size_t state_size_main, size_t state_size_drft); bool load(server_prompt & prompt, const server_tokens & tokens_new, llama_context * ctx_main, llama_context * ctx_drft, int32_t id_slot);